Lightning-AI / Lightning-AI/lightning-thunder

Models trained with FSDP + Thunder doesn't work with litgpt chat

Open
#895 5 comments 0 reactions 0 assignees View on GitHub
triage review
Dominant language
Python
Stars
1.5k
Forks
121
PR merge metrics
No merged PRs in 30d

Description

I was able to train Llama3-8b model with Thunder for a few steps and then save it. However when I try to use later `litgpt generate` or `litgpt chat` with the saved checkpoint I get an error about size mismatch. When I run the training in Eager mode everything works.

## 🐛 Bug

### To Reproduce
1. Please extract this archive and put all the files into selected directory (let's call it CHECKPOINT_DIR)
[Meta-Llama-3-8B-tuned.zip](https://github.com/user-attachments/files/16441914/Meta-Llama-3-8B-tuned.zip) . Here is the [license](https://github.com/meta-llama/llama-models/blob/main/models/llama3/LICENSE).

These are Llama-3B configuration files (no weights), they can be also downloaded by running:
`litgpt download meta-llama/Meta-Llama-3-8B`

2. Copy the benchmarking script from this repo located here `thunder/benchmarks/benchmark_litgpt.py` and add model saving in line 622:
``` # save weights
torch_dist.barrier()
states = benchmark.model.state_dict()
if global_rank == 0:
torch.save(states, "/lightning-thunder/checkpoints/meta-llama/Meta-Llama-3-8B-tuned/lit_model.pth")
```

To be sure that version of the script is the same, I'm also attaching the full, modified file (it's python code, but I can add only txt files here): [benchmark_litgpt.txt](https://github.com/user-attachments/files/16442033/benchmark_litgpt.txt)

Let's assume it's located in SCRIPT_DIR directory.

3. Start docker container on a node with 8xH100:
```
docker run --pull=always --gpus all --ipc=host --ulimit \
memlock=-1 --ulimit stack=67108864 -it \
-v ${CHECKPOINT_DIR}:/lightning-thunder/checkpoints/meta-llama/Meta-Llama-3-8B-tuned \
-v ${SCRIPT_DIR}:/repro
INTERNAL_IMAGE:nvidia internal container from 20240731
```
4. Install recent litgpt version:
```
python -m pip install litgpt==0.4.5
```
**For Eager**

5E. Run training for Eager (on dummy data so output won't make sense, but it's easier to run the reproduction instructions)
```
torchrun --standalone --max-restarts=0 --nproc-per-node=8 /repro/benchmark_litgpt.py --model_name Llama-3-8B --max_iters 10 --warmup_iters 2 --distributed_mode fsdp --shard_mode zero3 --bucketing_mode block
```
You should see new file lit_model.pth in checkpoint directory.

6E. Try to chat with the saved model:

```
litgpt chat /lightning-thunder/checkpoints/meta-llama/Meta-Llama-3-8B-tuned
```
It should run but return garbage.

**For Thunder**

5T. You can remove the lit_model.pth (but it will be overwritten anyway) and then run:
```
torchrun --standalone --max-restarts=0 --nproc-per-node=8 /repro/benchmark_litgpt.py --model_name Llama-3-8B --max_iters 10 --warmup_iters 2 --distributed_mode fsdp --shard_mode zero3 --bucketing_mode block --compile thunder
```

6T. Try to chat with the saved model:

```
litgpt chat /lightning-thunder/checkpoints/meta-llama/Meta-Llama-3-8B-tuned
```

There is an error:
> {'access_token': None,
> 'checkpoint_dir': PosixPath('/lightning-thunder/checkpoints/meta-llama/Meta-Llama-3-8B-tuned'),
> 'compile': False,
> 'max_new_tokens': 50,
> 'multiline': False,
> 'precision': None,
> 'quantize': None,
> 'temperature': 0.8,
> 'top_k': 200,
> 'top_p': 1.0}
> Traceback (most recent call last):
> File "/usr/local/bin/litgpt", line 8, in
> sys.exit(main())
> File "/usr/local/lib/python3.10/dist-packages/litgpt/__main__.py", line 71, in main
> CLI(parser_data)
> File "/usr/local/lib/python3.10/dist-packages/jsonargparse/_cli.py", line 119, in CLI
> return _run_component(component, init.get(subcommand))
> File "/usr/local/lib/python3.10/dist-packages/jsonargparse/_cli.py", line 204, in _run_component
> return component(**cfg)
> File "/usr/local/lib/python3.10/dist-packages/torch/utils/_contextlib.py", line 116, in decorate_context
> return func(*args, **kwargs)
> File "/usr/local/lib/python3.10/dist-packages/litgpt/chat/base.py", line 258, in main
> load_checkpoint(fabric, model, checkpoint_path)
> File "/usr/local/lib/python3.10/dist-packages/litgpt/utils.py", line 362, in load_checkpoint
> model.load_state_dict(state_dict, strict=strict)
> File "/usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py", line 2542, in load_state_dict
> raise RuntimeError(
> RuntimeError: Error(s) in loading state_dict for GPT:
> size mismatch for lm_head.weight: copying a param with shape torch.Size([16032, 4096]) from checkpoint, the shape in current model is torch.Size([128256, 4096]).
> size mismatch for transformer.wte.weight: copying a param with shape torch.Size([16032, 4096]) from checkpoint, the shape in current model is torch.Size([128256, 4096]).
> size mismatch for transformer.h.0.norm_1.weight: copying a param with shape torch.Size([512]) from checkpoint, the shape in current model is torch.Size([4096]).
> ...

[Complete output](https://github.com/user-attachments/files/16442319/output.txt)

### Expected behavior

We should be able to run model trained with Thunder with litgpt instructions.

### Environment

nvidia-smi output:
![image](https://github.com/user-attachments/assets/a11a6cb4-583f-47be-ba9d-95e678af5029)

Version of packages:

> lightning-thunder 0.2.0.dev0 /opt/pytorch/lightning-thunder
> lightning-utilities 0.11.6
> litgpt 0.4.5
> nvfuser 0.2.8+gitaf62096 /opt/pytorch/nvfuser
> pytorch-lightning 2.3.3
> torch 2.5.0a0+git83db609
> torchmetrics 1.4.0.post0
> torchvision 0.19.0a0+d23a6e1

Contributor guide

No contributing guide indexed for this repository

Research direction

Reproduce the mismatch with thunder/benchmarks/benchmark_litgpt.py using FSDP, zero3, and --compile thunder, then inspect litgpt/chat/base.py and litgpt/utils.py around checkpoint loading. Compare the saved state_dict shapes with the model created for litgpt chat. Done means the Thunder-trained checkpoint loads through the documented litgpt chat command without size-mismatch errors.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
compilers, 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.