All concepts
Distributed Training
Split training work across many GPUs or machines while keeping gradients synchronized.
Deep Learning · Advanced · ~8 min
In plain English
The model is too big for one machine, so you split either the data or the model itself across many, and make them agree after every step.
Why it's worth your time
The moment a model doesn't fit on one GPU, this stops being an optimization and becomes the only way to train at all.
If you remember three things
- Data parallel: same model everywhere, different batches, gradients averaged
- Model/tensor parallel: one model split across devices
- Communication, not compute, is usually the bottleneck
Overview
Splitting training across many GPUs or machines while keeping model replicas synchronized. In data parallelism each worker processes a shard and computes gradients, an all-reduce averages them across devices, and every replica applies the identical optimizer step. Scales throughput at the cost of communication and coordination.
How it works
- Start: Mini-batches Large datasets are split across workers.
- Mini-batches -> GPU Workers Each worker computes forward and backward passes on its shard.
- GPU Workers -> All-Reduce Gradients are averaged across devices.
- All-Reduce -> Optimizer Step Each replica applies the same update.
- Optimizer Step -> Scale Up Throughput improves, but communication, memory, and reproducibility become harder.
In an interview
Distributed training spreads work across GPUs to train faster or fit bigger models. Data parallelism replicates the model, gives each worker a different mini-batch shard, then all-reduces gradients so every replica updates identically. Beyond that, model, tensor, and pipeline parallelism partition the model itself. The bottleneck is inter-device communication, not compute.
Production defaults
- First reach
- data parallel (DDP). Simple, and enough until the model itself doesn't fit
- Doesn't fit
- FSDP / ZeRO — shard optimizer state, gradients, then parameters, in that order
- Batch scaling
- scale the learning rate with the global batch size, and warm up longer
What breaks
- 8 GPUs give 3× the speed — Communication-bound. Check interconnect, enable gradient bucketing/overlap, and raise per-device batch size.
- Loss differs from the single-GPU run — Batch-norm statistics or a seeding/sharding bug. Use SyncBatchNorm and verify each rank sees distinct data.