Lightning-AI / Lightning-AI/pytorch-lightning

Dataloader reload bug when loading from checkpoint

Open
#21,492 4 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

bug checkpointing lightningdatamodule ver: 2.5.x
Dominant language
Python
Stars
31.4k
Forks
3.8k
Avg merge
6d 7h
Merged PRs (30d)
6

Description

### Bug description

When loading from a checkpoint the setup_data function is called causing self._last_train_dl_reload_epoch to be updated and later the dataloader is not updated as expected causing reproducibility issues in training.

__________________________________________________________________________________________

I have trained a model for 50 epochs and checkpointed it.
I have a reload option every 50 epochs.
When loading the model the setup_data function is called and the dataloader is reset on epoch 49 (although normally it would be reset on the 51st epoch, so epoch 50). So when epoch 50 starts the dataloader is not reloaded as only one epoch has passed.

### What version are you seeing the problem on?

v2.5

### Reproduced in studio

_No response_

### How to reproduce the bug

```python
import torch
from lightning import Trainer, LightningModule
from torch.utils.data import Dataset, DataLoader

CHECKPOINT_PATH = "stage1.ckpt"
DEBUG = False

class FakeDataset(Dataset):
def __init__(self):
self.data = [torch.zeros(3) for _ in range(10)]
self.labels = [torch.zeros(1) for _ in range(10)]

def __len__(self):
return len(self.data)

def __getitem__(self, idx):
return self.data[idx], self.labels[idx]

class FakeDataset2(FakeDataset):
def __init__(self):
super().__init__()
self.data = [torch.ones(3) for _ in range(20)]
self.labels = [torch.ones(1) for _ in range(20)]

class SimpleModule(LightningModule):
def __init__(self):
super().__init__()
self.layer = torch.nn.Linear(3, 1)
self.first_stage_epochs: int = 2
self.second_stage_epochs: int = 10

def training_step(self, batch, batch_idx):
x, y = batch
if self.current_epoch < self.first_stage_epochs:
assert torch.all(x == 0), "Data in first stage should be zeros"
else:
assert torch.all(x == 1), "Data in second stage should be ones"
y_hat = self.layer(x)
loss = torch.nn.functional.mse_loss(y_hat, y)
self.log("train_loss", loss)
return loss

def configure_optimizers(self):
return torch.optim.SGD(self.layer.parameters(), lr=0.01)

def train_dataloader(self):
if self.current_epoch < self.first_stage_epochs:
dataset = FakeDataset()
else:
dataset = FakeDataset2()
return DataLoader(dataset, batch_size=1)

def on_train_epoch_end(self):
if self.current_epoch == self.first_stage_epochs - 1:
print("Completed first stage")
self.trainer.save_checkpoint(CHECKPOINT_PATH)

def on_train_start(self):
if self.current_epoch > 0 and DEBUG:
print("Resumed training for second stage")
# Fit Loop is calling the setup_data() during checkpoint restoration
# This is setting _last_train_dl_reload_epoch to current epoch - 1
self.trainer.fit_loop._last_train_dl_reload_epoch = 0

if __name__ == "__main__":
# Uncomment the following line for a quick fix
# DEBUG = True
dataset = FakeDataset()
model = SimpleModule()

trainer = Trainer(
max_epochs=model.first_stage_epochs + model.second_stage_epochs,
accelerator="gpu",
devices=1,
log_every_n_steps=1,
reload_dataloaders_every_n_epochs=model.first_stage_epochs,
)

#####################################################
# Train the first time both 1st stage and 2nd stage
#####################################################
trainer.fit(model)
print("Successfully trained first stage and second stage")
#####################################################
# Resume training from checkpoint and only train second stage
#####################################################
trainer.fit(model, ckpt_path=CHECKPOINT_PATH)

```

### Error messages and logs

```
# Error messages and logs here please
```

### Environment

Current environment

```
#- PyTorch Lightning Version (e.g., 2.5.0):
#- PyTorch Version (e.g., 2.5):
#- Python version (e.g., 3.12):
#- OS (e.g., Linux):
#- CUDA/cuDNN version:
#- GPU models and configuration:
#- How you installed Lightning(`conda`, `pip`, source):
```

### More info

_No response_

cc @ethanwharris @lantiga

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by tracing Trainer.fit checkpoint restoration through the Fit Loop's setup_data() and _last_train_dl_reload_epoch handling. Run the supplied reproduction script with reload_dataloaders_every_n_epochs set to first_stage_epochs and compare uninterrupted versus resumed training. Done means the resumed run reloads the dataloader at the same epoch as the uninterrupted run.

Written by the indexing model from the issue text.

Assessment

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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.