Skip to content

Data parallelism and ZeRO stages 1 to 3

ModuleL11.3 · side · Python · Pass 9 (optional) · 4 to 5 h
You buildpython/tinyllm/dist/zero.py: DDP (forward, sync_grads), ZeroOptimizer (gather, step, zero_grad, memory_bytes)
Contractcourse/contracts/py/tinyllm/dist/zero.pyi
Testscourse/tests/L11.3/ (what they check: section 4) · your own tests in python/tests/l11-3-zero/, rung R5, graded by mutation (threshold 0.80, every pitfall mutant required)
NeedsL11.2 Comm and chunk_bounds · L0.1 Tensor · L0.4 Module and layers · M10.2 SGD and M10.3 AdamW (single-process references and wrapped optimizers) · reading: M05.1 the memory plan, L11.1 accumulation
Used bythe capstone trainer’s --world 4 --zero 2 option (C1, joins the registry with the capstone)
Milestonenone; this optional side module is not part of MS-L11
Optional depthLi et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training (VLDB 2020); Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models (SC 2020); Zhao et al., PyTorch FSDP (VLDB 2023)
  • Data parallelism is gradient accumulation across processes: equal slices, mean losses, an averaged gradient, and every replica takes the same step (test_ddp_matches_single_process).
  • Replicas must start equal, so DDP broadcasts rank 0’s weights before the first step (test_ddp_broadcasts_rank0_weights).
  • ZeRO keeps one copy of each piece of training state across the ranks instead of pp: stage 1 the optimizer state, stage 2 also the gradients, stage 3 also the parameters, so per-rank memory falls toward 1/p1/p (test_hand_example_memory_by_stage, test_zero_memory_is_one_over_p).
  • Sharding changes who stores what, never the arithmetic: an elementwise optimizer on a chunk computes exactly the single process’s update for those entries (test_hand_example_zero_step, test_zero_matches_single_process).
Terminal window
ol start L11.3 # stubs zero.py into your repo
ol tests L11.3 # read the test catalog first
ol check L11.3 # course tests, then your tests graded by mutation
ol mutate L11.3 # the full mutation grade of your tests
ol check L11.3 --ref-deps # only if your L11.2 is not passing yet
ol diff L11.3 # after passing: your code against the reference

L11.2 gave you an all-reduce. The first thing anyone builds with it is data parallelism: pp workers, each with a full copy of the model, each computing the gradient of its own slice of the batch, then averaging. That turns pp processes into one large-batch trainer, and it is exactly L11.1’s accumulation spread over processes. But every replica also holds a full copy of the optimizer state: for AdamW, two moments per parameter, twice the size of the model. At scale that redundancy is what runs out of memory, and ZeRO removes it in three stages. On a laptop the capstone does not need either, so this module is optional; it is how every large model you will serve was trained, and the capstone’s --world 4 --zero 2 run must match the single-process run to show you have it right.

SymbolMeaningType / shape
pp, rrnumber of ranks, this rankint
w∈Rnw \in \mathbb{R}^nall parameters flattened in parameters() orderfloat64[n]
grg_rthe gradient of rank rr‘s local mean lossfloat64[n]
gˉ=1p∑rgr\bar g = \frac1p \sum_r g_rthe averaged gradientfloat64[n]
cr=[ar,br)c_r = [a_r, b_r)rank rr‘s chunk of ww (L11.2’s chunk_bounds)slice
β\betabytes per value (8 for float64)int
KKoptimizer state values per parameter (2 for AdamW’s moments, plus a master copy where kept)int

DDP. Split a batch of NN rows into pp equal slices. Each rank computes gr=∇Lrg_r = \nabla L_r of its slice’s mean loss. By L11.1’s argument with equal sizes, the full batch’s gradient is 1p∑rgr=gˉ\frac1p \sum_r g_r = \bar g. One all_reduce(op="mean") over all gradients, flattened into a single bucket, gives every rank gˉ\bar g with identical bits (L11.2). If every replica starts from the same weights and applies the same optimizer to the same gˉ\bar g, they stay identical forever, and they equal one process training on the whole batch up to the rounding of the ring’s sum. So DDP is three things: broadcast rank 0’s parameters at construction, average gradients after backward(), and nothing else (the optimizer is untouched).

The redundancy. Per rank, DDP with AdamW in float64 holds nn parameters, nn gradients, and 2n2n moments: 4nβ4n\beta bytes, the same as M05.1’s single-process plan, on every one of the pp ranks. Yet each rank uses only the result of the update, which is identical everywhere.

