Track A · Model internals

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:

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.

StrategySplitsSaves memory onCollective addedTypical placement
Data parallel (DDP)BatchNothing (activations per GPU shrink with local batch)All-reduce of grads, once per stepOutermost, across nodes
ZeRO-1/2/3, FSDPBatch, plus sharded stateOptimizer state / grads / paramsReduce-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 matmulParams, grads, optimizer state, most activationsAll-reduce (or RS+AG with SP) ×4 per layerInside the NVLink domain
Sequence parallel (SP)Sequence, in non-matmul regionsLayerNorm/dropout activationsReplaces TP all-reduce with RS+AGAlways paired with TP
Pipeline parallel (PP)LayersParams and state per stagePoint-to-point activation sendsAcross nodes (tolerates slow links)
Context parallel (CP)Sequence, including attentionActivations at long contextRing P2P of K/V, or all-gather of K/VIn node or across nearby nodes
Expert parallel (EP)MoE expertsExpert paramsAll-to-all dispatch + combineIn node or across a few nodes
Interview angle

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.

Go deeper

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:

ItemPrecisionBytes/paramWhy it exists
Working weightsBF162Used in forward/backward matmuls
GradientsBF16 (or FP32)2 (or 4)Output of backward, input to optimizer
Master weightsFP324Accumulates tiny updates without rounding loss
Adam first moment mFP324Momentum (EMA of grad)
Adam second moment vFP324EMA of grad², sets per-param step size
Total16 (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.

Common mistake

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.

Intuition

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).

Go deeper

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.

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.

Interview angle

"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.

Go deeper

Mixed precision: BF16, FP16, FP8, MX formats

The formats

FormatSign/Exp/Mantissa bitsDynamic rangeUse in training
FP321/8/23~1e-38 to 3e38Master weights, optimizer state, reductions
TF321/8/10 (internal)Same as FP32Tensor-core mode for "FP32" matmuls on Ampere+
FP161/5/10~6e-8 (subnormal) to 65504Older standard; needs loss scaling
BF161/8/7Same exponent range as FP32Today's default compute format
FP8 E4M31/4/3max 448Weights and activations (forward)
FP8 E5M21/5/2max 57344Gradients (wider range, less precision)
MXFP8 / MXFP6 / MXFP4FP8/6/4 elements + shared E8M0 scale per 32 elementsPer-block power-of-two scalingBlackwell-native block-scaled matmuls
NVFP4E2M1 elements + E4M3 scale per 16 elements + per-tensor FP32 scaleFiner blocks, non-power-of-two scaleInference, 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:

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.

PracticeMaturity (as of late 2026)
BF16 compute + FP32 master/optimizerWell established; the safe default everywhere
FP16 + dynamic loss scalingLegacy; mostly on V100-era hardware
FP8 GEMMs (per-tensor or block scaling) on H100/H200Established in production at frontier labs (DeepSeek-V3 is the best public evidence); still needs care
MXFP8 on BlackwellRecent; supported in Transformer Engine/Megatron, with NVIDIA-published recipes
NVFP4 / MXFP4 pretrainingEmerging and contested; demonstrated at 12B/10T scale, not yet a default
May be out of date

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.

Interview angle

"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.

Go deeper

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.

LevelTechnologyRough per-GPU bandwidthScope
On-chipHBM3/3e3.35 TB/s (H100), ~8 TB/s (B200)One GPU
Scale-upNVLink 4 + NVSwitch (Hopper)900 GB/s total (≈450 GB/s each direction)8 GPUs (HGX node)
Scale-upNVLink 5 (Blackwell), NVL72 rack1.8 TB/s totalUp to 72 GPUs in one NVLink domain
Scale-upNVLink 6 (Vera Rubin)~3 TB/s+ (preliminary)Rack-scale
Host linkPCIe Gen5 x16~64 GB/s each directionGPU ↔ CPU/NIC
Scale-outInfiniBand NDR / RoCE 400G400 Gb/s ≈ 50 GB/s per NIC, usually 1 NIC per GPUAcross nodes, thousands of GPUs
Scale-outInfiniBand XDR / 800G Ethernet800 Gb/s ≈ 100 GB/s per NICNewest 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.

