Distributed Training & GPU Systems
A frontier model doesn't fit on one GPU, and it wouldn't finish training on one even if it did. This page covers the systems side of training: where the bytes go, how numeric formats trade accuracy for speed, how GPUs talk to each other, and how data, tensor, pipeline, context and expert parallelism fit together into the "3D/4D/5D" setups used on clusters of tens of thousands of GPUs. Interviewers ask about this because it tells them whether you can reason about cost, memory and throughput from first principles instead of treating the training stack as a black box.
TL;DR: the 12 things to be able to say out loud
- Mixed-precision Adam costs about 16 bytes per parameter before activations: 2 (BF16 weights) + 2 (BF16 grads) + 12 (FP32 master copy, Adam m, Adam v). It's about 18 if gradients are kept in FP32. So a 7B model needs about 112 GB and a 70B model about 1.1 TB just for model state.
- Activations scale with batch × sequence × hidden × layers and usually dominate at long context. Activation checkpointing trades about 33% more compute (one extra forward pass) for keeping only layer inputs.
- BF16 has FP32's exponent range, so it needs no loss scaling. FP16 does need it. FP8 training (E4M3/E5M2) relies on scaling factors. The trend is toward finer-grained scaling: per-tensor, then per-block (DeepSeek-V3's 1×128 / 128×128 blocks), then hardware microscaling (MXFP8, NVFP4) on Blackwell.
- All-reduce = reduce-scatter + all-gather. With a ring, each GPU sends and receives about \(2\frac{N-1}{N}S\) bytes no matter how big N is. That is bandwidth-optimal, but latency grows with N. Trees fix the latency term.
- The interconnect is a steep hierarchy. NVLink inside a node gives roughly 10–20× more bandwidth than the InfiniBand/RoCE links between nodes. Parallelism strategies are placed onto that hierarchy deliberately.
- DDP replicates everything and all-reduces gradients in buckets, overlapping that with backward. ZeRO-1/2/3 (FSDP) shard optimizer state, then gradients, then parameters across data-parallel ranks. ZeRO-3 costs about 1.5× DDP's communication volume.
- Tensor parallelism (Megatron) splits the MLP column-wise then row-wise, which needs 2 all-reduces per layer in forward and 2 in backward, all on the critical path. That's why TP stays inside an NVLink domain (usually ≤8).
- Sequence parallelism turns TP's all-reduce into a reduce-scatter + all-gather and shards the LayerNorm/dropout activations along the sequence. You get the activation-memory savings for free.
- Pipeline parallelism has a bubble fraction of about \((p-1)/m\). 1F1B bounds activation memory, interleaving divides the bubble by v, and zero-bubble schedules split the backward pass to fill the gaps.
- Context parallelism / ring attention shards the sequence and rotates K/V blocks around a ring. Expert parallelism places MoE experts on different GPUs and needs two all-to-alls per MoE layer per direction.
- Llama 3 405B trained on 16K H100s with TP=8, PP=16, DP=128 (CP=16 for long context), ordered [TP, CP, PP, DP] so the chattiest dimension gets the fastest links, at 38–43% MFU.
- MFU = achieved model FLOPs/s ÷ peak. 35–55% is typical for well-tuned dense training. At 10K+ GPUs, hardware fails several times a day, so checkpointing, fast restart and elasticity are first-class design concerns.
The mental model: compute, memory, communication
Every distributed training decision trades off three resources:
- Compute: matmul throughput (H100 SXM peak dense BF16 is about 989 TFLOP/s NVIDIA H100). Training a dense transformer costs about \(6N\) FLOPs per token (2N forward, 4N backward), where N is the parameter count Kaplan+ 2020.
- Memory: HBM capacity (80 GB on H100, more on H200/B200) and HBM bandwidth (about 3.35 TB/s on H100 SXM).
- Communication: bandwidth and latency of the links between GPUs. These are tiered: NVLink inside a node, InfiniBand/RoCE between nodes.
Every parallelism strategy is a way of cutting the computation graph across devices. Each one cuts along a different axis (batch, hidden dimension, layers, sequence, experts), so each one saves a different kind of memory and adds a different communication pattern. If you can say what each strategy shards, which collective it adds, how many bytes that collective moves, and whether it sits on the critical path, you can answer almost any interview question in this area.
| Strategy | Splits | Saves memory on | Collective added | Typical placement |
|---|---|---|---|---|
| Data parallel (DDP) | Batch | Nothing (activations per GPU shrink with local batch) | All-reduce of grads, once per step | Outermost, across nodes |
| ZeRO-1/2/3, FSDP | Batch, plus sharded state | Optimizer state / grads / params | Reduce-scatter + all-gather (params gathered per layer in Z3) | Across nodes, or hybrid (shard in node, replicate across) |
| Tensor parallel (TP) | Hidden dim inside each matmul | Params, grads, optimizer state, most activations | All-reduce (or RS+AG with SP) ×4 per layer | Inside the NVLink domain |
| Sequence parallel (SP) | Sequence, in non-matmul regions | LayerNorm/dropout activations | Replaces TP all-reduce with RS+AG | Always paired with TP |
| Pipeline parallel (PP) | Layers | Params and state per stage | Point-to-point activation sends | Across nodes (tolerates slow links) |
| Context parallel (CP) | Sequence, including attention | Activations at long context | Ring P2P of K/V, or all-gather of K/V | In node or across nearby nodes |
| Expert parallel (EP) | MoE experts | Expert params | All-to-all dispatch + combine | In node or across a few nodes |
A common opener is "you have a 70B model and 256 H100s, how do you train it?" A strong answer starts with arithmetic (bytes of state per GPU, rough activation memory, FLOPs per step), then picks parallelism dimensions that fit the hardware topology, then names the bottleneck you'd watch (exposed communication, pipeline bubble, recompute overhead). Jumping straight to "use FSDP" without numbers is a weak answer.
- Hugging Face Ultra-Scale Playbook: the best single walkthrough of every parallelism dimension, with thousands of benchmark runs.
- How to Scale Your Model (JAX scaling book): rooflines, sharding notation and communication cost arithmetic. TPU-flavoured, but it generalizes.
Memory accounting for training
GPU memory during training holds four things: parameters, gradients, optimizer state and activations. A fifth category covers temporary buffers, communication buffers and allocator fragmentation, which together often cost 5–15% extra.
Model state: the 16-bytes-per-parameter rule
The standard recipe is mixed-precision Adam(W) Micikevicius+ 2017. Forward and backward run in BF16 (or FP16), and the optimizer keeps an FP32 master copy of the weights, because small updates added to a BF16 weight would be rounded away. Per parameter:
| Item | Precision | Bytes/param | Why it exists |
|---|---|---|---|
| Working weights | BF16 | 2 | Used in forward/backward matmuls |
| Gradients | BF16 (or FP32) | 2 (or 4) | Output of backward, input to optimizer |
| Master weights | FP32 | 4 | Accumulates tiny updates without rounding loss |
| Adam first moment m | FP32 | 4 | Momentum (EMA of grad) |
| Adam second moment v | FP32 | 4 | EMA of grad², sets per-param step size |
| Total | 16 (18 with FP32 grads) | The ZeRO paper writes this as \(2\Psi + 2\Psi + K\Psi\) with \(K=12\) Rajbhandari+ 2019 |
Variations come up a lot. Many frameworks accumulate gradients in FP32 for stability, especially across gradient-accumulation micro-steps, which gives 18 bytes. Some drop the separate BF16 weight copy and cast on the fly. 8-bit optimizers or factored optimizers (Adafactor-style) shrink the 12. Muon-style optimizers keep one momentum buffer rather than two. The 16–18 figure is the default to quote, and it's worth saying out loud which variant you're assuming.
Worked example: 7B and 70B
7B model
Weights BF16: 14 GB
Grads BF16: 14 GB
Adam + master FP32: 84 GB
Total ≈ 112 GB (126 GB with FP32 grads)
That doesn't fit on one 80 GB H100 even before activations. Options: ZeRO-1 over 8 GPUs gives 14 + 14 + 84/8 ≈ 38.5 GB per GPU. ZeRO-3 over 8 GPUs gives 112/8 = 14 GB per GPU, which leaves plenty of room for activations.
70B model
Weights BF16: 140 GB
Grads BF16: 140 GB
Adam + master FP32: 840 GB
Total ≈ 1.12 TB
The model state alone needs at least 14 H100s with perfect sharding. ZeRO-3/FSDP over 64 GPUs gives 17.5 GB per GPU. With TP=8 inside a node plus ZeRO-1 across 16 nodes: params and grads are 280/8 = 35 GB per GPU, optimizer state is 840/(8·16) ≈ 6.6 GB, about 42 GB in total before activations.
The ZeRO paper's reference example, 7.5B parameters on 64-way DP, gives 120 GB per GPU for plain DP, 31.4 GB with stage 1, 16.6 GB with stage 2 and 1.9 GB with stage 3 Rajbhandari+ 2019. It's worth memorizing the shape of that curve.
Activations
Activations are the tensors saved during forward so that backward can compute gradients. For a standard GPT-style layer with hidden size \(h\), sequence length \(s\), micro-batch \(b\), \(a\) attention heads and 16-bit activations, Korthikanti et al. derive the per-layer activation memory without any parallelism as Korthikanti+ 2022:
$$ M_{\text{act/layer}} \approx s\,b\,h\left(34 + \frac{5\,a\,s}{h}\right)\ \text{bytes} $$The \(34\,sbh\) term covers the linear-layer inputs, GeLU input, LayerNorm inputs and dropout masks. The \(5as^2b\) term is the attention score matrix (softmax output, dropout mask), which is quadratic in sequence length. FlashAttention removes most of that quadratic term because it never materializes the \(s\times s\) matrix: it recomputes it in tiles during backward and saves only the output and the softmax statistics. With TP degree \(t\) and sequence parallelism, everything is divided by \(t\).
Worked example (7B, Llama-like: \(h=4096\), \(L=32\), \(s=4096\), \(b=1\)): \(sbh = 16.8\)M elements. With FlashAttention, about \(34 \times 16.8\text{M} \approx 0.57\) GB per layer, so roughly 18 GB for 32 layers. Without FlashAttention, the score term adds \(5 \cdot 32 \cdot 4096/4096 = 160\) to the 34, about 3.3 GB per layer and over 100 GB in total. That's why naive attention made long-context training impractical. Activation memory grows linearly with micro-batch and (with FlashAttention) linearly with sequence length, so at 32K–128K context activations dwarf model state, and you need context parallelism or recomputation.
Forgetting activations, or treating them as fixed. Model state scales with N. Activations scale with \(b \cdot s \cdot h \cdot L\) and don't shrink under ZeRO/FSDP at all, because FSDP shards state, not activations. If someone says "FSDP will fix my OOM at 128K context", they're probably wrong. They need CP, recomputation or offload.
Model state is a fixed cost you can shard across DP ranks. Activations are a per-token cost you can only reduce by giving each GPU fewer tokens (smaller micro-batch, CP), fewer layers (PP), a thinner slice of each tensor (TP+SP), or by not storing them at all (recompute, offload).
- ZeRO (Rajbhandari+ 2019): section 3 is the cleanest memory breakdown in the literature.
- Reducing Activation Recomputation (Korthikanti+ 2022): the activation formula, sequence parallelism and selective recomputation.
Activation checkpointing (recomputation)
Activation checkpointing (also called gradient checkpointing or rematerialization) stores only a subset of activations during forward, usually each transformer block's input, and recomputes the rest during backward Chen+ 2016. Chen et al. showed that checkpointing every \(\sqrt{L}\) layers gives \(O(\sqrt{L})\) memory for one extra forward pass. In transformers, the usual practice is to checkpoint at block boundaries.
- Full recomputation: save only each layer's input (\(2sbh\) bytes per layer in BF16). Cost is about one extra forward pass, so per-token compute goes from \(6N\) to roughly \(8N\), about 33% more FLOPs.
- Selective recomputation: recompute only the parts that are cheap in FLOPs but expensive in memory, mainly the attention softmax/dropout region. Korthikanti et al. report a ~5× activation memory reduction at about 2–4% compute overhead for large GPT models Korthikanti+ 2022. FlashAttention already does this for attention internally, which is why "selective" in modern stacks often means picking which MLP or norm outputs to keep.
- Offloading: move saved activations to CPU memory over PCIe and fetch them back for backward. This helps when PCIe bandwidth can be overlapped with compute. It's increasingly attractive on Grace-Hopper/Grace-Blackwell systems with fast CPU–GPU links.
Recomputation is also why MFU and HFU differ (see MFU). The recomputed forward does real work on the hardware but doesn't count as useful model FLOPs.
"When would you not use activation checkpointing?" When memory isn't the binding constraint: the 33% compute tax is pure waste if you could have fit a bigger micro-batch or dropped a parallelism dimension instead. Good answers mention that recompute is often cheaper than adding TP or PP, which bring communication and bubbles, and that the decision should come from profiling, not habit.
- Training Deep Nets with Sublinear Memory Cost (Chen+ 2016): the original √L analysis.
- torch.utils.checkpoint docs: the API and its reentrant vs non-reentrant variants.
Mixed precision: BF16, FP16, FP8, MX formats
The formats
| Format | Sign/Exp/Mantissa bits | Dynamic range | Use in training |
|---|---|---|---|
| FP32 | 1/8/23 | ~1e-38 to 3e38 | Master weights, optimizer state, reductions |
| TF32 | 1/8/10 (internal) | Same as FP32 | Tensor-core mode for "FP32" matmuls on Ampere+ |
| FP16 | 1/5/10 | ~6e-8 (subnormal) to 65504 | Older standard; needs loss scaling |
| BF16 | 1/8/7 | Same exponent range as FP32 | Today's default compute format |
| FP8 E4M3 | 1/4/3 | max 448 | Weights and activations (forward) |
| FP8 E5M2 | 1/5/2 | max 57344 | Gradients (wider range, less precision) |
| MXFP8 / MXFP6 / MXFP4 | FP8/6/4 elements + shared E8M0 scale per 32 elements | Per-block power-of-two scaling | Blackwell-native block-scaled matmuls |
| NVFP4 | E2M1 elements + E4M3 scale per 16 elements + per-tensor FP32 scale | Finer blocks, non-power-of-two scale | Inference, and emerging pretraining |
The FP8 E4M3/E5M2 split was standardized by NVIDIA, Arm and Intel Micikevicius+ 2022. The MX formats were defined in the OCP Microscaling spec Rouhani+ 2023. NVFP4's 16-element blocks with E4M3 scales are described in NVIDIA's format introduction NVIDIA 2025.
FP16 and loss scaling vs BF16
FP16 has only 5 exponent bits, so small gradient values (often below \(2^{-24}\)) underflow to zero. The fix is loss scaling: multiply the loss by a factor \(S\) (e.g. \(2^{16}\)) before backward, which scales every gradient by \(S\) through the chain rule, then divide by \(S\) before the optimizer step Micikevicius+ 2017. Dynamic loss scaling raises \(S\) periodically and halves it (skipping that step) whenever an inf/NaN appears.
BF16 keeps FP32's 8 exponent bits, so its range matches FP32 and underflow is rarely a problem. It needs no loss scaling. The cost is precision: 7 mantissa bits means about 2–3 significant decimal digits, so you need FP32 master weights (a weight of 1.0 plus an update of 1e-4 rounds back to 1.0 in BF16) and FP32 accumulation inside matmuls and reductions. Since A100, BF16 has been the default for LLM training.
FP8 training: scaling granularity is everything
FP8 roughly doubles tensor-core throughput over BF16 on Hopper and Blackwell and halves the memory traffic for the quantized tensors. The hard part is range: E4M3 only spans about 18 binades, so each tensor must be multiplied by a scale that maps its values into representable range. The design choices:
- Per-tensor delayed scaling: the scale comes from an amax history over the last N steps. It's fast, because there's no extra pass over the data, but a sudden outlier gets clipped. This was the original Transformer Engine recipe on H100 TE docs.
- Per-tensor current scaling: compute amax of the current tensor first. It's more robust, at the cost of an extra reduction.
- Fine-grained block scaling: DeepSeek-V3 used tile-wise 1×128 groups for activations (per token, per 128 channels) and 128×128 blocks for weights, with E4M3 everywhere. It promoted partial sums to FP32 on CUDA cores every 128 elements, because Hopper's FP8 tensor-core accumulation has limited precision DeepSeek-AI 2024. Outliers then only damage their own block.
- Hardware microscaling (MXFP8): Blackwell tensor cores natively consume per-32-element E8M0 scales, so fine-grained scaling costs almost nothing. NVIDIA reports MXFP8 pretraining of an 8B model on 15T tokens matching BF16 once the scale-factor rounding scheme is changed from the OCP-suggested one, which diverged (see the paper for the exact rule) Mishra+ 2025.
What stays in higher precision even in "FP8 training": master weights and optimizer state (FP32, sometimes BF16 moments), embeddings and the output head, normalization layers, the attention softmax, and often the last few layers. Only the big GEMMs run in FP8.
FP4 and the state of play in 2026
NVIDIA trained a 12B model on 10T tokens in NVFP4 with loss and downstream accuracy comparable to an FP8 baseline. It took four tricks: random Hadamard transforms to spread outliers, 2D (block) weight quantization so forward and backward see consistent values, stochastic rounding on gradients, and keeping a few sensitive layers in higher precision NVIDIA 2025.
| Practice | Maturity (as of late 2026) |
|---|---|
| BF16 compute + FP32 master/optimizer | Well established; the safe default everywhere |
| FP16 + dynamic loss scaling | Legacy; mostly on V100-era hardware |
| FP8 GEMMs (per-tensor or block scaling) on H100/H200 | Established in production at frontier labs (DeepSeek-V3 is the best public evidence); still needs care |
| MXFP8 on Blackwell | Recent; supported in Transformer Engine/Megatron, with NVIDIA-published recipes |
| NVFP4 / MXFP4 pretraining | Emerging and contested; demonstrated at 12B/10T scale, not yet a default |
Low-precision training moved fast in 2025–2026: Blackwell-native MX formats, NVFP4 pretraining results, and many open models reporting FP8 training. Check current Transformer Engine/Megatron docs and recent technical reports for what's default now, especially for FP4.
"Why BF16 over FP16?" Range: BF16 has no underflow issues and needs no loss scaling. "Why FP32 master weights?" Updates are often smaller than BF16's spacing near the weight value, so they'd be rounded away. "Why does FP8 need per-block scaling?" A few outlier activations (which LLMs do have) blow up a per-tensor amax, crushing everything else into a handful of FP8 values or into underflow. Block scaling confines the damage to one block.
- FP8 Formats for Deep Learning (Micikevicius+ 2022): why two FP8 formats exist.
- Transformer Engine FP8/FP4 primer: delayed, current, block, MXFP8 and NVFP4 recipes.
- Recipes for Pre-training LLMs with MXFP8 (Mishra+ 2025): practical MXFP8 gotchas.
The interconnect hierarchy
A modern training cluster is a hierarchy of increasingly slow links. The rough per-GPU numbers below are for orientation. Exact figures depend on SKU and configuration, so check the vendor spec before quoting them as fact.
| Level | Technology | Rough per-GPU bandwidth | Scope |
|---|---|---|---|
| On-chip | HBM3/3e | 3.35 TB/s (H100), ~8 TB/s (B200) | One GPU |
| Scale-up | NVLink 4 + NVSwitch (Hopper) | 900 GB/s total (≈450 GB/s each direction) | 8 GPUs (HGX node) |
| Scale-up | NVLink 5 (Blackwell), NVL72 rack | 1.8 TB/s total | Up to 72 GPUs in one NVLink domain |
| Scale-up | NVLink 6 (Vera Rubin) | ~3 TB/s+ (preliminary) | Rack-scale |
| Host link | PCIe Gen5 x16 | ~64 GB/s each direction | GPU ↔ CPU/NIC |
| Scale-out | InfiniBand NDR / RoCE 400G | 400 Gb/s ≈ 50 GB/s per NIC, usually 1 NIC per GPU | Across nodes, thousands of GPUs |
| Scale-out | InfiniBand XDR / 800G Ethernet | 800 Gb/s ≈ 100 GB/s per NIC | Newest clusters |
NVIDIA lists 900 GB/s per GPU for Hopper NVLink, 1.8 TB/s for Blackwell, and a preliminary figure for Vera Rubin's NVLink 6 NVIDIA NVLink. Llama 3 405B trained over RoCE at 400 Gb/s between GPUs, while smaller Llama 3 models used InfiniBand Llama Team 2024. DeepSeek-V3 quotes about 160 GB/s NVLink versus 50 GB/s InfiniBand on its export-restricted H800 cluster, roughly 3.2×, and designed its MoE routing around that gap DeepSeek-AI 2024.
The key ratio: intra-node bandwidth is about 10× inter-node bandwidth per GPU on a standard H100 HGX system. Rack-scale NVLink (GB200/GB300 NVL72) moves that cliff from 8 GPUs to 72, which matters a lot for TP and EP placement. Scale-out fabrics are usually rail-optimized fat trees: GPU i of every node connects to the same leaf "rail" switch, so collectives among same-index GPUs across nodes take one hop.
Think of a node (or NVL72 rack) as one "super-GPU" with a fast internal bus, and the cluster as a network of super-GPUs. Put communication that happens every layer (TP, sometimes EP/CP) inside the super-GPU, and communication that happens once per step or once per micro-batch (DP, PP) across super-GPUs.
Hardware generations change every 1–2 years (B200/B300, GB300 NVL72, Vera Rubin, 800G/1.6T networking, AMD MI-series with Infinity Fabric, TPU ICI). Treat the bandwidths above as order-of-magnitude and check current vendor spec sheets.
Collective communication
The primitives
| Collective | Input → output (N ranks, S bytes total) | Where it shows up |
|---|---|---|
| Broadcast | One rank's S → everyone | Initial weight sync |
| Reduce | Everyone's S → sum on one rank | Rare in training |
| All-reduce | Everyone's S → everyone has the sum | DDP grads, TP activations |
| Reduce-scatter | Everyone's S → each rank holds 1/N of the sum | ZeRO/FSDP grads, SP |
| All-gather | Each rank's S/N shard → everyone has all S | FSDP params, SP, CP (K/V) |
| All-to-all | Each rank sends a distinct chunk to every other rank | MoE dispatch/combine, Ulysses SP |
| Send/recv (P2P) | Point-to-point | Pipeline stages, ring attention |
All-reduce = reduce-scatter followed by all-gather. This identity is the foundation of ZeRO: if you only need your shard of the summed gradient, which you do once the optimizer state is sharded, you can stop after the reduce-scatter and save half the traffic for that phase.
Cost model: ring vs tree
Use the α–β model: sending \(m\) bytes costs \(\alpha + m/B\), where \(\alpha\) is per-message latency (a few µs) and \(B\) is link bandwidth. In a ring all-reduce, the buffer is split into N chunks. In N−1 reduce-scatter steps, each rank sends one chunk to its neighbor and adds the chunk it receives. Then N−1 all-gather steps circulate the finished chunks:
$$ T_{\text{ring all-reduce}} \approx 2(N-1)\,\alpha \;+\; 2\,\frac{N-1}{N}\,\frac{S}{B} $$The bandwidth term approaches \(2S/B\) and doesn't grow with N. That's bandwidth-optimal: every byte has to leave and arrive somewhere. The latency term grows linearly with N, which hurts at thousands of ranks and for small messages. Reduce-scatter and all-gather alone each cost half: \((N-1)\alpha + \frac{N-1}{N}\frac{S}{B}\).
Tree algorithms (NCCL uses a double binary tree, so both halves of the bandwidth are used) get latency of \(O(\log N)\) with near-optimal bandwidth for large messages. NCCL picks ring, tree or other algorithms (such as NVLink SHARP on NVSwitch systems, which does the reduction inside the switch) based on message size and topology.
nccl-tests reports algorithm bandwidth (S / time) and bus bandwidth (algBW × \(2(N-1)/N\) for all-reduce). Bus bandwidth is comparable to raw link speed regardless of N, so it's the number to compare against hardware peak nccl-tests docs.
Back-of-envelope: DDP gradient all-reduce for 7B
BF16 gradients: \(S = 14\) GB. Ring traffic per GPU is about \(2S = 28\) GB. Over InfiniBand at 50 GB/s per GPU that's ≈0.56 s per step if nothing overlaps. Inside one node over NVLink at about 450 GB/s, it's ≈0.06 s. If a step's compute takes 2–5 s, the cross-node all-reduce can be hidden behind backward. If the step is short (small per-GPU batch), it can't, and you start needing bigger batches, gradient accumulation or hierarchical reduction.
NCCL and friends
NCCL is NVIDIA's collective library. It detects the topology (NVLink, PCIe, NICs), builds rings and trees, and runs collectives as CUDA kernels on dedicated streams, which is why communication competes with compute for SMs. Equivalents: RCCL (AMD), Gloo (CPU), and XLA collectives over ICI on TPUs. PyTorch exposes them through torch.distributed process groups NCCL GitHub.
Thinking ring all-reduce time grows with the number of GPUs. Per-GPU bytes are essentially constant, \(2S(N-1)/N\). What grows is latency (with N) and the slowest link in the ring, which is usually the inter-node hop. Scaling DP from 64 to 1024 GPUs barely changes all-reduce bandwidth cost but can hurt latency and straggler sensitivity.
Expect "derive the cost of ring all-reduce" or "why is reduce-scatter + all-gather the same as all-reduce?" Strong answers write the α–β formula, note that bandwidth cost is independent of N, explain when trees win (small messages, large N), and mention that the bottleneck is the slowest link (inter-node), not the average one.
- nccl-tests PERFORMANCE.md: algBW vs busBW, with derivations per collective.
- NCCL user guide: environment variables, algorithms and debugging.
- JAX scaling book: chapters on collective costs and sharded matmuls.
Data parallelism (DDP)
Every rank holds a full replica of the model and processes a different slice of the global batch. After backward, gradients are averaged with all-reduce, so every replica applies the same update. The global batch is micro-batch × gradient-accumulation steps × DP degree.
PyTorch DDP makes this efficient in three ways Li+ 2020:
- Gradient bucketing: instead of one all-reduce per parameter (latency-bound) or one giant all-reduce at the end (no overlap), gradients are packed into buckets of about 25 MB by default. A bucket's all-reduce fires as soon as all its gradients are ready.
- Compute/communication overlap: backward produces gradients from the last layer to the first, and buckets are ordered roughly in reverse, so the all-reduce of late layers runs while early layers are still computing gradients. Only the last bucket's communication is exposed.
no_syncduring gradient accumulation: skip the all-reduce on intermediate micro-batches and sync only on the last one.
DDP's limit is memory: every GPU needs the full 16 bytes/param, so DDP alone tops out at roughly a 1–3B model on an 80 GB GPU once activations are included. Its strength is simplicity and the lowest communication volume per step of any approach that does something useful.
"Why not all-reduce each gradient as it's produced?" Too many small messages, all latency. "Why not one all-reduce at the end?" No overlap. Buckets balance the two. A bonus point: very large DP degrees push the global batch above the critical batch size, beyond which more data per step stops helping optimization. That's a statistical limit on DP, not a systems one.
ZeRO and FSDP: sharded data parallelism
Plain DP stores N identical copies of the optimizer state, gradients and parameters. ZeRO (Zero Redundancy Optimizer) removes that redundancy in three stages Rajbhandari+ 2019. With \(\Psi\) parameters, DP degree \(N_d\) and \(K=12\):
| Stage | Shards | Memory per GPU | Comm per step (per GPU, in units of Ψ elements) | Notes |
|---|---|---|---|---|
| DDP | Nothing | \((2+2+K)\Psi = 16\Psi\) | \(2\Psi\) (all-reduce) | Baseline |
| ZeRO-1 (Pos) | Optimizer state | \(4\Psi + K\Psi/N_d\) | \(2\Psi\) (RS grads + AG updated params) | Free lunch: same comm as DDP |
| ZeRO-2 (Pos+g) | + Gradients | \(2\Psi + (2+K)\Psi/N_d\) | \(2\Psi\) | Grads reduce-scattered as produced |
| ZeRO-3 (Pos+g+p) / FSDP | + Parameters | \(16\Psi/N_d\) | \(3\Psi\) (AG fwd + AG bwd + RS grads) | 1.5× DDP comm; params gathered per layer just in time |
How ZeRO-3 / FSDP runs a step
Prefetching the next unit's all-gather while the current one computes hides most of the cost when per-GPU compute is large enough. Keeping parameters gathered between forward and backward (reshard_after_forward=False, effectively ZeRO-2 for parameters) removes the second all-gather at the cost of memory. Llama 3 did exactly that: shard optimizer state and gradients, but don't reshard parameters after forward Llama Team 2024.
FSDP1 vs FSDP2, HSDP, and ZeRO variants
- FSDP (v1) flattened each wrapped module's parameters into one
FlatParameterand sharded that. It was efficient but awkward with per-parameter features (mixed requires_grad, per-param optimizers) and with torch.compile Zhao+ 2023. - FSDP2 (
fully_shard) shards each parameter individually along dim 0 as aDTensor. It composes cleanly with TP on a 2D device mesh, with torch.compile, and with FP8 all-gathers, and it's the recommended API in current PyTorch and torchtitan PyTorch fully_shard Liang+ 2024. - HSDP (hybrid sharding): shard within a node (or a group of nodes) and replicate across groups. Parameter all-gathers stay on fast links, and only a gradient all-reduce crosses the slow tier. This is the usual answer to "ZeRO-3 is too chatty across 1000 nodes."
- ZeRO++ cuts ZeRO-3 cross-node traffic with quantized weight all-gathers, a secondary in-node parameter copy, and quantized gradient reduce-scatter Wang+ 2023. ZeRO-Offload/Infinity pushes state to CPU or NVMe Rajbhandari+ 2021.
Saying ZeRO-3 is "model parallelism." It isn't: every GPU still runs the full computation of every layer on its own data. Parameters are only stored sharded and are gathered just in time. So ZeRO-3 doesn't reduce activation memory, and a single layer's full weights (plus prefetch buffers) still have to fit on one GPU.
Expect "what does each ZeRO stage shard, and what does it cost?" The crisp answer: Z1 shards optimizer state for free (RS + AG = all-reduce volume). Z2 also shards gradients for free. Z3 also shards parameters and costs 1.5× the communication, because parameters must be gathered in both forward and backward. Bonus: FSDP = ZeRO-3 in PyTorch, and HSDP fixes cross-node chattiness.
- PyTorch FSDP (Zhao+ 2023): design, rate limiting, prefetch and production lessons.
- FSDP2 fully_shard docs: the current API.
- ZeRO++ (Wang+ 2023): communication-reduction tricks for ZeRO-3 at scale.
Tensor parallelism (Megatron-style)
Tensor parallelism splits individual weight matrices across GPUs, so each GPU computes a slice of every layer Shoeybi+ 2019. The trick is to arrange the splits so that communication is needed only once per sub-block.
The MLP
The MLP computes \(Z = \mathrm{GeLU}(XA)\,B\). Split \(A\) by columns, \(A=[A_1, A_2]\). Each GPU computes \(Y_i = \mathrm{GeLU}(XA_i)\) independently, because GeLU is elementwise and the column split keeps whole output features together. No communication is needed here. (A row split of A would need a sum before the nonlinearity, since GeLU(a+b) ≠ GeLU(a)+GeLU(b).) Then split \(B\) by rows, \(B = [B_1; B_2]\), so \(Z = Y_1B_1 + Y_2B_2\). Each GPU computes its partial \(Y_iB_i\), and one all-reduce sums them.
Attention
Attention parallelizes naturally over heads. The Q, K, V projections are column-parallel (each GPU gets \(a/t\) heads), attention runs locally per head with no communication, and the output projection is row-parallel, followed by one all-reduce. With GQA, the KV heads are also divided across GPUs, which caps TP at the number of KV heads unless they're replicated.
Communication cost and why TP stays inside a node
Per transformer layer: 2 all-reduces in forward (one after attention, one after MLP) and 2 in backward, each of size \(b\cdot s\cdot h\) activations. These are synchronous and on the critical path: the next operation needs the result, so overlapping them is hard. (Recent work does overlap them by chunking the GEMM, but it's not free.) Rough ratio: a layer does about \(24\,bsh^2\) forward FLOPs and moves about \(2 \times 2 \times bsh\) bytes per all-reduce pair, so compute per byte scales with \(h\). With only about 50 GB/s per GPU across nodes, the all-reduces would take as long as the math. Over 450+ GB/s NVLink they're manageable. Hence the rule: TP ≤ GPUs per NVLink domain, typically 8 on HGX, potentially larger on NVL72.
TP's benefits: it shards parameters, gradients, optimizer state and activations (within the parallel regions) by \(t\), with no pipeline bubble. Its costs: the communication above, and smaller per-GPU GEMMs that run at lower efficiency as \(t\) grows. Megatron's PTD-P paper found TP within nodes plus PP across nodes the best combination, reaching about 52% of peak on 3072 A100s for a 1T model Narayanan+ 2021.
"Why is the first MLP matrix split by columns and the second by rows?" Column split lets the elementwise GeLU run without communication. Row split then produces partial sums that need exactly one all-reduce. Splitting the first by rows would force a sync before GeLU. Follow-up: "why not TP=64 across nodes?" Four critical-path all-reduces per layer over 10× slower links, plus tiny GEMMs.
Sequence parallelism (with TP)
Plain TP leaves some activations replicated: LayerNorm, dropout and residual adds operate on the full \(b\cdot s\cdot h\) tensor on every TP rank. That's the \(10\,sbh\)-ish part of the 34 that TP doesn't divide. Megatron sequence parallelism shards these regions along the sequence dimension instead Korthikanti+ 2022.
Each TP all-reduce becomes a reduce-scatter (into sequence shards) and an all-gather (before the next TP region). Since all-reduce = RS + AG, communication volume is unchanged, and all activations are now divided by \(t\): per-layer memory becomes \(\frac{sbh}{t}\left(34 + \frac{5as}{h}\right)\). In practice SP is always turned on with TP.
Confusing Megatron "sequence parallelism" (shards only the non-matmul regions, always paired with TP) with context parallelism or the older Colossal-AI "sequence parallelism" Li+ 2021 and DeepSpeed-Ulysses, which shard the sequence through attention too. Say which one you mean.
- Megatron-LM (Shoeybi+ 2019): the original TP layouts and f/g operators.
- Efficient Large-Scale Training with Megatron-LM (Narayanan+ 2021): PTD-P composition, interleaved pipelines, scaling to 1T.
- PyTorch Tensor Parallel docs: ColwiseParallel/RowwiseParallel/SequenceParallel on DTensor.
Pipeline parallelism
Pipeline parallelism assigns contiguous blocks of layers to different GPUs ("stages"). Stage \(i\) sends its output activations to stage \(i+1\) in forward and receives gradients back in backward. It's point-to-point communication of only \(b\cdot s\cdot h\) per micro-batch per boundary, which is tiny compared with TP's traffic. That's why PP is the parallelism you can stretch across slow inter-node links.
The problem is the bubble. With one batch, only one stage works at a time. GPipe splits the batch into \(m\) micro-batches and pipelines them Huang+ 2018. With \(p\) stages and per-micro-batch forward+backward time \(t\) per stage, the ideal time is \(m\,t\) and the fill/drain overhead is \((p-1)\,t\):
$$ \text{bubble fraction} = \frac{(p-1)}{m} \quad\text{(relative to ideal compute)}, \qquad \text{or}\ \ \frac{p-1}{m+p-1}\ \text{of total time} $$So you want \(m \gg p\), typically \(m \ge 4p\), to keep the bubble under about 20%. But \(m\) is bounded by global batch ÷ (DP × micro-batch size), which ties PP to the batch size. That's one reason very large DP plus deep PP becomes a problem at fixed global batch.
Schedules
| Schedule | Idea | Bubble | Activation memory per stage |
|---|---|---|---|
| GPipe | All forwards, then all backwards; synchronous flush | \((p-1)/m\) | \(m\) micro-batches in flight (bad) |
| 1F1B (PipeDream-Flush) Harlap+ 2018 | Warm up with \(p-1-i\) forwards, then alternate one forward, one backward | \((p-1)/m\) | At most \(p\) in flight, independent of \(m\) |
| Interleaved 1F1B Narayanan+ 2021 | Each GPU holds \(v\) non-contiguous chunks of layers ("virtual stages") | \((p-1)/(v\,m)\) | Slightly more; \(v\)× more P2P messages |
| Zero Bubble (ZB-H1/H2) Qi+ 2023 | Split backward into B (grad w.r.t. input, on critical path) and W (grad w.r.t. weights, deferrable) and use W to fill bubbles | Roughly 1/3 of 1F1B (H1) down to near zero (H2) | H1 ≈ 1F1B; H2 needs more |
| DualPipe (DeepSeek-V3) DeepSeek-AI 2024 | Bidirectional: feed micro-batches from both ends; overlap MoE all-to-all and PP comm inside forward/backward chunk pairs | Reduced vs 1F1B / ZB | Keeps 2 copies of params per stage |
Practical PP headaches: stage balancing (the embedding layer and the large-vocabulary output head and loss make the first and last stages heavier, so frameworks put fewer transformer blocks there), the dependence on enough micro-batches, more complex code (schedules, P2P deadlocks), and the interaction with gradient accumulation and FSDP (re-gathering parameters for every micro-batch is wasteful, so PP is usually combined with ZeRO-1 rather than ZeRO-3).
"Derive the bubble fraction." Show fill and drain of \(p-1\) slots against \(m\) useful slots. "What does 1F1B buy over GPipe?" The same bubble, but activation memory bounded by \(p\) instead of \(m\). "How do you shrink the bubble?" More micro-batches, interleaving (paying more communication), or zero-bubble scheduling (splitting backward into input-grad and weight-grad). Saying "1F1B removes the bubble" is a red flag.
- GPipe (Huang+ 2018): micro-batching and re-materialization.
- Zero Bubble Pipeline Parallelism (Qi+ 2023): the B/W split.
- torch.distributed.pipelining: PyTorch-native schedules (GPipe, 1F1B, interleaved, ZB).
Context parallelism and ring attention
At 128K+ tokens, even one sequence's activations overflow a GPU after TP+SP. Context parallelism (CP) shards the sequence across \(c\) GPUs for the entire layer, attention included. Everything except attention is token-local (MLP, norms, projections), so it needs no communication. Attention needs every query to see every earlier key and value.
Ring attention
Each GPU holds its query chunk \(Q_i\) and its K/V chunk. Over \(c\) steps, K/V chunks rotate around a ring. At each step a GPU computes blockwise attention of \(Q_i\) against the K/V chunk it currently holds and merges the result using FlashAttention's online-softmax rescaling (running max and normalizer). Sending the next K/V chunk overlaps with computing the current one Liu+ 2023. Communication is hidden as long as computing one block (∝ \((s/c)^2 h\)) takes longer than transferring a K/V chunk (∝ \((s/c)\,h_{kv}\)), which gets easier as chunks grow.
Causal load imbalance: with a causal mask, the GPU holding the last chunk attends to everything, while the first attends almost nothing. The standard fix is to split the sequence into \(2c\) pieces and give rank \(i\) pieces \(i\) and \(2c-1-i\) (zig-zag), so every rank does the same work. Llama 3 used this kind of 2×CP chunking, but chose an all-gather of K/V instead of a ring. With GQA, K/V are much smaller than Q, and all-gather makes document masks easier to support Llama Team 2024.
DeepSpeed-Ulysses (all-to-all sequence parallelism)
An alternative: keep activations sequence-sharded outside attention, then use an all-to-all to re-shard from "all heads, s/c tokens" to "a/c heads, all tokens", run ordinary attention per head group, and all-to-all back Jacobs+ 2023. Per-GPU communication stays constant when \(s\) and \(c\) grow together. The degree is limited by the number of (KV) heads. Hybrids (Ulysses inside a node, ring across nodes) are common.
| Ring attention (P2P) | All-gather K/V | Ulysses (all-to-all) | |
|---|---|---|---|
| Comm pattern | c−1 neighbor sends of K/V, overlapped | One all-gather of K/V per layer | 2 all-to-alls per attention |
| Scales past #heads? | Yes | Yes | No (≤ heads) |
| Peak memory | Low (one remote chunk at a time) | Full K/V gathered | Full sequence for a/c heads |
| Best when | Very long sequences, slower links | GQA (small K/V), complex masks | Fast all-to-all, many heads |
"How would you train with 1M-token context?" Strong answer: FlashAttention (no \(s^2\) memory), TP+SP inside the node, CP across GPUs with ring or all-gather K/V and zig-zag load balancing, aggressive recomputation, and a progressive context-extension schedule (most tokens at short context, a final phase at long context), not 1M from step one.
- Ring Attention (Liu+ 2023): blockwise ring computation.
- DeepSpeed Ulysses (Jacobs+ 2023): the all-to-all approach.
- ring-flash-attention: a readable implementation, including zig-zag variants.
Expert parallelism for MoE
In a Mixture-of-Experts layer, a router sends each token to its top-\(k\) of \(E\) expert MLPs. Expert parallelism places different experts on different GPUs Lepikhin+ 2020. Each MoE layer then needs:
Per-GPU volume per MoE layer is roughly tokens × k × h × bytes, each way. Unlike all-reduce, all-to-all doesn't benefit from ring pipelining. Every pair talks, so it's sensitive to bisection bandwidth and to imbalance. Key issues:
- Load imbalance: a popular expert becomes a straggler. Fixes include auxiliary load-balancing losses, capacity factors that drop overflow tokens (Switch Transformer Fedus+ 2021), and DeepSeek-V3's auxiliary-loss-free bias adjustment.
- Topology-aware routing: DeepSeek-V3 limits each token to at most 4 nodes, sends cross-node traffic over IB once per node, then forwards over NVLink. It used 64-way EP across 8 nodes and overlapped all-to-all with compute using DualPipe DeepSeek-AI 2024.
- Composition: attention layers are dense and typically use DP/TP, while MoE layers use EP. Often the EP group is carved out of the DP group, so the same GPUs act as DP replicas for attention and as EP shards for experts.
MoE converts parameter count into communication. You get a big model's capacity at a small model's FLOPs per token, but every MoE layer becomes an all-to-all shuffle. On fast rack-scale NVLink (NVL72), wide EP becomes much cheaper, which is part of why MoE and rack-scale systems are co-evolving.
- GShard (Lepikhin+ 2020): expert parallelism and automatic sharding.
- DeepSeek-V3 technical report: section 3 on infrastructure (DualPipe, all-to-all kernels, FP8).
- PyTorch blog: training MoEs at scale: practical EP/FSDP composition.
Composing 3D/4D/5D parallelism
Real runs combine several dimensions on a device mesh: a logical N-dimensional array of ranks where each axis is one parallelism type. The total GPU count is the product: \(N = d_{dp} \times d_{pp} \times d_{cp} \times d_{tp}\) (× EP folded into DP for MoE layers). The placement rule: the most communication-intensive dimension goes innermost, on the fastest links.
Case study: Llama 3 405B
Meta trained Llama 3 405B with 4D parallelism ordered [TP, CP, PP, DP] from innermost to outermost, with DP implemented as FSDP that shards optimizer state and gradients but doesn't reshard parameters after forward Llama Team 2024:
| GPUs | TP | CP | PP | DP | Seq len | Tokens/batch | TFLOP/s per GPU | BF16 MFU |
|---|---|---|---|---|---|---|---|---|
| 8,192 | 8 | 1 | 16 | 64 | 8,192 | 16M | 430 | 43% |
| 16,384 | 8 | 1 | 16 | 128 | 8,192 | 16M | 400 | 41% |
| 16,384 | 8 | 16 | 16 | 8 | 131,072 | 16M | 380 | 38% |
Things to notice. TP=8 matches the 8-GPU NVLink node. PP=16 spans nodes. Going from 8K to 16K GPUs at a fixed 16M-token batch halved the per-DP-rank batch and cost about 2 points of MFU. For the long-context phase, CP=16 was carved out of DP (DP dropped from 128 to 8) to keep the global batch in tokens constant. Sanity check (using the reported ~15.6T training tokens): \(6 \times 405\text{B} \times 15.6\text{T} \approx 3.8\times10^{25}\) FLOPs; at about 400 TFLOP/s × 16,384 GPUs that's about \(5.8\times10^6\) s, roughly 67 days of perfect running. That's the right order of magnitude for a frontier pretraining run.
Case study: DeepSeek-V3 (671B MoE)
DeepSeek-V3 is a useful contrast: no tensor parallelism at all. It used 16-way PP, 64-way EP across 8 nodes and ZeRO-1 DP on H800s. Heavy engineering on all-to-all kernels and DualPipe made TP unnecessary, and the whole run took 2.788M H800 GPU-hours DeepSeek-AI 2024. The lesson: the "right" mesh depends on the model (dense vs MoE), the interconnect (H800's reduced NVLink) and engineering investment, not on a fixed formula.
How to choose a parallelism strategy
A practical decision procedure, roughly following the Ultra-Scale Playbook's advice HF Playbook 2025:
- Fit one replica. Can model state plus activations at micro-batch 1 fit on one GPU? If yes, use DDP (or ZeRO-1/2 to free memory for bigger micro-batches). If no, go to step 2.
- Shard state first. Use FSDP/ZeRO-3, or HSDP across many nodes. It's the least invasive option, with no model code changes.
- If a single layer or the activations still don't fit, or FSDP communication is exposed at large scale, add TP+SP inside the node (≤8 on HGX).
- If the model still doesn't fit across one node's TP group, or DP has grown so large that per-rank batch is tiny, add PP across nodes.
- For long sequences, add CP. For MoE, add EP.
- Then tune for throughput: micro-batch size, recompute policy, bucket sizes, overlap. Measure MFU and exposed communication with a profiler.
| Situation | Recommended starting point | Why |
|---|---|---|
| ≤ ~1B params, any GPU count | DDP (+ ZeRO-1) | Fits easily; lowest comm; simplest |
| 1–13B, up to a few hundred GPUs | FSDP2 (ZeRO-3) or ZeRO-2, + activation checkpointing as needed | State sharded; no model surgery; comm overlaps well |
| 30–70B dense, hundreds of GPUs | FSDP/HSDP + TP=8 inside node (+SP); or TP=8 + PP=2–4 + ZeRO-1 | A single layer is large; TP cuts activations and state; FSDP across nodes |
| 100B–1T dense, thousands of GPUs | TP=8 (in node) × PP (8–16, across nodes) × DP/FSDP; interleaved or ZB schedule | Too large for FSDP alone; PP tolerates slow links |
| Long context (≥64K) | Add CP (ring / all-gather K/V), zig-zag balancing; FlashAttention; selective recompute | Activations ∝ s; only CP divides attention-region activations |
| MoE | EP (in node or few nodes) + DP/FSDP for dense parts; PP across; TP optional | Experts dominate params; all-to-all must stay on fast links |
| Slow interconnect (no IB, cloud Ethernet) | Favor PP + HSDP; avoid cross-node TP; larger micro-batches; gradient compression | Minimize per-layer cross-node traffic |
| TPU pods | JAX/XLA GSPMD sharding over a 2–3D mesh (data, fsdp, model) | ICI torus favors sharding annotations over hand-written PP |
Maximizing every dimension "because it scales." Each dimension adds overhead: TP adds critical-path all-reduces and shrinks GEMMs, PP adds bubbles and complexity, CP adds attention communication, ZeRO-3 adds 50% traffic. The goal is the smallest model-parallel footprint that fits, with everything else as DP.
Design questions ("train a 400B model on 4,096 H100s") reward a structured answer: (1) model state ≈ 6.4 TB, so minimum sharding is about 80 GPUs before activations. (2) TP=8 in node. (3) PP=8–16 to reach a per-GPU footprint that fits with activations. (4) The remainder is DP, and check that micro-batches per pipeline satisfy \(m \ge 4p\) given the global batch. (5) Estimate time from \(6ND\) and an assumed 40% MFU. (6) Plan for failures and checkpointing.
- The Llama 3 Herd of Models: section 3.3 is a superb public account of 16K-GPU infrastructure.
- Ultra-Scale Playbook: the "finding the best configuration" chapter.
- TorchTitan (Liang+ 2024): composable FSDP2 + TP + PP + CP in PyTorch.
MFU and HFU: measuring efficiency
Model FLOPs Utilization (MFU), popularized by PaLM, is the ratio of the FLOPs the model needs per second at your observed throughput to the hardware's peak Chowdhery+ 2022:
$$ \text{MFU} = \frac{\text{tokens/s} \times F_{\text{token}}}{N_{\text{GPU}} \times \text{peak FLOP/s}}, \qquad F_{\text{token}} \approx 6N + 12\,L\,h\,s $$The \(6N\) term covers the parameter matmuls (forward + backward). The second term approximates attention's \(QK^\top\) and \(AV\) FLOPs, which grow with sequence length. Some definitions halve it for causal masking, so state which one you use. Hardware FLOPs Utilization (HFU) also counts FLOPs actually executed, such as recomputation, so HFU ≥ MFU. With full recomputation, HFU is about 8/6 of MFU. MFU is the honest metric, because it can't be inflated by doing redundant work.
| Setting | Reported utilization |
|---|---|
| PaLM 540B on 6144 TPU v4 | 46.2% MFU (57.8% HFU) |
| Megatron PTD-P 1T on 3072 A100 | ~52% of peak |
| Llama 3 405B on 8K–16K H100 (BF16) | 38–43% MFU |
| Typical well-tuned dense LLM, H100, BF16 | 35–55% |
| MoE / long context / small models | Usually lower (all-to-all, attention, small GEMMs) |
Where the missing 50–60% goes: exposed communication (TP all-reduces, FSDP gathers, the last DDP bucket), pipeline bubbles, memory-bound kernels (norms, activations, optimizer step, attention softmax), kernel launch and Python overhead, stragglers and synchronization, data loading, and the fact that marketed peak FLOP/s (and the power and thermal limits behind them) aren't sustainable on real GEMM shapes. Note that MFU computed against FP8 peak looks worse than against BF16 peak for the same wall-clock speedup, so always state the reference precision.
"Your run is at 25% MFU. How do you debug it?" Profile one step (PyTorch profiler / Nsight Systems). Look for exposed communication gaps between compute kernels, the bubble size, GPU idle time waiting on the data loader, and memory-bound kernels you could fuse. Then check that GEMM shapes are tensor-core friendly (multiples of 64/128), that recompute isn't overused, and that micro-batch size isn't tiny. Name concrete fixes: larger buckets, prefetching, fewer TP ranks, more micro-batches, torch.compile, fused kernels.
Fault tolerance at scale
At 16K GPUs, failures are routine. During a 54-day snapshot of Llama 3 405B pretraining, Meta saw 466 job interruptions, 419 of them unexpected, with about 78% of the unexpected ones attributed to confirmed hardware issues, faulty GPUs and HBM being the largest categories. Even so, it kept effective training time above 90% Llama Team 2024. That's roughly one interruption every 3 hours. Because synchronous training is all-or-nothing, one bad GPU stops all 16K.
Failure modes
- Hard failures: GPU falls off the bus, HBM ECC errors, NIC/link flaps, host crashes. Usually surfaced as NCCL timeouts or hangs.
- Stragglers: a slow GPU (thermal throttling, a bad link) slows everyone, because every collective waits for the slowest rank.
- Silent data corruption (SDC): wrong arithmetic with no error. Detected via loss spikes, gradient-norm anomalies or cross-replica checksums.
- Numerical instability: loss spikes and divergence. Typical response is to roll back to a checkpoint and skip or reorder data batches.
Checkpointing
The optimal checkpoint interval follows the classic Young/Daly approximation: \(\tau^* \approx \sqrt{2\,C\,M}\), where \(C\) is the time to write a checkpoint and \(M\) is the job's mean time between failures. Example: \(C = 1\) min and \(M = 3\) h gives \(\tau^* \approx \sqrt{2 \cdot 1 \cdot 180} \approx 19\) min. Expected loss per failure is about half an interval plus restart time. Engineering moves:
- Sharded (distributed) checkpoints: each rank writes only its shard in parallel, with metadata that allows resharding on load to a different mesh (PyTorch DCP) PyTorch DCP. A 405B checkpoint with optimizer state is several TB.
- Asynchronous checkpointing: copy state to CPU memory quickly (training blocks only for the device-to-host copy), then persist to storage in the background PyTorch blog.
- In-memory / peer checkpoints: keep recent snapshots in other hosts' RAM for near-instant recovery, with periodic durable checkpoints to object storage.
- Fast restart: hot-spare nodes, health checks before job start, automated bad-node exclusion, and cached NCCL setup. Restart overhead (often minutes) can matter more than checkpoint write time.
Elastic and fault-tolerant training
Elastic training lets the job continue with fewer or more workers instead of waiting for a replacement. torchrun's elastic mode re-forms the process group within a min:max node range torchrun docs. Newer approaches such as torchft tolerate failures at step granularity by treating DP replica groups as independently recoverable: a failed replica drops out of the gradient all-reduce and rejoins later, while the others keep training torchft. The catch is that changing DP size changes the effective batch, and model-parallel groups (TP/PP) still can't lose a member without reconfiguring.
Fault-tolerance tooling (torchft, in-memory checkpointing, semi-synchronous approaches like DiLoCo-style local SGD across islands) is an active 2025–2026 area. Check current PyTorch, Megatron and lab reports for production practice.
"How often should you checkpoint a 10K-GPU job?" Estimate the job MTBF (per-GPU failure rate × GPU count, so hours, not weeks), estimate checkpoint cost (TB of state ÷ aggregate storage bandwidth, or seconds with async), apply \(\sqrt{2CM}\), then discuss restart time, hot spares and goodput (useful training time ÷ wall time) as the real metric.
- Llama 3 paper §3.3.4: reliability statistics and operational challenges.
- PyTorch Distributed Checkpoint: sharded save/load and resharding.
- torchft: per-step fault tolerance for PyTorch.
Frameworks
| Framework | What it is | Strengths | Notes |
|---|---|---|---|
| PyTorch native (DDP, FSDP2, DTensor, TP, pipelining, DCP) | Composable building blocks on DeviceMesh + DTensor | Standard, compiles with torch.compile, flexible | You assemble the pieces yourself |
| torchtitan | PyTorch's reference LLM pretraining stack GitHub | FSDP2 + TP/SP + PP + CP (+ EP), Float8, async checkpointing, clean code | Good for learning how the native pieces compose |
| Megatron-LM / Megatron-Core | NVIDIA's high-performance training library GitHub | Best-in-class TP/SP/PP/CP/EP on NVIDIA, FP8/MXFP8 via Transformer Engine | Most frontier-scale open recipes build on it |
| DeepSpeed | Microsoft's library GitHub | ZeRO 1/2/3, offload/Infinity, ZeRO++, Ulysses | Easy to bolt onto existing models; common for fine-tuning |
| Transformer Engine | NVIDIA kernels + FP8/FP4 recipes GitHub | FP8 delayed/current/block scaling, MXFP8, NVFP4 | Used by Megatron, torchtitan, others |
| JAX / XLA (GSPMD, shard_map) | Annotate array shardings over a Mesh; the compiler inserts collectives JAX docs | Very concise; same code from 1 to thousands of chips; TPU-native | Used by Google and others (MaxText-style stacks) |
In JAX, you write a single-device program, place arrays on a named mesh (e.g. axes ('data', 'fsdp', 'model')) with PartitionSpecs, and the XLA SPMD partitioner propagates shardings and inserts all-gathers, reduce-scatters and all-reduces. That's the same set of collectives you'd write by hand in Megatron. shard_map gives manual per-device control when the compiler's choices aren't good enough.
Framework capabilities change quickly (FSDP2 replacing FSDP1, torchtitan's EP and FP8 support, Megatron-Core MoE features, DeepSpeed's moves). Verify current feature support in each project's docs before claiming it in an interview.
Interview question bank
How much GPU memory do you need to train a 13B model with Adam in mixed precision, and how would you fit it on 8×80 GB?
Model state is about 16 bytes/param × 13B ≈ 208 GB (234 GB with FP32 grads), before activations. That exceeds one 80 GB GPU but is well under 640 GB in total. With ZeRO-3/FSDP over 8 GPUs, state is about 26 GB per GPU, leaving about 50 GB for activations, buffers and fragmentation. Alternatively, ZeRO-2 gives 2·13 + 14·13/8 ≈ 49 GB, which is tighter. Then size activations: for h=5120, L=40, s=4096 with FlashAttention, about 34·s·h bytes ≈ 0.7 GB per layer per sequence, about 28 GB per micro-batch of 1, so micro-batch 1 (or activation checkpointing for bigger micro-batches). A strong answer states its assumptions (gradient precision, FlashAttention) and leaves 10–15% headroom.
Where does the "16 bytes per parameter" number come from? When is it different?
2 bytes of BF16 weights + 2 bytes of BF16 gradients + 12 bytes of FP32 optimizer state (4 for the master weights, 4 for Adam's first moment, 4 for the second). It's 18 if gradients are accumulated in FP32. It's less with 8-bit optimizers, BF16 moments, Adafactor-style factored second moments, or single-buffer optimizers like Muon or SGD with momentum. Inference only needs the weights (2 bytes/param in BF16, plus the KV cache). Activations come on top and are often the larger term at long context.
Explain ZeRO stages 1, 2 and 3. What does each cost in communication?
All three are data parallelism that removes redundant copies of model state. Stage 1 shards optimizer state: gradients are reduce-scattered, each rank updates its shard of the master weights, and the updated BF16 weights are all-gathered. RS + AG equals one all-reduce, so communication is the same as DDP (2Ψ). Stage 2 also shards gradients: each rank keeps only its reduce-scattered slice, still 2Ψ. Stage 3 (FSDP) also shards parameters, which must be all-gathered before each layer's forward and again in backward, plus a gradient reduce-scatter: 3Ψ, 1.5× DDP. Memory per GPU goes 16Ψ → 4Ψ+12Ψ/N → 2Ψ+14Ψ/N → 16Ψ/N. ZeRO doesn't shard activations.
Derive the communication cost of ring all-reduce. Why is it called bandwidth-optimal?
Split the S-byte buffer into N chunks. Reduce-scatter takes N−1 steps in which each rank sends one chunk of S/N to its neighbor and adds the chunk it receives. After that, each rank owns one fully reduced chunk. All-gather takes another N−1 steps to circulate the finished chunks. Each rank sends 2(N−1)·S/N bytes, so time ≈ 2(N−1)α + 2(N−1)/N · S/B. The bandwidth term approaches 2S/B independent of N, and any all-reduce algorithm must move at least about that much per rank, which is why it's bandwidth-optimal. The latency term grows linearly with N, so tree algorithms (O(log N) latency) win for small messages and huge N.
Back-of-envelope: how long does a gradient all-reduce for a 70B model take across nodes, and can it be hidden?
BF16 gradients are 140 GB. Ring all-reduce moves about 2× that per GPU, about 280 GB. At about 50 GB/s per GPU over 400G InfiniBand, that's about 5.6 s if done naively in pure DP. That's why nobody runs pure DDP for 70B: TP=8 shrinks each GPU's gradient to 17.5 GB (about 0.7 s), and with PP the per-stage gradient shrinks further. Hiding it requires backward compute per step to exceed the communication time, which depends on tokens per GPU per step. You'd also use bucketing, overlap, HSDP (in-node reduce-scatter first), or gradient accumulation to amortize.
Walk through Megatron tensor parallelism for the MLP and attention. How many collectives per layer?
MLP: the first weight A is split by columns, so each GPU computes GeLU(X·A_i) independently (GeLU is elementwise). The second weight B is split by rows, so each GPU produces a partial Y_i·B_i and one all-reduce sums them. Attention: Q/K/V projections are column-split by heads, each GPU computes attention for its heads locally, and the output projection is row-split, followed by one all-reduce. So 2 all-reduces per layer in forward and 2 in backward (the conjugate f operator all-reduces input gradients). Each is b·s·h activations and sits on the critical path, which is why TP lives inside an NVLink domain.
Why does TP usually stop at 8? When might it go higher?
TP needs 4 synchronous all-reduces of activation-sized tensors per layer. Inside an HGX node, NVLink gives about 450 GB/s per direction. Across nodes you get about 50 GB/s, roughly 10× slower, so cross-node TP makes communication dominate. Larger TP also shrinks each GPU's GEMMs (lower tensor-core efficiency), and GQA's small number of KV heads caps head-splitting. TP can go higher on rack-scale NVLink domains (GB200/GB300 NVL72), or for very wide models where the GEMMs stay large.
What is Megatron sequence parallelism and why is it "free"?
Under TP, LayerNorm, dropout and residual operations still see the full b·s·h tensor on every TP rank, so those activations are replicated. SP shards those regions along the sequence dimension. The all-reduce at the end of each TP region becomes a reduce-scatter (producing sequence shards), and an all-gather restores the full sequence before the next TP region. Since all-reduce = RS + AG, communication volume is unchanged, while activation memory in those regions drops by t. It's always enabled alongside TP in modern stacks.
Derive the pipeline bubble fraction. How do 1F1B, interleaving and zero-bubble schedules change it?
With p stages and m micro-batches each taking t (forward+backward) per stage, useful work per stage is m·t, but filling and draining the pipeline adds (p−1)·t idle time, so bubble/ideal = (p−1)/m. 1F1B doesn't change the bubble but bounds in-flight activations at p instead of m. Interleaving with v virtual stages per GPU cuts the bubble to (p−1)/(v·m) at the cost of v× more P2P messages. Zero-bubble schedules split backward into input-gradient (critical) and weight-gradient (deferrable) work and use the latter to fill the bubbles, approaching zero bubble at some memory cost.
You have PP=16 and a global batch of 4M tokens at sequence length 4096 with DP=64. Is the bubble acceptable?
4M / 4096 ≈ 1024 sequences per step. Divided over DP=64 that's 16 sequences per pipeline. With micro-batch size 1, m = 16, so bubble/ideal = 15/16 ≈ 94%: terrible. You'd want m ≥ 4p = 64, which needs a larger global batch (often not allowed for optimization reasons), smaller DP (more TP or more PP per replica changes the math), interleaving (v=4 gives about 23%), or zero-bubble schedules. This is the real tension of PP at scale: bubbles depend on micro-batches per replica, and DP growth shrinks that number. Llama 3 lost some MFU going from DP=64 to DP=128 at fixed batch for this kind of reason.
How does ring attention work, and what is the causal load-imbalance problem?
The sequence is split into c chunks, one per GPU. Each GPU keeps its query chunk and passes K/V chunks around a ring. In each of c steps it computes attention of its queries against the current K/V chunk while the next chunk is in flight, merging partial results with online-softmax rescaling (running max and sum, as in FlashAttention). With causal masking, later chunks attend to more keys than earlier ones, so the last GPU does about c× the work of the first. The fix is zig-zag sharding: split into 2c pieces and give rank i pieces i and 2c−1−i, so every rank gets an equal share of the triangle.
Ring attention vs DeepSpeed-Ulysses: when would you pick each?
Ulysses uses two all-to-alls per attention to switch from sequence-sharded to head-sharded layout, then runs ordinary full-sequence attention per head group. It's simple and efficient with fast all-to-all, but its degree is capped by the number of (KV) heads. Ring attention uses neighbor P2P with overlap, scales beyond the head count, and has low peak memory, but needs large enough chunks to hide communication and careful load balancing. All-gathering K/V (Llama 3's choice) is simplest when GQA makes K/V small. Hybrids (Ulysses inside a node, ring across nodes) are common.
What communication does expert parallelism add, and how do you keep it from dominating?
Each MoE layer does an all-to-all to dispatch tokens to the GPUs holding their top-k experts, local expert computation, and an all-to-all to combine results back: two in forward and two in backward. Volume is roughly tokens × k × hidden × bytes each way. To keep it manageable: keep the EP group on fast links (in node or rack), limit how many nodes a token can route to (DeepSeek-V3: at most 4), deduplicate cross-node sends, overlap all-to-all with compute (DualPipe, micro-batch overlap), balance load (aux loss or bias-based balancing, capacity factors), and use FP8 for dispatch payloads.
BF16 vs FP16: why did the field move to BF16, and is loss scaling still needed?
FP16 has 5 exponent bits (max 65504, subnormals down to about 6e-8), so small gradients underflow and large activations can overflow. It needs dynamic loss scaling: multiply the loss by S before backward, unscale before the update, and back off on inf/NaN. BF16 has 8 exponent bits, the same range as FP32, so underflow and overflow are basically gone and no loss scaling is needed. The cost is only 7 mantissa bits, which is why master weights, optimizer state and matmul accumulation stay in FP32. Since A100/TPU support, BF16 has been the default for LLMs.
How does FP8 training work, and why does scaling granularity matter?
The large GEMMs run in FP8, typically E4M3 for weights and activations and E5M2 (or E4M3) for gradients, with FP32 accumulation. Master weights, optimizer state, norms, softmax, embeddings and the LM head stay in higher precision. Because FP8's range is tiny, each tensor (or block) is multiplied by a scale chosen from its absolute max. Per-tensor delayed scaling uses an amax history and is fast but fragile with outliers. Finer granularity contains outliers: DeepSeek-V3 used 1×128 tiles for activations and 128×128 blocks for weights, and Blackwell's MXFP8 has hardware per-32-element power-of-two scales. Gains are up to about 2× matmul throughput and less memory traffic, with risk of divergence if the recipe is wrong.
What are MX formats and NVFP4, and where do they stand in 2026?
The OCP Microscaling (MX) formats store blocks of 32 low-precision elements (FP8, FP6 or FP4) with one shared 8-bit power-of-two (E8M0) scale, so per-block scaling is built into the data type and executed natively by Blackwell tensor cores. NVFP4 uses FP4 (E2M1) elements with smaller 16-element blocks, an E4M3 (non-power-of-two) block scale and a per-tensor FP32 scale, for better accuracy at 4 bits. MXFP8 pretraining has been shown to match BF16 at 8B/15T scale with the right rounding. NVFP4 pretraining matched FP8 at 12B/10T with Hadamard transforms, stochastic rounding and some high-precision layers. MXFP8 is becoming practical; FP4 pretraining is recent and not yet a default. Check current reports.
Define MFU vs HFU. Estimate the MFU of a run that processes 3,000 tokens/s per H100 for a 70B dense model.
MFU = (model FLOPs needed per second at the observed throughput) / peak FLOP/s. HFU also counts executed-but-redundant work such as recomputation, so HFU ≥ MFU. For 70B: about 6 × 70e9 = 4.2e11 FLOPs/token (ignoring attention). × 3,000 tokens/s = 1.26e15 FLOP/s per GPU, which is more than H100's ~989 TFLOP/s dense BF16 peak, so the number is impossible. A plausible figure is around 900–1,300 tokens/s per GPU, which gives about 40–55% MFU. Pointing out that a quoted throughput is physically impossible is exactly what interviewers want.
Estimate how long it takes to pretrain a 70B model on 15T tokens with 4,096 H100s.
Total compute ≈ 6 × 70e9 × 15e12 = 6.3e24 FLOPs. Effective per-GPU throughput at 40% MFU on 989 TFLOP/s BF16 ≈ 400 TFLOP/s. Cluster: 4,096 × 4e14 = 1.6e18 FLOP/s. Time ≈ 6.3e24 / 1.6e18 ≈ 3.9e6 s ≈ 45 days of perfect running. Add about 10% for failures, restarts and evaluation, so roughly 7 weeks. Mention the sensitivity: FP8 could shorten it, and lower MFU or poor goodput lengthens it proportionally.
Design: you must train a 70B dense model with 32K context on 512 H100s (64 nodes). Choose a parallelism layout.
State is about 1.12 TB, so sharding is required. Activations at 32K are large: with FlashAttention, about 34·s·h ≈ 34 × 32768 × 8192 ≈ 9 GB per layer per sequence before TP, times 80 layers, so TP+SP and probably CP are needed. Proposal: TP=8 inside each node (with SP) divides activations and weights by 8. CP=2 across a node pair halves the sequence per GPU. The remaining 32-way dimension is FSDP/HSDP over nodes (or PP=4 × ZeRO-1 DP=8 if FSDP gathers prove exposed). Add selective or full recomputation as needed to fit. Check memory: params/grads/opt per GPU ≈ 1.12 TB/(8·32) ≈ 4.4 GB with ZeRO-3 over DP. Activations per layer ≈ 9 GB/(8·2) ≈ 0.56 GB, × 80 ≈ 45 GB without recompute, so recompute some layers. Then validate with a profiler and tune micro-batching.
FSDP vs tensor parallelism: both shard parameters, so what's the real difference?
FSDP shards parameter storage but gathers full weights before computing, so every GPU does the full layer computation on its own data. Its communication is parameter-sized, happens once per layer per pass, and can be prefetched and overlapped. It doesn't reduce activations. TP shards the computation itself: each GPU computes a slice of every matmul, which also divides activations, but it requires activation-sized all-reduces on the critical path every layer. FSDP scales well across nodes when per-GPU batch is large enough. TP needs NVLink. They compose naturally as a 2D mesh (TP inside the node, FSDP across nodes).
How do you choose the checkpoint interval for a 16K-GPU job, and what else matters for goodput?
Estimate MTBF from observed failure rates. Llama 3 saw 419 unexpected interruptions in 54 days, about one every 3 hours. Estimate checkpoint cost C: synchronous writes of several TB might take minutes, async checkpointing blocks training only for the GPU-to-CPU copy. Young/Daly gives τ ≈ √(2CM): with C = 1 min and M = 180 min, τ ≈ 19 min. Beyond that, goodput depends on detection time (NCCL timeouts can be long, so use health checks and watchdogs), restart time (scheduling, node replacement, NCCL init, loading checkpoints), hot spares, and avoiding repeated rollbacks for loss spikes. Elastic or replica-level fault tolerance (torchft) can avoid full restarts for DP replica failures.
Why do DDP implementations bucket gradients, and what determines a good bucket size?
An all-reduce per parameter tensor means thousands of small messages dominated by latency and launch overhead. One all-reduce at the end of backward means no overlap with computation. Buckets (PyTorch defaults to about 25 MB) are launched as soon as all their gradients are ready, so communication of later layers overlaps with backward compute of earlier layers. Bigger buckets amortize latency but delay the start of communication and leave a larger exposed tail. Smaller buckets overlap better but are latency-bound. The best size depends on network latency and bandwidth and on layer sizes, and is usually tuned empirically.
Why did DeepSeek-V3 train without tensor parallelism, and what does that teach?
DeepSeek-V3 is an MoE where most parameters live in experts, so EP (64-way across 8 nodes) plus 16-way PP and ZeRO-1 DP sharded the model enough. Their H800s have reduced NVLink bandwidth (about 160 GB/s quoted), making TP's critical-path all-reduces expensive. They invested instead in custom all-to-all kernels, node-limited routing and DualPipe to overlap communication with compute. The lesson: parallelism choices follow model architecture and interconnect, and engineering effort can substitute for a dimension. There's no universal recipe.
What is the order of mesh dimensions in Llama 3 and why?
[TP, CP, PP, DP] from innermost to outermost. TP has the most frequent, latency-sensitive, critical-path communication (all-reduces every layer), so it gets the NVLink domain. CP exchanges K/V every attention layer and sits next. PP sends activations only at stage boundaries, so it tolerates inter-node links. DP/FSDP communicates once or twice per step, can be overlapped, and is the most tolerant of slow, long-distance links. The general principle: order by communication frequency × volume ÷ ability to overlap.