Training a model on a single GPU: straightforward.
Training on 8 GPUs across 2 machines: a whole new category of headaches.
I spent two weeks debugging a distributed training job that worked on 1 GPU but produced garbage on 4. The culprit: a subtle issue with batch normalization statistics not being properly synchronized across GPUs.
Distributed training lessons from the trenches:
Start with data parallelism (same model on each GPU, different data). It's the simplest and handles 80% of cases. FSDP (Fully Sharded Data Parallel) in PyTorch is the modern standard.
Gradient synchronization is where things go wrong. All-reduce operations need to complete before the optimizer step. If they don't, each GPU diverges silently.
Learning rate scaling: when you multiply GPUs, you often need to adjust the learning rate. Linear scaling with warmup is a good starting point.
Mixed precision training (bf16 or fp16) is basically free performance — 2x speedup with minimal quality impact. Use it by default.
Monitor each GPU independently. A single GPU with a hardware issue can corrupt the entire training run. Check loss values per GPU, not just the aggregate.
And the practical tip that saves the most time: test your distributed training code at small scale first. Run on 2 GPUs with a tiny dataset for a few steps. If that works, scale up. If not, you've found the bug in minutes instead of hours.