Intuition

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.

May be out of date

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

CollectiveInput → output (N ranks, S bytes total)Where it shows up
BroadcastOne rank's S → everyoneInitial weight sync
ReduceEveryone's S → sum on one rankRare in training
All-reduceEveryone's S → everyone has the sumDDP grads, TP activations
Reduce-scatterEveryone's S → each rank holds 1/N of the sumZeRO/FSDP grads, SP
All-gatherEach rank's S/N shard → everyone has all SFSDP params, SP, CP (K/V)
All-to-allEach rank sends a distinct chunk to every other rankMoE dispatch/combine, Ulysses SP
Send/recv (P2P)Point-to-pointPipeline 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.

Ring all-reduce, N=4, buffer split into chunks a,b,c,d reduce-scatter (N-1 = 3 steps) all-gather (3 steps) GPU0 ─► GPU1 ─► GPU2 ─► GPU3 ─┐ each GPU now owns one fully ▲ │ reduced chunk; pass the finished └────────────────────────────┘ chunks around the ring 3 more times after RS: GPU0 has Σd, GPU1 has Σa, GPU2 has Σb, GPU3 has Σc after AG: every GPU has Σa Σb Σc Σd bytes sent per GPU = 2·(N-1)/N · S = 1.5·S for N=4

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.

Common mistake

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.

Interview angle

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.

Go deeper

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:

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.

Interview angle

"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\):

