Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
5 changes: 4 additions & 1 deletion src/litdata/utilities/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,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 +115,9 @@ 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"]
from copy import deepcopy
Comment thread
deependujha marked this conversation as resolved.
Outdated

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]}