Lightning-AI / Lightning-AI/litgpt

Weird error when using activation checkpointing for FSDPStrategy

Open
#805 2 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
13.7k
Forks
1.5k
Avg merge
15h 37m
Merged PRs (30d)
1

Description

I'm training tinyllama with 8 A40s.
Everything goes very smooth until I want to increase the micro batch size for better computation to communication ratio.

I follow the official tutorial of lit gpt by passing `activation_checkpointing_policy={Block}` into FSDPStrategy. The modified setup is also attached below.

```
def setup(
devices: int = 8,
train_data_dir: Path = Path("data/redpajama_sample"),
val_data_dir: Optional[Path] = None,
precision: Optional[str] = None,
tpu: bool = False,
resume: Union[bool, Path] = False,
) -> None:
precision = precision or get_default_supported_precision(training=True, tpu=tpu)

if devices > 1:
if tpu:
...
else:
strategy = FSDPStrategy(
auto_wrap_policy={Block},
activation_checkpointing_policy={Block},
state_dict_type="full",
limit_all_gathers=True,
cpu_offload=False,
sharding_strategy="FULL_SHARD",
)
else:
strategy = "auto"
```

But I got some strange errors about the activation checkpointing.
Could someone shed some light on this, anything informative is a big help for me.

```
Traceback (most recent call last):
File "pretrain/tinyllama.py", line 424, in
CLI(setup)
File "/usr/local/lib/python3.8/dist-packages/jsonargparse/_cli.py", line 96, in CLI
return _run_component(components, cfg_init)
File "/usr/local/lib/python3.8/dist-packages/jsonargparse/_cli.py", line 181, in _run_component
return component(**cfg)
File "pretrain/tinyllama.py", line 108, in setup
main(fabric, train_data_dir, val_data_dir, resume)
File "pretrain/tinyllama.py", line 160, in main
train(fabric, state, train_dataloader, val_dataloader, monitor, resume)
File "pretrain/tinyllama.py", line 244, in train
fabric.backward(loss / gradient_accumulation_steps)
File "/usr/local/lib/python3.8/dist-packages/lightning/fabric/fabric.py", line 422, in backward
self._strategy.backward(tensor, module, *args, **kwargs)
File "/usr/local/lib/python3.8/dist-packages/lightning/fabric/strategies/strategy.py", line 192, in backward
self.precision.backward(tensor, module, *args, **kwargs)
File "/usr/local/lib/python3.8/dist-packages/lightning/fabric/plugins/precision/fsdp.py", line 126, in backward
super().backward(tensor, model, *args, **kwargs)
File "/usr/local/lib/python3.8/dist-packages/lightning/fabric/plugins/precision/precision.py", line 107, in backward
tensor.backward(*args, **kwargs)
File "/usr/local/lib/python3.8/dist-packages/torch/_tensor.py", line 492, in backward
torch.autograd.backward(
File "/usr/local/lib/python3.8/dist-packages/torch/autograd/__init__.py", line 251, in backward
Variable._execution_engine.run_backward( # Calls into the C++ engine to run the backward pass
File "/usr/local/lib/python3.8/dist-packages/torch/utils/checkpoint.py", line 1075, in unpack_hook
frame.check_recomputed_tensors_match(gid)
File "/usr/local/lib/python3.8/dist-packages/torch/utils/checkpoint.py", line 812, in check_recomputed_tensors_match
raise CheckpointError(
torch.utils.checkpoint.CheckpointError: torch.utils.checkpoint: A different number of tensors was saved during the original forward and recomputation.
Number of tensors saved during forward: 27
Number of tensors saved during recomputation: 8
```

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 with the activation_checkpointing_policy setup in pretrain/tinyllama.py and trace the failure from fabric.backward through the FSDP precision and torch.utils.checkpoint stack shown in the traceback. Reproduce the error with the provided FSDPStrategy configuration; the investigation is complete when the cause and required change for consistent checkpoint recomputation are established.

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
Needs clarification
Newbie friendliness
28/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.