Distributed training
Be able to explain data, tensor and pipeline parallelism, and run training on several GPUs.
Prerequisites
- DTrain a neural network in PyTorchrequired
- EOptimisers: momentum, Adam, schedulingrequired
Intuition
Three ways of splitting training over several GPUs — they solve different problems:
| The strategy | What is split | It solves | The communication |
|---|---|---|---|
| Data parallelism | the batch | slow training | an all-reduce of the gradients every step |
| Tensor parallelism | every matrix | the model does not fit on one card | a lot, within every layer → it needs a fast link |
| Pipeline parallelism | the layers | the model does not fit | less, but it gives «bubbles» of waiting |
| ZeRO / FSDP | the optimiser state, the gradients, the weights | memory in data parallelism | more than plain DP, less than TP |
The rule: start with data parallelism (DDP). If the memory is not enough: add ZeRO/FSDP. If the model still does not fit: tensor parallelism within a node, pipeline between nodes.
Formal
ZeRO in three stages (Rajbhandari et al. 2019) — each stage shards one more thing across all N GPUs:
| The stage | Shards | The memory per GPU (7B, Adam, bf16) |
|---|---|---|
| — (plain DDP) | nothing | 14 + 14 + 56 = 84 GB |
| ZeRO-1 | the optimiser state | 14 + 14 + 56/N |
| ZeRO-2 | + the gradients | 14 + (14+56)/N |
| ZeRO-3 / FSDP | + the parameters | 84/N + temporaries |
With 8 GPUs and ZeRO-3 that becomes ~11 GB per card instead of 84.
The effective batch size = batch_per_gpu × the number of gpus × accumulation_steps. It should be kept constant when you change the setup, otherwise you are changing two things at once and the results cannot be compared.
Gradient accumulation simulates a large batch on little memory: run k microbatches, accumulate the gradients, update once. Mathematically nearly equivalent to a batch k times larger (the difference lies in the batch norm statistics, if any are used).
The most common practical mistake: not scaling the learning rate when the effective batch grows, or forgetting that only rank 0 should log and save checkpoints — otherwise they all overwrite each other.
Code
# FSDP in PyTorch (the equivalent of ZeRO-3)
import torch, torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.utils.data.distributed import DistributedSampler
dist.init_process_group("nccl")
rank, world = dist.get_rank(), dist.get_world_size()
torch.cuda.set_device(rank)
model = FSDP(model.cuda(), device_id=rank)
sampler = DistributedSampler(dataset, num_replicas=world, rank=rank, shuffle=True)
loader = torch.utils.data.DataLoader(dataset, batch_size=4, sampler=sampler)
ACC = 4 # the effective batch = 4 * world * 4
for epoch in range(epochs):
sampler.set_epoch(epoch) # otherwise the same order every epoch
for i, (xb, yb) in enumerate(loader):
loss = criterion(model(xb.cuda()), yb.cuda()) / ACC
loss.backward()
if (i + 1) % ACC == 0:
opt.step(); opt.zero_grad()
if rank == 0:
save_checkpoint(model, epoch) # only one process writes
Mastery means
- Explains data, tensor and pipeline parallelism
- Chooses a strategy to suit the model size and the hardware
- Knows about ZeRO/FSDP and gradient accumulation
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — arXiv (open access; licence per article)
- PyTorch — FSDP (BSD-3) — BSD-3-Clause