diff --git a/test/nodes/test_loader.py b/test/nodes/test_loader.py index 774ec3297..74255c058 100644 --- a/test/nodes/test_loader.py +++ b/test/nodes/test_loader.py @@ -53,3 +53,26 @@ def test_loader_equal_state_dict_on_save_load_map(self) -> None: length = 10 node = MapStyleWrapper(DummyMapDataset(length), sampler=range(length)) self._test_loader_correct_state_dict_at_midpoint(node, length) + + def test_loader_restores_before_none_item(self) -> None: + loader = Loader(IterableWrapper([0, None, 2])) + iterator = iter(loader) + self.assertEqual(next(iterator), 0) + state_dict = loader.state_dict() + + restored = Loader(IterableWrapper([0, None, 2])) + restored.load_state_dict(state_dict) + + self.assertEqual(list(restored), [None, 2]) + + def test_loader_reset_clears_cached_state_dict(self) -> None: + checkpoint_source = Loader(IterableWrapper([0, 1, 2])) + self.assertEqual(next(iter(checkpoint_source)), 0) + state_dict = checkpoint_source.state_dict() + + loader = Loader(IterableWrapper([0, 1, 2])) + self.assertTrue(iter(loader).has_next()) + loader.load_state_dict(state_dict) + iter(loader) + + self.assertEqual(loader.state_dict(), state_dict) diff --git a/torchdata/nodes/loader.py b/torchdata/nodes/loader.py index b934feea8..9ffa210fa 100644 --- a/torchdata/nodes/loader.py +++ b/torchdata/nodes/loader.py @@ -98,6 +98,7 @@ def __init__( self.loader = loader self.root = loader.root self._cached_item = None + self._has_cached_item = False self._cached_state_dict: Optional[Dict[str, Any]] = None self._num_yielded = 0 @@ -110,22 +111,26 @@ def reset(self, initial_state: Optional[Dict[str, Any]] = None): self.root.reset(None) self._num_yielded = 0 self._cached_item = None + self._has_cached_item = False + self._cached_state_dict = None def has_next(self) -> bool: - if self._cached_item is None: + if not self._has_cached_item: try: # Cache the current state dict self._cached_state_dict = self.state_dict() # Load and save the next item self._cached_item = next(self) + self._has_cached_item = True except StopIteration: pass - return self._cached_item is not None + return self._has_cached_item def next(self): - if self._cached_item is not None: + if self._has_cached_item: item = self._cached_item self._cached_item = None + self._has_cached_item = False self._cached_state_dict = None else: item = next(self.root)