Lightning-AI / Lightning-AI/pytorch-lightning

`num_training_batches` is `inf` in `configure_optimizers`

Open
#16,060 6 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

bug data handling loops
Dominant language
Python
Stars
31.4k
Forks
3.8k
Avg merge
6d 7h
Merged PRs (30d)
6

Description

### Bug description

The value of `num_training_batches` is `inf` when referenced in `configure_optimizers()`. It seems that it doesn't actually get its correct value until some point later. This causes a very hard-to-find issue because the training runs without error, except the loss is `nan`.

Something inside `optim.lr_scheduler.CyclicLR` actually sets the `lr` of the `optimizer` to `nan`.

It would be nice if:
* This value was available `configure_optimizers()` was called, or
* There was a warning if accessing it before it's set

### How to reproduce the bug

```python
import os

import torch
from torch.utils.data import DataLoader, Dataset

from pytorch_lightning import LightningModule, Trainer

class RandomDataset(Dataset):
def __init__(self, size, length):
self.len = length
self.data = torch.randn(length, size)

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

def __len__(self):
return self.len

class BoringModel(LightningModule):
def __init__(self):
super().__init__()
self.layer = torch.nn.Linear(32, 2)

def forward(self, x):
return self.layer(x)

def training_step(self, batch, batch_idx):
loss = self(batch).sum()
self.log("train_loss", loss)
return {"loss": loss}

def validation_step(self, batch, batch_idx):
loss = self(batch).sum()
self.log("valid_loss", loss)

def test_step(self, batch, batch_idx):
loss = self(batch).sum()
self.log("test_loss", loss)

def configure_optimizers(self):
optimizer = torch.optim.SGD(self.layer.parameters(), lr=0.1)
print(f"{optimizer.param_groups[0]['lr'] = }") # 0.1
lr_scheduler = torch.optim.lr_scheduler.CyclicLR(
optimizer=optimizer,
base_lr=0.01,
max_lr=0.1,
step_size_up=self.trainer.num_training_batches * 1, # problematic!
step_size_down=self.trainer.num_training_batches * 2, # problematic!
cycle_momentum=False,
)
print(f"{optimizer.param_groups[0]['lr'] = }") # nan
return [optimizer], [lr_scheduler]

def run():
train_data = DataLoader(RandomDataset(32, 64), batch_size=2)
val_data = DataLoader(RandomDataset(32, 64), batch_size=2)
test_data = DataLoader(RandomDataset(32, 64), batch_size=2)

model = BoringModel()
trainer = Trainer(
default_root_dir=os.getcwd(),
limit_train_batches=1,
limit_val_batches=1,
limit_test_batches=1,
num_sanity_val_steps=0,
max_epochs=1,
enable_model_summary=False,
enable_checkpointing=False,
)
trainer.fit(model, train_dataloaders=train_data, val_dataloaders=val_data)
trainer.test(model, dataloaders=test_data)

if __name__ == "__main__":
run()
```

### Error messages and logs

The main hint something is wrong is actually tensorboard printing "NaN or Inf found in input tensor" - but even that doesn't come with a trace telling me who's printing this.

### Environment

Current environment

```
#- Lightning Component (e.g. Trainer, LightningModule, LightningApp, LightningWork, LightningFlow):
#- PyTorch Lightning Version (e.g., 1.5.0):
#- Lightning App Version (e.g., 0.5.2):
#- PyTorch Version (e.g., 1.10):
#- Python version (e.g., 3.9):
#- OS (e.g., Linux):
#- CUDA/cuDNN version:
#- GPU models and configuration:
#- How you installed Lightning(`conda`, `pip`, source):
#- Running environment of LightningApp (e.g. local, cloud):
```

### More info

_No response_

cc @justusschock @awaelchli @carmocca

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 running the supplied reproduction and tracing when Trainer.num_training_batches is initialized relative to configure_optimizers(). Inspect the training setup path and add a regression test for the early access behavior. Done means the value is available there or access produces a clear warning, without the scheduler receiving inf and producing NaN learning rates.

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.