diff --git a/src/litdata/streaming/combined.py b/src/litdata/streaming/combined.py index c2882958d..044b8923f 100644 --- a/src/litdata/streaming/combined.py +++ b/src/litdata/streaming/combined.py @@ -191,7 +191,7 @@ def __init__( self._is_done = False if num_samples_yielded is not None: - self._num_samples_yielded = num_samples_yielded + self._num_samples_yielded = deepcopy(num_samples_yielded) for _ in range(sum(num_samples_yielded)): choice_indexes: list[int] = [index for index in self._dataset_indexes if index is not None] choice_weights: list[float] = [w for w in self._weights if w is not None] diff --git a/src/litdata/streaming/parallel.py b/src/litdata/streaming/parallel.py index 435e43bd3..ef261e11d 100644 --- a/src/litdata/streaming/parallel.py +++ b/src/litdata/streaming/parallel.py @@ -291,10 +291,17 @@ def get_num_samples_yielded( output[i] = sum(s for (s, c) in zip(num_samples_yielded, num_cycles) if c == cycles[i]) return output, cycles + def reset_state_dict(self) -> None: + """Reset the state of the dataset.""" + super().reset_state_dict() + self._num_cycles = None + def load_state_dict(self, state_dict: dict[str, Any]) -> None: + if not state_dict: + return super().load_state_dict(state_dict) if self._use_streaming_dataloader: - self._num_cycles = state_dict["num_cycles"] + self._num_cycles = deepcopy(state_dict["num_cycles"]) def state_dict( self, num_workers: int, batch_size: int, num_samples_yielded: list[int] | None = None @@ -327,8 +334,8 @@ def __init__( ) -> None: self._datasets = datasets self._dataset_iters = [iter(dataset) for dataset in datasets] - self._num_samples_yielded = num_samples_yielded or [0 for _ in range(len(datasets))] - self._num_cycles = num_cycles or [0 for _ in range(len(datasets))] + self._num_samples_yielded = deepcopy(num_samples_yielded) or [0 for _ in range(len(datasets))] + self._num_cycles = deepcopy(num_cycles) or [0 for _ in range(len(datasets))] self._length = length self._use_streaming_dataloader = use_streaming_dataloader self._transform = transform diff --git a/src/litdata/utilities/base.py b/src/litdata/utilities/base.py index 7e24ece2d..11285c8fe 100644 --- a/src/litdata/utilities/base.py +++ b/src/litdata/utilities/base.py @@ -13,6 +13,7 @@ from abc import ABC, abstractmethod from collections.abc import Iterator, Sequence +from copy import deepcopy from typing import Any from torch.utils.data import IterableDataset @@ -80,6 +81,7 @@ def set_drop_last(self, drop_last: bool) -> None: def reset_state_dict(self) -> None: """Reset the state of the dataset.""" + self._num_samples_yielded = None for dataset in self._datasets: dataset.reset_state_dict() @@ -114,7 +116,7 @@ def load_state_dict(self, state_dict: dict[str, Any]) -> None: # Used to iterate over the sampler to avoid sampling the same samples if self._use_streaming_dataloader: - self._num_samples_yielded = state_dict["num_samples_yielded"] + self._num_samples_yielded = deepcopy(state_dict["num_samples_yielded"]) def _get_len(self, d: Any) -> int: # mypy: ``self.batch_size`` can be a ``Sequence[int]`` now, but the diff --git a/tests/streaming/test_combined.py b/tests/streaming/test_combined.py index dc2b970e1..588be1f36 100644 --- a/tests/streaming/test_combined.py +++ b/tests/streaming/test_combined.py @@ -703,3 +703,66 @@ def test_combined_rejects_topology_change_on_resume(tmpdir): loader_b = StreamingDataLoader(dataset_b, num_workers=4, batch_size=2) with pytest.raises(ValueError, match="support resume only"): loader_b.load_state_dict(state) + + +def test_combined_dataset_reset_state_dict_after_checkpoint_resume(tmpdir): + data_dir_1 = os.path.join(tmpdir, "data_1") + data_dir_2 = os.path.join(tmpdir, "data_2") + os.makedirs(data_dir_1) + os.makedirs(data_dir_2) + for path, count, offset in ((data_dir_1, 10, 0), (data_dir_2, 12, 100)): + cache = Cache(input_dir=path, chunk_size=2) + for i in range(count): + cache[i] = i + offset + cache.done() + cache.merge() + + def make_loader(): + dataset = CombinedStreamingDataset( + datasets=[ + StreamingDataset(input_dir=data_dir_1, shuffle=True), + StreamingDataset(input_dir=data_dir_2, shuffle=True), + ], + seed=42, + iterate_over_all=True, + ) + return StreamingDataLoader(dataset, num_workers=0, batch_size=2) + + # Baseline: run 2 full epochs without checkpoint resume + ref_loader = make_loader() + list(ref_loader) + ref_epoch_2 = [batch.tolist() for batch in ref_loader] + + # Resumed: break mid-epoch 1, restore, finish epoch 1, then run epoch 2 + loader = make_loader() + ckpt = None + for idx, _ in enumerate(loader): + if idx == 2: + ckpt = deepcopy(loader.state_dict()) + break + + assert ckpt is not None + loader.load_state_dict(ckpt) + assert loader.restore + + # Finish epoch 1 + list(loader) + assert not loader.restore + + # Epoch 2 should reset _num_samples_yielded and match the baseline epoch 2 batches and state + resumed_epoch_2 = [] + epoch_2_mid_ckpt = None + for idx, batch in enumerate(loader): + resumed_epoch_2.append(batch.tolist()) + if idx == 2: + epoch_2_mid_ckpt = deepcopy(loader.state_dict()) + + assert resumed_epoch_2 == ref_epoch_2 + assert sum(loader.state_dict()["num_samples_yielded"][0]) == len(loader.dataset) + + # A mid-epoch checkpoint saved in epoch 2 must still be restorable + assert epoch_2_mid_ckpt is not None + loader.load_state_dict(epoch_2_mid_ckpt) + assert loader.restore + resumed_epoch_2_tail = [batch.tolist() for batch in loader] + assert resumed_epoch_2_tail == ref_epoch_2[3:] diff --git a/tests/streaming/test_parallel.py b/tests/streaming/test_parallel.py index 5b03f8480..72036a75e 100644 --- a/tests/streaming/test_parallel.py +++ b/tests/streaming/test_parallel.py @@ -1193,3 +1193,33 @@ def test_parallel_dataset_complete_iteration_resume_without_dataloader(tmp_path_ assert all(x == y for x, y in zip(sample, expected[i])) elif not resume and length is not None: assert all(x == y for x, y in zip(sample, samples[i])) + + +def test_parallel_dataset_reset_state_dict_after_checkpoint_resume(tmp_path_factory): + _, _, pardset, dataloader, _ = prepare_parallel_dataset_and_dataloder( + tmp_path_factory, parlen=16, len1=10, len2=12, batch_size=2, num_workers=0, shuffle=False, resume=False + ) + + ckpt = None + for idx, _ in enumerate(dataloader): + if idx == 2: + ckpt = deepcopy(dataloader.state_dict()) + break + + assert ckpt is not None + dataloader.load_state_dict(ckpt) + assert dataloader.restore + + # Finish epoch 1 + remaining_epoch_1 = list(dataloader) + assert len(remaining_epoch_1) == 5 + assert not dataloader.restore + + # Epoch 2 must reset _num_samples_yielded and _num_cycles on the dataset + epoch_2_batches = list(dataloader) + assert len(epoch_2_batches) == 8 + assert pardset._num_samples_yielded is None + assert pardset._num_cycles is None + state_epoch_2 = dataloader.state_dict() + assert state_epoch_2["num_samples_yielded"] == {0: [6, 4]} + assert state_epoch_2["num_cycles"] == {0: [1, 1]}