Lightning-AI / Lightning-AI/pytorch-lightning

Pytorch FSDPStrategy saving checkpoint behavior work correctly?

Open
#20,100 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

bug checkpointing fabric strategy: fsdp
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

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 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.