Lightning-AI / Lightning-AI/litData

`StreamingDataLoader` state is not loaded from checkpoint when resuming training

Open
#775 1 comment 1 reaction 0 assignees View on GitHub
bug help wanted
Dominant language
Python
Stars
614
Forks
106
Avg merge
15h 8m
Merged PRs (30d)
22

Description

## 🐛 Bug

The answer given in #249 suggests that `Trainer` seamlessly integrates with `StreamingDataLoader`, such that when saving checkpoints, the state of the `StreamingDataLoader` is included, and when resuming training, the `StreamingDataLoader` state is loaded. However this does not seem to be the case. See example code below.

If this is expected, i.e. if users are required to implement extra logic in order to achieve this behavior, then an example of how to achieve this behavior should be added in the README, as suggested in the original issue.

### To Reproduce

```python
import torch
import torch.nn as nn
from lightning import LightningDataModule, LightningModule, Trainer
from litdata import StreamingDataLoader, StreamingDataset
from litdata.streaming import Cache

class MyTrainingException(Exception):
pass

class MyStreamingDataLoader(StreamingDataLoader):
def load_state_dict(self, obj):
raise ValueError("StreamingDataLoader.load_state_dict called! Hooray!")

class MyLightningDataModule(LightningDataModule):
def prepare_data(self):
cache = Cache("temp/", chunk_size=1)
dset_len = 10
for i in range(dset_len):
cache[i] = i
cache.done()
cache.merge()

def setup(self, stage):
self.dset = StreamingDataset("temp/")

def train_dataloader(self):
return MyStreamingDataLoader(self.dset)

class MyLightningModule(LightningModule):
def __init__(self):
super().__init__()
self.param = nn.Parameter(torch.randn(1))

def training_step(self, batch):
if self.current_epoch == 2:
raise MyTrainingException
return batch * self.param

def configure_optimizers(self):
return torch.optim.Adam(self.parameters())

dm = MyLightningDataModule()
model = MyLightningModule()
trainer = Trainer(max_epochs=10, default_root_dir="temp/")
try:
trainer.fit(model, dm)
except MyTrainingException:
print("Training crashed as expected on epoch 2.")
# resume training
dm = MyLightningDataModule()
model = MyLightningModule()
trainer = Trainer(max_epochs=10, default_root_dir="temp/")
# the call below should raise ValueError("StreamingDataLoader.load_state_dict called! Hooray!")
# but it raises MyTrainingException again, which means the dataloader state is not loaded from the checkpoint
trainer.fit(model, dm, ckpt_path="temp/lightning_logs/version_0/checkpoints/epoch=1-step=20.ckpt")
```

Contributor guide

Open the contributing guide

Research direction

Start by running the supplied reproducer and tracing Trainer.fit checkpoint restoration to StreamingDataLoader.load_state_dict. Done means the overridden method is called when resuming from the shown checkpoint; if that behavior is not expected, document the required integration in the README instead.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
data-engineering, machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
43/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.