All posts
// / Blog

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.

#DistributedTraining#GPU#DeepLearning#PyTorch#MachineLearning#HPC