Skip to content
AI-grafen
FAI engineeringModel training and fine-tuning· about 90 min· fast-moving, sources checked often· verified 2026-09-20· EN

Distributed training

Be able to explain data, tensor and pipeline parallelism, and run training on several GPUs.

Prerequisites

Intuition

Three ways of splitting training over several GPUs — they solve different problems:

The strategyWhat is splitIt solvesThe communication
Data parallelismthe batchslow trainingan all-reduce of the gradients every step
Tensor parallelismevery matrixthe model does not fit on one carda lot, within every layer → it needs a fast link
Pipeline parallelismthe layersthe model does not fitless, but it gives «bubbles» of waiting
ZeRO / FSDPthe optimiser state, the gradients, the weightsmemory in data parallelismmore 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 stageShardsThe memory per GPU (7B, Adam, bf16)
— (plain DDP)nothing14 + 14 + 56 = 84 GB
ZeRO-1the optimiser state14 + 14 + 56/N
ZeRO-2+ the gradients14 + (14+56)/N
ZeRO-3 / FSDP+ the parameters84/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

All the sources and licences