Stage 1: shard the optimizer state. Rank rr updates only its chunk crc_r. It keeps a master copy w[cr]w[c_r] and the optimizer state for those entries only, and wraps them in one flat “parameter” the unchanged optimizer can step. A step: all-reduce the gradients (as DDP), take gˉ[cr]\bar g[c_r], step, then all_gather the updated chunks so every rank’s model has the full new ww. AdamW and SGD are elementwise (each entry’s update depends only on that entry’s gradient and state), so the chunked update equals the single-process update entry by entry. Optimizer memory per rank: Kn/pK n / p.

Stage 2: shard the gradients. After the all-reduce each rank still holds all nn averaged gradients but uses only crc_r. Replace the all-reduce by a reduce_scatter, which leaves rank rr with just ∑rgr[cr]\sum_r g_r[c_r]; divide by pp; drop the full .grad arrays. Gradient memory per rank: n/pn/p (in a real system the full gradients are reduce-scattered bucket by bucket during backward, so they never all exist at once). The traffic is the same: a reduce-scatter plus an all-gather is exactly an all-reduce.

Stage 3: shard the parameters. Between steps each rank keeps only w[cr]w[c_r] and releases the rest. Before a forward pass, gather() all-gathers the chunks into full parameters; after the step they are released again. Memory per rank: everything over pp, at the cost of one more all-gather per step (real systems gather layer by layer just before use and free right after).

state per rankDDPstage 1stage 2stage 3
parametersnnnnnnn/pn/p
gradientsnnnnn/pn/pn/pn/p
optimizer (moments, master)KnKnKn/pKn/pKn/pKn/pKn/pKn/p

Your emulation’s bookkeeping. memory_bytes() reports what the rank holds after a step: the parameters’ .data (in stage 3 between steps, only the chunk), the .grad arrays plus the chunk’s gradient, and every array in the wrapped optimizer’s state_dict() plus, in stages 1 and 2, the chunk’s master copy (in stage 3 the chunk is the parameter storage, counted once). Stage 1 keeps the averaged full gradients in .grad until zero_grad(), as a training loop that logs gradient norms would.

One parameter w=[0,1,…,9]w = [0, 1, \dots, 9] (n=10n = 10) on p=4p = 4 ranks. chunk_bounds(10, 4) gives chunks of 3, 3, 2, 2: rank 0 owns entries 0 to 2, rank 3 entries 8 and 9. Rank rr‘s loss is ∑i(wi−r)2/2\sum_i (w_i - r)^2/2, so gr=w−rg_r = w - r and the averaged gradient is gˉ=w−1.5\bar g = w - 1.5 (the mean of 0, 1, 2, 3).

The step. AdamW with η=0.5\eta = 0.5, no weight decay. On the first step, m=(1−β1)gˉm = (1 - \beta_1)\bar g and v=(1−β2)gˉ2v = (1 - \beta_2)\bar g^2; after bias correction m^=gˉ\hat m = \bar g and v^=∣gˉ∣\sqrt{\hat v} = |\bar g|, so each entry moves by η gˉ/∣gˉ∣=±0.5\eta\, \bar g / |\bar g| = \pm 0.5 (the ϵ=10−8\epsilon = 10^{-8} changes the eighth digit). Entries 0 and 1 have gˉ<0\bar g < 0 and move up; the others move down:

w=[0.5, 1.5, 1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5].w = [0.5,\ 1.5,\ 1.5,\ 2.5,\ 3.5,\ 4.5,\ 5.5,\ 6.5,\ 7.5,\ 8.5].

Every stage gives this on every rank: rank 0 computes entries 0 to 2, rank 3 entries 8 and 9, and the all-gather (or stage 3’s gather()) assembles the rest.

The memory after the step, in bytes (float64, β=8\beta = 8):

stage 1, rank 0stage 1, rank 3stage 2, rank 0stage 2, rank 3stage 3, rank 0stage 3, rank 3one process
params80808080241680
grads80 + 24 = 10480 + 16 = 962416241680
optimizer2·24 + 24 = 722·16 + 16 = 4872484832160

Rank 0’s stage 3 total is 96 bytes against 320 for one process: about 1/p1/p, up to the unequal chunks. These are test_hand_example_zero_step and test_hand_example_memory_by_stage.