StageShardsMemory per GPUComm per step (per GPU, in units of Ψ elements)Notes
DDPNothing\((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
Per-GPU memory, 4 data-parallel ranks params (2Ψ) grads (2Ψ) optimizer (12Ψ) DDP 16Ψ on every GPU ZeRO-1 4Ψ + 12Ψ/4 = 7Ψ ZeRO-2 2Ψ + 14Ψ/4 = 5.5Ψ ZeRO-3 16Ψ/4 = 4Ψ (+ one gathered layer at a time) Dashed outline = memory freed relative to DDP. Comm: DDP, Z1, Z2 move 2Ψ per step; Z3 moves 3Ψ.
ZeRO stages for Nd=4. Each stage shards one more category of model state across the data-parallel group.

How ZeRO-3 / FSDP runs a step

for each layer (or "FSDP unit") in forward: all-gather full params of this unit from all DP ranks # Ψ_unit traffic compute forward free the gathered params (reshard_after_forward=True) # keep only own 1/N shard for each unit in reverse during backward: all-gather params again # Ψ_unit traffic compute grads reduce-scatter grads -> each rank keeps 1/N of summed grad # Ψ_unit traffic free gathered params and full grads optimizer step on the local 1/N shard only (no communication)

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

Common mistake

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.

Interview angle

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.

Go deeper

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.

X replicated (b·s × h) f: identity fwd all-reduce bwd GeLU(X·A₁)A₁: h × 4h/2 (cols) GeLU(X·A₂)A₂: h × 4h/2 (cols) Y₁·B₁B₁: 4h/2 × h (rows) Y₂·B₂B₂: 4h/2 × h (rows) g: all-reduce(sum partials) fwd Z GPU 0 (top row) and GPU 1 (bottom row). No communication between the two matmuls; one all-reduce of b·s·h at the end.
Megatron tensor-parallel MLP with TP=2. Column-parallel first GEMM, row-parallel second GEMM. The conjugate operators f and g are an identity and an all-reduce, swapped between forward and backward.

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.

Interview angle

"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.

[LayerNorm, dropout: seq-sharded s/t] │ all-gather (along seq) g ▼ [Attention or MLP: TP, hidden-sharded] │ reduce-scatter (along seq) ḡ ▼ [dropout + residual: seq-sharded s/t]

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.

Common mistake

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.

Go deeper

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.

GPipe (p=4, m=4): all forwards, then all backwards stage 0 stage 1 stage 2 stage 3 F1 B1 F2 B2 F3 B3 F4 B4 F1 B1 F2 B2 F3 B3 F4 B4 F1 B1 F2 B2 F3 B3 F4 B4 F1 B1 F2 B2 F3 B3 F4 B4 1F1B (p=4, m=8): warm-up, steady 1-forward-1-backward, cool-down stage 0 stage 1 stage 2 stage 3 F1 F2 F3 F4 F1 F2 F3 F1 F2 F1 B1 F2 B2 B1 F3 B2 F3 B3 B1 F4 B2 F4 B3 F4 B4 B1 F5 B2 F6 F5 B3 F6 F5 B4 F6 F5 B5 F6 B6 B3 F7 B4 F7 B5 F7 B6 F7 B7 B4 F8 B5 F8 B6 F8 B7 F8 B8 B5 B6 B7 B8 B7 B8 B8 Empty dashed slots = bubble. Same bubble fraction as GPipe for equal m, but each stage holds at most p in-flight micro-batches instead of m.
Pipeline schedules with equal forward/backward cost per slot (a simplification: backward is really about 2× forward). Both have \((p-1)\) slots of bubble per stage at each end. 1F1B's win is memory, not bubble.

Schedules

ScheduleIdeaBubbleActivation memory per stage
GPipeAll forwards, then all backwards; synchronous flush\((p-1)/m\)\(m\) micro-batches in flight (bad)
1F1B (PipeDream-Flush) Harlap+ 2018Warm 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+ 2021Each 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+ 2023Split backward into B (grad w.r.t. input, on critical path) and W (grad w.r.t. weights, deferrable) and use W to fill bubblesRoughly 1/3 of 1F1B (H1) down to near zero (H2)H1 ≈ 1F1B; H2 needs more
DualPipe (DeepSeek-V3) DeepSeek-AI 2024Bidirectional: feed micro-batches from both ends; overlap MoE all-to-all and PP comm inside forward/backward chunk pairsReduced vs 1F1B / ZBKeeps 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).

Interview angle

"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.

Go deeper

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/VUlysses (all-to-all)
Comm patternc−1 neighbor sends of K/V, overlappedOne all-gather of K/V per layer2 all-to-alls per attention
Scales past #heads?YesYesNo (≤ heads)
Peak memoryLow (one remote chunk at a time)Full K/V gatheredFull sequence for a/c heads
Best whenVery long sequences, slower linksGQA (small K/V), complex masksFast all-to-all, many heads
Interview angle

"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.

Go deeper

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:

tokens (sharded by DP/EP rank) │ router: top-k experts per token ▼ all-to-all DISPATCH : send each token's hidden vector to the GPU(s) owning its experts ▼ local expert MLPs (grouped GEMM over the tokens received) ▼ all-to-all COMBINE : send outputs back to the token's home GPU, weighted sum ▼ backward: the same two all-to-alls in reverse → 4 all-to-alls per MoE layer per step

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:

Intuition

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.

Go deeper

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.

Node 0dp=0, pp stage=0 (one TP=8 group, NVLink) tp0 tp1 tp2 tp3 tp4 tp5 tp6 tp7 Node 1dp=0, pp stage=1 (one TP=8 group, NVLink) tp0 tp1 tp2 tp3 tp4 tp5 tp6 tp7 Node 2dp=1, pp stage=0 (one TP=8 group, NVLink) tp0 tp1 tp2 tp3 tp4 tp5 tp6 tp7 Node 3dp=1, pp stage=1 (one TP=8 group, NVLink) tp0 tp1 tp2 tp3 tp4 tp5 tp6 tp7 PP: send/recv activations DP/FSDP: reduce-scatter / all-gather between same-index GPUs (same rail) Mesh shape [dp=2, pp=2, tp=8] = 32 GPUs. Innermost dim (TP) = fastest links; outermost (DP) = slowest, least frequent traffic.
A 32-GPU device mesh with shape [dp=2, pp=2, tp=8]. TP lives inside each NVLink node, PP crosses nodes with P2P, and DP crosses nodes between GPUs with the same local index.

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:

GPUsTPCPPPDPSeq lenTokens/batchTFLOP/s per GPUBF16 MFU
8,1928116648,19216M43043%
16,38481161288,19216M40041%
16,384816168131,07216M38038%

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:

  1. 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.
  2. Shard state first. Use FSDP/ZeRO-3, or HSDP across many nodes. It's the least invasive option, with no model code changes.
  3. 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).
  4. 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.
  5. For long sequences, add CP. For MoE, add EP.
  6. Then tune for throughput: micro-batch size, recompute policy, bucket sizes, overlap. Measure MFU and exposed communication with a profiler.
