๐Ÿ”— Deep Learning ยท Lecture 32 of 38

Distributed Training: Data, Model, Pipeline and Sharded Parallelism

Large models and datasets need many accelerators. We explain data parallelism with all-reduce, sharded data parallelism (ZeRO/FSDP), tensor and pipeline model parallelism, and how they combine at scale.

Training modern foundation models requires thousands of GPUs working together for weeks. Even university projects often benefit from spreading training across several GPUs. Distributed training is fundamentally about two constraints: time (we want training to finish sooner) and memory (the model and its training state may not fit on one device). Different parallelism strategies address each.

Data parallelism#

The simplest and most common approach: replicate the model on every device, give each device a different slice of the mini-batch, and average gradients before updating.

  1. Each of $N$ workers holds a full model copy.
  2. Each computes forward and backward passes on its local batch $\mathcal{B}_k$.
  3. Gradients are averaged across workers with an all-reduce operation:
$$ \mathbf{g} = \frac{1}{N}\sum_{k=1}^{N}\mathbf{g}_k $$
  1. Every worker applies the identical update, so replicas stay synchronised.

The effective batch size is $N \times$ the per-device batch, so the learning rate and warm-up usually need adjusting.

All-reduce#

A naive approach sends all gradients to one parameter server โ€” a bandwidth bottleneck. Ring all-reduce arranges workers in a ring; each sends and receives chunks so that the per-worker communication volume is about $2\frac{N-1}{N}$ times the gradient size โ€” nearly independent of $N$. Libraries such as NCCL implement efficient all-reduce over NVLink and InfiniBand.

DDP in PyTorch#

DistributedDataParallel overlaps communication with computation: it all-reduces gradients in buckets as soon as they are computed during the backward pass.

python
# launch with:  torchrun --nproc_per_node=4 train.py
import os, torch, torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data.distributed import DistributedSampler

dist.init_process_group("nccl")
rank = int(os.environ["LOCAL_RANK"]); torch.cuda.set_device(rank)

model = MyModel().cuda(rank)
model = DDP(model, device_ids=[rank])
sampler = DistributedSampler(train_ds, shuffle=True)          # each rank gets a distinct shard
loader = torch.utils.data.DataLoader(train_ds, batch_size=64, sampler=sampler, num_workers=4)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)

for epoch in range(epochs):
    sampler.set_epoch(epoch)                                   # reshuffle differently each epoch
    for xb, yb in loader:
        loss = torch.nn.functional.cross_entropy(model(xb.cuda(rank)), yb.cuda(rank))
        opt.zero_grad(); loss.backward(); opt.step()           # all-reduce happens in backward
    if rank == 0:
        torch.save(model.module.state_dict(), f"ckpt_{epoch}.pt")   # save from one rank only
dist.destroy_process_group()

Sharded data parallelism: ZeRO and FSDP#

Plain data parallelism replicates everything โ€” weights, gradients and optimiser states โ€” on every GPU. With Adam in mixed precision that is roughly 16 bytes per parameter per GPU: a 10-billion-parameter model needs about 160 GB per GPU, exceeding any single device.

ZeRO (Zero Redundancy Optimizer, Rajbhandari et al., 2020) removes the redundancy by sharding state across the $N$ data-parallel workers:

StageShardedMemory per GPU (approx.)
ZeRO-1Optimiser states~4 + 12/N bytes/param
ZeRO-2+ Gradients~2 + 14/N
ZeRO-3 / FSDP+ Parameters~16/N

In ZeRO-3 / PyTorch FSDP (Fully Sharded Data Parallel), each layer's parameters are gathered from all shards just before they are needed (all-gather), used, and released; gradients are reduce-scattered back to their owning shards. Memory scales down with $N$ at the cost of extra communication. Offloading optimiser states or parameters to CPU memory or NVMe extends this further.

Model parallelism#

When a single layer or the activations are too large, or to reduce per-device compute, split the model itself.

Tensor parallelism#

Split individual weight matrices across devices. For a linear layer $\mathbf{Y} = \mathbf{X}\mathbf{W}$, partition $\mathbf{W}$ by columns so each device computes part of the output; Megatron-LM pairs a column-split first MLP layer with a row-split second layer so only one all-reduce is needed per MLP block. Attention heads split naturally across devices. Tensor parallelism requires very fast interconnects and is typically used within a server.

Pipeline parallelism#

Place consecutive groups of layers ("stages") on different devices. Naively, only one stage works at a time โ€” a large pipeline bubble. Micro-batching (GPipe) splits each batch into micro-batches that flow through the stages concurrently; schedules like 1F1B (one-forward-one-backward) further reduce idle time and memory.

Sequence / context parallelism#

Split long sequences across devices so attention over very long contexts fits in memory.

Expert parallelism#

In mixture-of-experts models, place different experts on different devices and route tokens to them.

3-D parallelism#

Frontier-scale training combines these: tensor parallelism inside a node, pipeline parallelism across nodes, and (sharded) data parallelism across replicas of the pipeline. Choosing the right combination depends on model size, sequence length, cluster topology and interconnect bandwidth.

StrategySolvesCommunicationTypical scope
Data parallel (DDP)TimeGradient all-reduce per stepAny
ZeRO / FSDPMemory (states, params)All-gather + reduce-scatterMany GPUs
Tensor parallelMemory + per-layer computeFrequent, within layersWithin a node
Pipeline parallelMemory (layers)Activations between stagesAcross nodes

Practical concerns#

  • Reproducibility and randomness: seed each rank; use DistributedSampler; only rank 0 logs and saves.
  • Batch normalisation: statistics are per-GPU unless you use SyncBatchNorm.
  • Fault tolerance: long runs on thousands of GPUs will see failures โ€” checkpoint frequently and support automatic restart.
  • Communication efficiency: gradient compression, overlapping communication with computation, and topology-aware placement.
  • Tools: PyTorch DDP/FSDP, DeepSpeed, Megatron-LM, Hugging Face Accelerate, JAX with pjit/sharding annotations.
JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ”— Deep Learning

Mixed-Precision Training and GPU Efficiency

Training in 16-bit arithmetic roughly halves memory and can multiply throughput. We explain float16 and bfloat16, loss scaling, automatic mixed precision, and other practical techniques to make GPUs work harder.

Advancedโฑ 5 min#127
๐Ÿ”— Deep Learning

Graph Neural Networks: Learning on Relational Data

Molecules, social networks, road maps and knowledge graphs are graphs. GNNs learn from them by passing messages between neighbours. We derive message passing, GCN and GAT layers, and survey node, edge and graph-level tasks.

Advancedโฑ 6 min#129
๐Ÿ”— Deep Learning

Debugging Neural Network Training: A Systematic Recipe

Neural networks fail silently โ€” they train, but badly. We present a systematic recipe for finding bugs, from data inspection and overfitting a single batch to monitoring activations, gradients and learning curves.

Intermediateโฑ 6 min#126