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

  1. Start: Mini-batches Large datasets are split across workers.
  2. Mini-batches -> GPU Workers Each worker computes forward and backward passes on its shard.
  3. GPU Workers -> All-Reduce Gradients are averaged across devices.
  4. All-Reduce -> Optimizer Step Each replica applies the same update.
  5. 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.

Watch it explained

How DDP works || Distributed Data Parallel || Quick explained — Developers Hutt, 3:20

Related