SituationRecommended starting pointWhy
≤ ~1B params, any GPU countDDP (+ ZeRO-1)Fits easily; lowest comm; simplest
1–13B, up to a few hundred GPUsFSDP2 (ZeRO-3) or ZeRO-2, + activation checkpointing as neededState sharded; no model surgery; comm overlaps well
30–70B dense, hundreds of GPUsFSDP/HSDP + TP=8 inside node (+SP); or TP=8 + PP=2–4 + ZeRO-1A single layer is large; TP cuts activations and state; FSDP across nodes
100B–1T dense, thousands of GPUsTP=8 (in node) × PP (8–16, across nodes) × DP/FSDP; interleaved or ZB scheduleToo large for FSDP alone; PP tolerates slow links
Long context (≥64K)Add CP (ring / all-gather K/V), zig-zag balancing; FlashAttention; selective recomputeActivations ∝ s; only CP divides attention-region activations
MoEEP (in node or few nodes) + DP/FSDP for dense parts; PP across; TP optionalExperts 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 compressionMinimize per-layer cross-node traffic
TPU podsJAX/XLA GSPMD sharding over a 2–3D mesh (data, fsdp, model)ICI torus favors sharding annotations over hand-written PP
Common mistake

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.

Interview angle

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.

Go deeper

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.

SettingReported utilization
PaLM 540B on 6144 TPU v446.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, BF1635–55%
MoE / long context / small modelsUsually 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.

Interview angle

"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

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:

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.

May be out of date

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.

Interview angle

"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.

Go deeper

Frameworks

FrameworkWhat it isStrengthsNotes
PyTorch native (DDP, FSDP2, DTensor, TP, pipelining, DCP)Composable building blocks on DeviceMesh + DTensorStandard, compiles with torch.compile, flexibleYou assemble the pieces yourself
torchtitanPyTorch's reference LLM pretraining stack GitHubFSDP2 + TP/SP + PP + CP (+ EP), Float8, async checkpointing, clean codeGood for learning how the native pieces compose
Megatron-LM / Megatron-CoreNVIDIA's high-performance training library GitHubBest-in-class TP/SP/PP/CP/EP on NVIDIA, FP8/MXFP8 via Transformer EngineMost frontier-scale open recipes build on it
DeepSpeedMicrosoft's library GitHubZeRO 1/2/3, offload/Infinity, ZeRO++, UlyssesEasy to bolt onto existing models; common for fine-tuning
Transformer EngineNVIDIA kernels + FP8/FP4 recipes GitHubFP8 delayed/current/block scaling, MXFP8, NVFP4Used by Megatron, torchtitan, others
JAX / XLA (GSPMD, shard_map)Annotate array shardings over a Mesh; the compiler inserts collectives JAX docsVery concise; same code from 1 to thousands of chips; TPU-nativeUsed 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.

May be out of date

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.