Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/litdata/streaming/combined.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
13 changes: 10 additions & 3 deletions src/litdata/streaming/parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
4 changes: 3 additions & 1 deletion src/litdata/utilities/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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
Expand Down
63 changes: 63 additions & 0 deletions tests/streaming/test_combined.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:]
30 changes: 30 additions & 0 deletions tests/streaming/test_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]}
Loading