Lightning-AI / Lightning-AI/pytorch-lightning

incorrect global_step with multiple optimizers and automatic_optimization=False

Open
#17,958 25 comments 18 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

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

Description

### Bug description

Hello,

I encountered a bug when training with `automatic_optimization = False` and two optimizers.

In summary: the `global_step` attribute of the trainer and the lightning module is tracking the total number of calls to `optimizer.step()` (in my case, two per `training_step`), rather than the total number of iterations of the dataloader.

This conflicts with the notion of `step` in arguments like `log_every_n_steps` and `val_check_interval` in the trainer. Case in point, if we call
```python
self.log("global_step", self.global_step)
```
inside `training_step`, with `CSVLogger`, `log_every_n_steps=10`, and two `optimizer.step()`s per `training_step`, the CSV logs show:
```
global_step,epoch,step
20.0,0,9
40.0,0,19
60.0,0,29
80.0,0,39
100.0,0,49
```
Note how `global_step` conflicts with `step`, and in fact is twice the expected value, since we have two optimizers.

I have attached a complete code example that replicates the issue.

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

v2.0

### How to reproduce the bug

```python
import pytorch_lightning as pl
import torch

from pytorch_lightning.loggers import CSVLogger
from torch.utils.data import TensorDataset, IterableDataset, DataLoader

SEMVER = tuple(int(x) for x in pl.__version__.split("."))
assert SEMVER >= (2, 0, 3)

class LinearRegression(pl.LightningModule):

def __init__(self):
super().__init__()
self.gamma = torch.nn.Parameter(torch.ones(()))
self.beta = torch.nn.Parameter(torch.zeros(()))
self.automatic_optimization = False

def forward(self, x):
return self.gamma * x + self.beta

def configure_optimizers(self):
gamma_opt = torch.optim.SGD([self.gamma], lr=1e-2)
beta_opt = torch.optim.SGD([self.beta], lr=1e-2)
return gamma_opt, beta_opt

def training_step(self, batch, batch_idx):

# Two optimizers.
gamma_opt, beta_opt = self.optimizers()

# Forward pass with loss.
inputs, targets = batch
predictions = self(inputs)
loss = torch.nn.functional.mse_loss(predictions, targets)

# Backprop through entire graph but only update gamma.
gamma_opt.zero_grad()
self.manual_backward(loss, retain_graph=True)
gamma_opt.step()

# Backprop through partial graph and only update beta.
beta_opt.zero_grad()
self.manual_backward(loss, inputs=[self.beta])
beta_opt.step()

# Log the global step.
self.log("global_step_train", self.global_step)

def validation_step(self, batch, batch_idx):

# Forward pass with loss.
inputs, targets = batch
predictions = self(inputs)
loss = torch.nn.functional.mse_loss(predictions, targets)

# Log the global step.
self.log("global_step_val", self.global_step)

class IterableTensorDataset(IterableDataset):

def __init__(self, inputs, targets):
self.inputs, self.targets = inputs, targets

def __iter__(self):
while True:
i = torch.randint(self.inputs.shape[0], size=())
yield self.inputs[i], self.targets[i]

def load_dataset(gamma=2.0, beta=-1.0, sigma=0.2):
inputs = torch.linspace(-1, 1, 201)
targets = gamma * inputs + beta
targets += sigma * torch.randn_like(targets)
indices = torch.randperm(inputs.shape[0])
pivot = inputs.shape[0] // 2
train_inds, val_inds = indices[:pivot], indices[pivot:]
return (
(inputs[train_inds], targets[train_inds]),
(inputs[val_inds], targets[val_inds]))

def main():
train_data, val_data = load_dataset()
train_set = IterableTensorDataset(*train_data)
val_set = TensorDataset(*val_data)
train_loader = DataLoader(train_set, batch_size=4)
test_loader = DataLoader(val_set, batch_size=1)
trainer = pl.Trainer(
log_every_n_steps=10,
val_check_interval=100,
logger=CSVLogger("./logs"),
enable_progress_bar=False,
max_steps=1000,
)
model = LinearRegression()
trainer.fit(model, train_loader, test_loader)
return (
model.gamma.data.detach().cpu().item(),
model.beta.data.detach().cpu().item())

if __name__ == "__main__":
print(main())
```

### Error messages and logs

```
global_step_train,epoch,step,global_step_val
20.0,0,9,
40.0,0,19,
60.0,0,29,
80.0,0,39,
100.0,0,49,
120.0,0,59,
140.0,0,69,
160.0,0,79,
180.0,0,89,
200.0,0,99,
,0,99,200.0
220.0,0,109,
240.0,0,119,
260.0,0,129,
280.0,0,139,
300.0,0,149,
320.0,0,159,
340.0,0,169,
360.0,0,179,
380.0,0,189,
400.0,0,199,
,0,199,400.0
420.0,0,209,
440.0,0,219,
460.0,0,229,
480.0,0,239,
500.0,0,249,
520.0,0,259,
540.0,0,269,
560.0,0,279,
580.0,0,289,
600.0,0,299,
,0,299,600.0
620.0,0,309,
640.0,0,319,
660.0,0,329,
680.0,0,339,
700.0,0,349,
720.0,0,359,
740.0,0,369,
760.0,0,379,
780.0,0,389,
800.0,0,399,
,0,399,800.0
820.0,0,409,
840.0,0,419,
860.0,0,429,
880.0,0,439,
900.0,0,449,
920.0,0,459,
940.0,0,469,
960.0,0,479,
980.0,0,489,
1000.0,0,499,
,0,499,1000.0
```

### Environment

Current environment

* CUDA:
- GPU: None
- available: False
- version: None
* Lightning:
- lightning-utilities: 0.9.0
- pytorch-lightning: 2.0.4
- torch: 2.0.1
- torchmetrics: 0.11.4
* Packages:
- aiohttp: 3.8.4
- aiosignal: 1.3.1
- async-timeout: 4.0.2
- attrs: 23.1.0
- certifi: 2023.5.7
- charset-normalizer: 3.1.0
- filelock: 3.12.2
- frozenlist: 1.3.3
- fsspec: 2023.6.0
- idna: 3.4
- jinja2: 3.1.2
- lightning-utilities: 0.9.0
- markupsafe: 2.1.3
- mpmath: 1.3.0
- multidict: 6.0.4
- networkx: 3.1
- numpy: 1.25.0
- packaging: 23.1
- pip: 23.0.1
- pytorch-lightning: 2.0.4
- pyyaml: 6.0
- requests: 2.31.0
- setuptools: 67.6.0
- sympy: 1.12
- torch: 2.0.1
- torchmetrics: 0.11.4
- tqdm: 4.65.0
- typing-extensions: 4.7.0
- urllib3: 2.0.3
- wheel: 0.38.4
- yarl: 1.9.2
* System:
- OS: Darwin
- architecture:
- 64bit
-
- processor: i386
- python: 3.10.11
- release: 20.6.0
- version: Darwin Kernel Version 20.6.0: Thu Mar 9 20:39:26 PST 2023; root:xnu-7195.141.49.700.6~1/RELEASE_X86_64

### More info

If this is the intended behavior, it should be reconciled with the trainer's notion of step. Arguments like `log_every_n_steps` and `val_check_interval` use a different definition of step.

cc @borda

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 provided two-optimizer reproduction with PyTorch Lightning 2.0.4 and inspect how global_step changes during training_step, optimizer.step(), logging, and validation. Compare it with the trainer's log_every_n_steps and val_check_interval behavior; done means global_step has one consistent interpretation across the trainer, LightningModule, and CSVLogger output without breaking manual optimization.

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
38/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.