From be7f7d4e680cc0cb70278a7deb5d346c572647c3 Mon Sep 17 00:00:00 2001 From: aswanth-07 Date: Sun, 23 Aug 2026 18:40:50 +0530 Subject: [PATCH 1/2] Fix Loader restoration before None items --- test/nodes/test_loader.py | 11 +++++++++++ torchdata/nodes/loader.py | 10 +++++++--- 2 files changed, 18 insertions(+), 3 deletions(-) diff --git a/test/nodes/test_loader.py b/test/nodes/test_loader.py index 774ec3297..69c0253cb 100644 --- a/test/nodes/test_loader.py +++ b/test/nodes/test_loader.py @@ -53,3 +53,14 @@ 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]) diff --git a/torchdata/nodes/loader.py b/torchdata/nodes/loader.py index b934feea8..52b0f3282 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,25 @@ 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 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) From b88467a9953897302059000bb8edd06fe0a36086 Mon Sep 17 00:00:00 2001 From: aswanth-07 Date: Fri, 4 Sep 2026 01:55:55 +0530 Subject: [PATCH 2/2] Clear Loader cached state on reset --- test/nodes/test_loader.py | 12 ++++++++++++ torchdata/nodes/loader.py | 1 + 2 files changed, 13 insertions(+) diff --git a/test/nodes/test_loader.py b/test/nodes/test_loader.py index 69c0253cb..74255c058 100644 --- a/test/nodes/test_loader.py +++ b/test/nodes/test_loader.py @@ -64,3 +64,15 @@ def test_loader_restores_before_none_item(self) -> None: 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 52b0f3282..9ffa210fa 100644 --- a/torchdata/nodes/loader.py +++ b/torchdata/nodes/loader.py @@ -112,6 +112,7 @@ def reset(self, initial_state: Optional[Dict[str, Any]] = 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 not self._has_cached_item: