huggingface / huggingface/trl

GKDTrainer + FSDP results in RuntimeError: Expected all tensors to be on the same device, but found at least two devices

Open
#2,580 9 comments 0 reactions 1 assignee Claimed by @kashif View on GitHub
🏋 GKD 🐛 bug
Dominant language
Python
Stars
19.3k
Forks
3k
Avg merge
1d 20h
Merged PRs (30d)
194

Description

### System Info

- Platform: Linux-5.4.0-99-generic-x86_64-with-glibc2.35
- Python version: 3.10.12
- PyTorch version: 2.4.1
- CUDA device(s): A100-SXM-80GB, A100-SXM-80GB, A100-SXM-80GB, A100-SXM-80GB
- Transformers version: 4.47.1
- Accelerate version: 1.2.1
- Accelerate config: not found
- Datasets version: 3.2.0
- HF Hub version: 0.27.0
- TRL version: 0.13.0
- bitsandbytes version: not installed
- DeepSpeed version: not installed
- Diffusers version: not installed
- Liger-Kernel version: not installed
- LLM-Blender version: not installed
- OpenAI version: not installed
- PEFT version: 0.14.0

### Information

- [ ] The official example scripts
- [x] My own modified scripts

### Tasks

- [ ] An officially supported task in the `examples` folder
- [x] My own task or dataset (give details below)

### Reproduction

```python
device = "cuda" if torch.cuda.is_available() else "cpu"

# Model and tokenizer setup
model_name = "Qwen/Qwen2.5-0.5B-Instruct"

# Load models
model = AutoModelForCausalLM.from_pretrained(
model_name,
)
teacher_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-3B-Instruct")

training_args = GKDConfig(
output_dir=model_save_path + "-checkpoints",
num_train_epochs=5,
per_device_train_batch_size=32,
per_device_eval_batch_size=32,
evaluation_strategy="steps",
eval_steps=100,
logging_steps=100,
save_steps=5000,
learning_rate=1e-5,
lr_scheduler_type="cosine",
weight_decay=0.01,
warmup_ratio=0.1,
dataloader_num_workers=4,
fsdp="hybrid_shard auto_wrap",
fsdp_transformer_layer_cls_to_wrap="Qwen2DecoderLayer",
fsdp_config={
"activation_checkpointing": True
}
)

trainer = GKDTrainer(
model=model,
teacher_model=teacher_model,
args=training_args,
processing_class=tokenizer,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
)

trainer.train()
```
and run with script `torchrun --nproc_per_node=4 train.py`
outputs:

```
Traceback (most recent call last):
File "/data/Train/Qwen/Train.py", line 92, in
trainer.train()
File "/data/translation_env/lib/python3.10/site-packages/transformers/trainer.py", line 2164, in train
return inner_training_loop(
File "/data/translation_env/lib/python3.10/site-packages/transformers/trainer.py", line 2524, in _inner_training_loop
tr_loss_step = self.training_step(model, inputs, num_items_in_batch)
File "/data/translation_env/lib/python3.10/site-packages/trl/trainer/gkd_trainer.py",line 304, in training_step
loss = super().training_step(model, inputs, num_items_in_batch)
File "/data/translation_env/lib/python3.10/site-packages/transformers/trainer.py", line 3654, in training_step
loss = self.compute_loss(model, inputs, num_items_in_batch=num_items_in_batch)
File "/data/translation_env/lib/python3.10/site-packages/trl/trainer/gkd_trainer.py",line 229, in compute_loss
outputs_teacher = self.teacher_model(
[rank3]: Traceback (most recent call last):
[rank3]: File "/data/Train/Qwen/Train.py", line 92, in
[rank3]: trainer.train()
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/transformers/trainer.py", line 2164, in train
[rank3]: return inner_training_loop(
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/transformers/trainer.py", line 2524, in _inner_training_loop
[rank3]: tr_loss_step = self.training_step(model, inputs, num_items_in_batch)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/trl/trainer/gkd_trainer.py", line 304, in training_step
[rank3]: loss = super().training_step(model, inputs, num_items_in_batch)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/transformers/trainer.py", line 3654, in training_step
[rank3]: loss = self.compute_loss(model, inputs, num_items_in_batch=num_items_in_batch)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/trl/trainer/gkd_trainer.py", line 229, in compute_loss
[rank3]: outputs_teacher = self.teacher_model(
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1553, in _wrapped_call_impl
[rank3]: return self._call_impl(*args, **kwargs)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1562, in _call_impl
[rank3]: return forward_call(*args, **kwargs)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/transformers/models/qwen2/modeling_qwen2.py", line 1165, in forward
[rank3]: outputs = self.model(
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1553, in _wrapped_call_impl
[rank3]: return self._call_impl(*args, **kwargs)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1562, in _call_impl
[rank3]: return forward_call(*args, **kwargs)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/transformers/models/qwen2/modeling_qwen2.py", line 854, in forward
[rank3]: inputs_embeds = self.embed_tokens(input_ids)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1553, in _wrapped_call_impl
[rank3]: return self._call_impl(*args, **kwargs)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1562, in _call_impl
[rank3]: return forward_call(*args, **kwargs)
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/modules/sparse.py", line 164, in forward
[rank3]: return F.embedding(
[rank3]: File "/data/translation_env/lib/python3.10/site-packages/torch/nn/functional.py", line 2267, in embedding
[rank3]: return torch.embedding(weight, input, padding_idx, scale_grad_by_freq, sparse)
[rank3]: RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cpu and cuda:3! (when checking argument for argument index in method wrapper_CUDA__index_select)
```

### Expected behavior

Running a GKDTrainer with FSDP torchrun seems to be resulting in expected all tensors to be on the same device error.

### Checklist

- [x] I have checked that my issue isn't already filed (see [open issues](https://github.com/huggingface/trl/issues?q=is%3Aissue))
- [x] I have included my system information
- [x] Any code provided is minimal, complete, and reproducible ([more on MREs](https://docs.github.com/en/get-started/writing-on-github/working-with-advanced-formatting/creating-and-highlighting-code-blocks))
- [x] Any code provided is properly formatted in code blocks, (no screenshot, [more on code blocks](https://docs.github.com/en/get-started/writing-on-github/working-with-advanced-formatting/creating-and-highlighting-code-blocks))
- [x] Any traceback provided is complete

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.