Lightning-AI / Lightning-AI/pytorch-lightning
Pytorch FSDPStrategy saving checkpoint behavior work correctly?
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 31.4k
- Forks
- 3.8k
- Avg merge
- 6d 7h
- Merged PRs (30d)
- 6
Description
### Bug description
The experiment crashes after I resume it from the checkpoint.
After comparing [fabric/strategies/fsdp.py](https://github.com/Lightning-AI/pytorch-lightning/blob/cf348673eda662cc2e9aa71a72a19b8774f85718/src/lightning/fabric/strategies/fsdp.py#L419), [pytorch/strategies/fsdp.py](https://github.com/Lightning-AI/pytorch-lightning/blob/cf348673eda662cc2e9aa71a72a19b8774f85718/src/lightning/pytorch/strategies/fsdp.py#L560), and [Related Pytorch test code](https://github.com/pytorch/pytorch/blob/v2.3.1/test/distributed/checkpoint/test_fsdp_optim_state.py), the saving logic on `state_dict_type=full` seems work different.
As far as I know the pytorch FSDP strategy just save checkpoints with `torch.save()` for full state dict.
Can you investigate on this issue?
### What version are you seeing the problem on?
v2.2
### How to reproduce the bug
```python
trainer:
accelerator: cuda
strategy:
class_path: lightning.pytorch.strategies.fsdp.FSDPStrategy
init_args:
sharding_strategy: SHARD_GRAD_OP
state_dict_type: full
precision: 16-mixed
callbacks:
- class_path: lightning.pytorch.callbacks.ModelCheckpoint
```
### Error messages and logs
```
Traceback (most recent call last):
[rank1]: File "/home/jh/my_lightning_project/main.py", line 9, in
[rank1]: cli_main()
[rank1]: File "/home/jh/my_lightning_project/main.py", line 6, in cli_main
[rank1]: cli = LightningCLI(MyLightningModule, MyDataModule, save_config_kwargs={"overwrite": True})
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/cli.py", line 394, in __init__
[rank1]: self._run_subcommand(self.subcommand)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/cli.py", line 701, in _run_subcommand
[rank1]: fn(**fn_kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/trainer/trainer.py", line 543, in fit
[rank1]: call._call_and_handle_interrupt(
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/trainer/call.py", line 47, in _call_and_handle_interrupt
[rank1]: return trainer_fn(*args, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/trainer/trainer.py", line 579, in _fit_impl
[rank1]: self._run(model, ckpt_path=ckpt_path)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/trainer/trainer.py", line 986, in _run
[rank1]: results = self._run_stage()
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/trainer/trainer.py", line 1030, in _run_stage
[rank1]: self.fit_loop.run()
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/loops/fit_loop.py", line 205, in run
[rank1]: self.advance()
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/loops/fit_loop.py", line 363, in advance
[rank1]: self.epoch_loop.run(self._data_fetcher)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/loops/training_epoch_loop.py", line 140, in run
[rank1]: self.advance(data_fetcher)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/loops/training_epoch_loop.py", line 250, in advance
[rank1]: batch_output = self.automatic_optimization.run(trainer.optimizers[0], batch_idx, kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/loops/optimization/automatic.py", line 190, in run
[rank1]: self._optimizer_step(batch_idx, closure)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/loops/optimization/automatic.py", line 268, in _optimizer_step
[rank1]: call._call_lightning_module_hook(
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/trainer/call.py", line 167, in _call_lightning_module_hook
[rank1]: output = fn(*args, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/core/module.py", line 1308, in optimizer_step
[rank1]: optimizer.step(closure=optimizer_closure)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/core/optimizer.py", line 153, in step
[rank1]: step_output = self._strategy.optimizer_step(self._optimizer, closure, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/strategies/strategy.py", line 238, in optimizer_step
[rank1]: return self.precision_plugin.optimizer_step(optimizer, model=model, closure=closure, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/lightning/pytorch/plugins/precision/fsdp.py", line 163, in optimizer_step
[rank1]: step_output = self.scaler.step(optimizer, **kwargs) # type: ignore[arg-type]
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/amp/grad_scaler.py", line 453, in step
[rank1]: retval = self._maybe_opt_step(optimizer, optimizer_state, *args, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/amp/grad_scaler.py", line 351, in _maybe_opt_step
[rank1]: retval = optimizer.step(*args, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/optim/lr_scheduler.py", line 75, in wrapper
[rank1]: return wrapped(*args, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/optim/optimizer.py", line 391, in wrapper
[rank1]: out = func(*args, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/optim/optimizer.py", line 76, in _use_grad
[rank1]: ret = func(self, *args, **kwargs)
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/optim/adamw.py", line 188, in step
[rank1]: adamw(
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/optim/adamw.py", line 340, in adamw
[rank1]: func(
[rank1]: File "/home/jh/.cache/pypoetry/virtualenvs/my-project-2CIT3KWG-py3.10/lib/python3.10/site-packages/torch/optim/adamw.py", line 550, in _multi_tensor_adamw
[rank1]: torch._foreach_lerp_(device_exp_avgs, device_grads, 1 - beta1)
[rank1]: RuntimeError: output with shape [] doesn't match the broadcast shape [1]
```
### Environment
Current environment
* CUDA:
- GPU:
- NVIDIA RTX A6000
- NVIDIA RTX A6000
- NVIDIA RTX A6000
- NVIDIA RTX A6000
- NVIDIA RTX A6000
- NVIDIA RTX A6000
- NVIDIA RTX A6000
- NVIDIA RTX A6000
- available: True
- version: 12.1
* Lightning:
- lightning: 2.3.3
- lightning-utilities: 0.11.5
- pytorch-lightning: 2.3.3
- torch: 2.3.1
- torchmetrics: 1.4.0.post0
- torchscale: 0.3.0
- torchsummary: 1.5.1
- torchvision: 0.18.1
* Packages:
...
* System:
- OS: Linux
- architecture:
- 64bit
- ELF
- processor: x86_64
- python: 3.10.12
- release: 5.15.0-76-generic
- version: #83-Ubuntu SMP Thu Jun 15 19:16:32 UTC 2023
Used poetry to install packages.
### More info
This issue is also posted on pytorch: https://github.com/pytorch/pytorch/issues/130810
cc @lantiga @justusschock
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by comparing the full-state checkpoint handling in src/lightning/fabric/strategies/fsdp.py and src/lightning/pytorch/strategies/fsdp.py with PyTorch's test/distributed/checkpoint/test_fsdp_optim_state.py. Reproduce the provided FSDPStrategy configuration on Lightning 2.3.3 and inspect the resume path. Done means full-state checkpoints can be saved and resumed without the reported optimizer shape mismatch.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems, machine-learning
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 35/100