Model Parallelism
When a model does not fit one GPU even for one step: shard its optimizer state, gradients and weights with ZeRO and FSDP, split its matrices with tensor parallelism, its layers with pipeline parallelism, its experts with expert parallelism, and combine them into a 3D plan for Llama 3.1 70B on 32 H100s.
An interactive AI Infrastructure lesson: 22 steps, about 34 minutes, on a live simulation in your browser.
The 1.5B GPT is done and the team wants Llama 3.1 70B trained on the same pod of 32 H100s. Before anyone submits a job, do the memory maths. The plan op does exactly that and runs nothing.
Mixed-precision training with Adam keeps 16 bytes per parameter: 2 for bf16 weights, 2 for bf16 gradients, 12 for the fp32 master weights and the two Adam moments. 70.6 B × 16 = 1.13 TB of model state, plus 91 GB of activations for one 4,096-token sequence through 80 layers: 1.22 TB on a GPU with 80 GB.
What you will learn
It does not fit
- Llama 70B on one GPU: 1.22 TB: Training memory is 16 bytes per parameter before activations. Divide by the GPU's memory and you have the least number of GPUs the model state must be split across.
- And Llama 405B?: The split sets the per-GPU state, and the GPU count sets the floor: no parallelism scheme puts 6.49 TB into 32 × 80 GB.
Sharding state: ZeRO and FSDP
- ZeRO: stop storing everything N times: ZeRO-1, 2 and 3 shard the optimizer state, then the gradients, then the weights across the data-parallel ranks. Model state per GPU falls toward 16 bytes × parameters / N; activations do not shrink at all.
- What each ZeRO stage costs on the wire: ZeRO-1 and 2 save memory for free; ZeRO-3 / FSDP adds a weight all-gather in both passes, about 1.5× the data-parallel traffic, which must also hide behind compute.
- FSDP over all 32 GPUs: Activations belong to the work, not to the model: ZeRO and FSDP never shard them. Only splitting the work of a layer (tensor parallelism) or the layers (pipeline) or recomputing them reduces them.
- Activation checkpointing: maths for memory: Activation checkpointing trades a third more compute for most of the activation memory. It is the first knob to turn when a plan is short by tens of gigabytes.
- Drill: ZeRO-3 state per GPU: ZeRO-3 model state per GPU = 16 bytes × parameters ÷ data-parallel ranks.
Tensor parallelism
- Tensor parallel: split every matrix: Tensor parallelism divides every layer's matrices across GPUs and pays with blocking all-reduces of the activations, four per layer. It saves memory and compute per GPU, and costs latency on every layer.
- Break it: tensor parallel across servers: Each tensor-parallel GPU all-reduces the whole activation tensor, so TP communication per GPU grows with the TP degree while its maths does not shrink. Keep TP inside the fastest, lowest-latency domain.
- Tensor parallel lives inside NVLink: Map the chattiest parallelism to the fastest link: tensor parallel on NVLink or Infinity Fabric, pipeline and data parallel across the network.
Pipeline parallelism
- Pipeline parallel: split the layers: A pipeline with p stages and m micro-batches idles for (p−1)/(m+p−1) of the time. Pipeline parallelism is cheap on the network and expensive in bubbles.
- Shrinking the bubble: More micro-batches shrink the bubble toward zero and shrink activation memory, but each micro-batch gets smaller and less efficient, and the global batch limits how many there can be.
- Break it: switch recomputation off: In a pipeline, the first stage holds activations for up to p micro-batches at once. That, not the average stage, sets the memory limit.
- Drill: size the bubble
Experts and 3D parallelism
- Expert parallel for mixture-of-experts: Mixture-of-experts decouples compute (active parameters) from memory (total parameters). Expert parallelism splits the experts across GPUs and pays with all-to-all token exchanges.
- 3D parallelism: a plan for 70B on 32 GPUs: Pick TP to fill the NVLink domain, PP to make the model fit with a small bubble, and let DP (with ZeRO) take the rest of the GPUs. Check memory headroom, not just speed.
Megatron, DeepSpeed, FSDP
- Megatron-LM: the 3D plan as flags: In Megatron you choose TP and PP; DP is what is left of the world size. The global batch, micro-batch size and DP together fix the number of micro-batches.
- DeepSpeed: ZeRO as a config file: ZeRO stage 1, 2 or 3 is a config value, not a code change; the framework inserts the reduce-scatters and all-gathers.
- PyTorch FSDP in code: FSDP is ZeRO-3 as a PyTorch wrapper: shard per transformer block, gather a block's weights just before it runs, free them after.
- Drill: tensor parallel in Megatron: Megatron's two model-parallel knobs are --tensor-model-parallel-size and --pipeline-model-parallel-size; data parallelism is whatever divides the rest.
Recap & playground
- Cheat sheet
- Playground: re-plan a running job