❓ Questions & Help
I've managed to get model parallelism working on gpt2 for forward inference by modifying the GPT2Model class and adding a few lines to the generate method to ensure that tensors that need to be on the same device always are. It automatically distributes the blocks evenly across any number of GPUs that are detected. I had to add an additional argument to Trainer (model_parallel) to avoid conflicting distribute behavior. Unfortunately, I'm stuck on backprop, specifically in Trainer.training_step on the line loss.backward().
loss is tensor(71.5152, device='cuda:3', grad_fn=<NllLossBackward>)
The error is:
RuntimeError: expected device cuda:3 but got device cuda:0 (compute_types at ..\aten\src\ATen\native\TensorIterator.cpp:246)
(no backtrace available)
So something somewhere is on the wrong device. It would be a miracle if someone knows how to fix this, but more realistically I'm hoping for a list of things that might be wrong which I can check. Can do a code review with someone from the transformers team. This could be the pattern to enable model parallelism on all PyTorch transformers.
❓ Questions & Help
I've managed to get model parallelism working on
gpt2for forward inference by modifying theGPT2Modelclass and adding a few lines to thegeneratemethod to ensure that tensors that need to be on the same device always are. It automatically distributes the blocks evenly across any number of GPUs that are detected. I had to add an additional argument toTrainer(model_parallel) to avoid conflicting distribute behavior. Unfortunately, I'm stuck on backprop, specifically inTrainer.training_stepon the lineloss.backward().loss is
tensor(71.5152, device='cuda:3', grad_fn=<NllLossBackward>)The error is:
So something somewhere is on the wrong device. It would be a miracle if someone knows how to fix this, but more realistically I'm hoping for a list of things that might be wrong which I can check. Can do a code review with someone from the transformers team. This could be the pattern to enable model parallelism on all PyTorch transformers.