python/tinyllm/dist/zero.py
class DDP(Module):
def __init__(self, model: Module, comm: Comm) # broadcasts rank 0's parameters
def forward(self, *args, **kwargs) # model(*args, **kwargs)
def sync_grads(self) -> None # one bucket, all_reduce(op="mean")
class ZeroOptimizer:
def __init__(self, opt_cls, params, comm: Comm, stage: Literal[1, 2, 3], **kw)
def gather(self) -> None # stage 3: full parameters before forward
def step(self) -> None # reduce, update the chunk, gather or release
def zero_grad(self) -> None
def memory_bytes(self) -> dict[str, int] # params, grads, optimizer

A ZeRO training step on each rank: zo.gather(), zo.zero_grad(), forward and backward() on the rank’s slice, zo.step(). With DDP: opt.zero_grad(), forward and backward(), ddp.sync_grads(), opt.step().

TestKINDChecksWhy it matters downstream
test_hand_example_zero_stepunitsection 3’s step, every stage, every rankyou and the test agree on the chunking
test_hand_example_memory_by_stageunitsection 3’s byte tablethe stage definitions
test_ddp_matches_single_processdifferential2 and 4 ranks, 5 momentum-SGD steps, 10−1010^{-10}, bitwise-equal gradientsDDP is just a bigger batch
test_ddp_broadcasts_rank0_weightsunitranks initialized differently end equal to rank 0’s runreplicas never drift
test_zero_matches_single_processdifferentialstages 1 to 3 on 2 and 4 ranks, 4 AdamW steps with decaysharding never changes training
test_zero_memory_is_one_over_ppropertya 44-parameter model on 4 ranks: exact shares per stagethe reason to use ZeRO
test_zero_rejects_bad_stageboundarystage 0 or 4, no parametersconfiguration errors surface early
PitfallSymptomCaught by
1. not broadcasting the initial weightsreplicas with different seeds never agreetest_ddp_broadcasts_rank0_weights (mutant s01)
2. summing gradients instead of averagingthe learning rate is effectively pp times largertest_ddp_matches_single_process (mutant s02)
3. synchronizing only some parametersreplicas silently drift aparttest_ddp_matches_single_process (mutant s03)
4. a rank starting from the wrong chunkits master copy belongs to its neighbourtest_hand_example_zero_step (mutant s04)
5. reduce-scatter without dividing by ppstage 2 and 3 steps on the summed gradienttest_zero_matches_single_process (mutant s05)
6. no all-gather after the updateother ranks’ chunks go stale in your modeltest_hand_example_zero_step (mutant s06)
7. stage 2 keeping the full gradientsno gradient memory savedtest_hand_example_memory_by_stage (mutant s07)
8. stage 3 keeping the full parametersno parameter memory savedtest_zero_memory_is_one_over_p (mutant s08)
9. unflattening in the wrong orderparameters swapped between tensorstest_zero_matches_single_process (mutant s09)
10. dropping the optimizer’s hyperparametersthe default learning rate and decay instead of yourstest_zero_matches_single_process (mutant s10)
DirectionModuleHow it uses this
BackL11.2broadcast, all_reduce, reduce_scatter, all_gather, and chunk_bounds
BackL0.1parameters are Tensors; .grad is what gets reduced
BackL0.4DDP is a Module wrapping yours; parameters() fixes the flat order
BackM10.2SGD in the DDP tests
BackM10.3AdamW wrapped by ZeroOptimizer
ForwardC1(optional) --world 4 --zero 2 trains the capstone on 4 local processes and must match the single-process loss
Your pieceProduction equivalentWhat it addsWhere to look
DDPtorch.nn.parallel.DistributedDataParallelgradient buckets all-reduced during backward (overlapping communication with compute), unused-parameter detectiontorch/nn/parallel/distributed.py, torch/csrc/distributed/c10d/reducer.cpp
ZeroOptimizer stages 1 and 2DeepSpeed ZeRObucketed reduce-scatter during backward, CPU and NVMe offload (ZeRO-Offload, ZeRO-Infinity)deepspeed/runtime/zero/stage_1_and_2.py
stage 3PyTorch FSDP2per-layer all-gather just before use, prefetching the next layer, mixed-precision shardstorch/distributed/fsdp/
clippingsharded global-norm clippingan all-reduce of squared chunk norms before the stepFullyShardedDataParallel.clip_grad_norm_ in torch/distributed/fsdp/