Distributed PyTorch — DDP, FSDP, Tensor & Pipeline Parallelism
Past ~7B parameters, multi-GPU training failure modes change: rank-only OOM, unscaled LR, NCCL hangs from mismatched collectives, TP across a slow fabric. This lesson is the training mental model — DDP, ZeRO/FSDP, TP, PP, 3D grids — and why you should not import those intuitions into serving.
- 1Gist
- 2Maps
- 3Q&A
- 4Sandbox
Voice readout needs Web Speech Synthesis in this browser.
How you grow past one GPU
Prefer
DDP until OOM, then FSDP, then TP/PP on purpose
Replicate-and-sync is simple. Shard state when the model does not fit. Shard matrices only inside an NVLink domain. Shard depth across nodes.
- DDP memory is params + grads + optimizer + activations — Adam fp32 is ~16 bytes/param before activations.
- FSDP is PyTorch-native ZeRO-3: all-gather, compute, free; reduce-scatter in backward.
- 3D maps onto hardware: TP in-node, PP cross-node, DP/FSDP outermost.
Alternative
Turn on TP=8 across Ethernet and hope
TP is latency-bound. Across a slow fabric you add GPUs and lose throughput. NCCL hangs look like deadlocks but are mismatched collective order.
- find_unused_parameters and bucket-size mismatches are the usual hang culprits.
- Scale LR with world size or the loss diverges.
- Serving replicas are not DDP — they do not sync gradients.
Decision order
Fit first. Add latency-sensitive axes only when FSDP cannot hit the throughput target.
- 1
DDP until it OOMs
Full replica per rank. All-reduce gradient buckets. Scale LR with world size. - 2
FSDP / ZeRO
Shard opt / grads / params. Wrap per transformer block. Activation checkpoint if needed. - 3
Add TP in-node
Column then row parallel so an MLP needs one all-reduce each way. Stay on NVLink. - 4
Add PP across nodes
Micro-batches, 1F1B, watch the bubble. Rest of the cluster becomes DP.
Overview
A model that fits on one GPU is a toy. Past ~7B parameters — or a batch big enough to keep A100 / H100 tensor cores busy — failure modes change shape: OOM on rank 3 only, loss diverging because you forgot to scale LR, NCCL hangs that look like deadlocks, throughput that drops when you add GPUs because you went too aggressive on tensor parallelism across a slow interconnect.
This lesson is training-focused. The inference multi-GPU story (replicas behind a router, or tensor parallel for latency) is contrasted at the end because people routinely import training intuitions into serving and get burned.
Data parallel: DDP vs ZeRO / FSDP
DDP (DistributedDataParallel) is the baseline. Every rank holds a full copy of the model, reads a different slice of the global batch, runs forward / backward locally, then averages gradients across ranks with one all-reduce per parameter bucket.
Memory per rank = params + grads + optimizer state + activations. For Adam in fp32 that is roughly 16 bytes/param of state before activations — a 7B model is ~112 GB of state, which is why plain DDP dies well before 7B.
DDP mechanics worth remembering:
- Bucketing — gradients are grouped into ~25 MB buckets so the all-reduce overlaps with the backward pass. The last bucket cannot overlap, so it is serialized cost.
- Gradient views — DDP registers autograd hooks on parameters and reduces a flat bucket, not individual tensors.
find_unused_parameters— set it when parts of the graph are conditional; it costs a graph traversal every step. Prefer restructuring the forward to avoid it.- LR scaling — DDP with world size N means an effective batch N× larger. Use linear or sqrt scaling, plus warmup, or you will get divergence.
ZeRO / FSDP attack the memory wall by sharding replicated state instead of replicating it:
- ZeRO-1: optimizer state sharded. ZeRO-2: + gradients. ZeRO-3: + parameters.
- FSDP (FullyShardedDataParallel) is PyTorch-native ZeRO-3. Parameters live sharded; before each layer’s forward, FSDP issues an all-gather to materialize full weights, computes, then frees them; in backward it re-gathers and reduce-scatters gradients so each rank keeps only its shard.
- Communication cost goes up (~1.5×) but memory per rank drops ~linearly with world size. FSDP is the default when the model does not fit in DDP but fits in total cluster memory.
- Tune with
sharding_strategy(FULL_SHARD,SHARD_GRAD_OP,HYBRID_SHARD),auto_wrap_policy(wrap per transformer block, not per parameter), andactivation_checkpointingto trade compute for activation memory.
Rule of thumb: DDP until it OOMs, then FSDP. Reach for tensor / pipeline parallelism only when FSDP alone cannot hit your throughput target, because they add latency-sensitive communication inside the forward pass.
Tensor parallelism
Tensor parallel (TP) shards individual weight matrices across ranks and computes the layer cooperatively:
- Column-parallel linear: split output features. Each rank computes a slice of Y = XW; no communication in forward, all-reduce in backward.
- Row-parallel linear: split input features. Each rank computes a partial sum that must be all-reduced in forward.
- Megatron pairing: column-parallel then row-parallel, so a full MLP block needs exactly one all-reduce forward and one all-reduce backward. Attention heads are split naturally (head-parallel); LayerNorm and dropout are replicated.
- Embeddings are split along the vocab dimension (with a cross-entropy that avoids materializing full logits).
TP is latency-bound, not bandwidth-bound: it fires many small all-reduces per layer. That is why it stays inside an NVLink domain (typically TP ≤ 8, one node) and why it is the wrong tool across a slow Ethernet fabric. TP also shards optimizer state and gradients for free, since each rank only owns part of the weights.
Pipeline parallelism
Pipeline parallel (PP) assigns contiguous layer blocks to stages and streams micro-batches through them. Rank 0 holds layers 0–7, rank 1 holds 8–15, and so on. Communication is tiny — just activations at stage boundaries — but it is a strict point-to-point dependency, so the pipeline idles unless you fill it.
Key ideas:
- Micro-batching splits the global batch into M chunks; stage k works on micro-batch i while stage k+1 works on i−1.
- Bubble fraction ≈ (P − 1) / (M + P − 1) for a P-stage, M-micro-batch 1F1B schedule. More micro-batches = smaller bubble but higher activation memory.
- Schedules: GPipe (all forwards, then all backwards — huge memory), 1F1B (interleaved, the default), interleaved / virtual stages (each rank takes multiple non-adjacent chunks, shrinking the bubble at the cost of more p2p traffic).
torch.distributed.pipeliningprovidespipeline()andSchedule1F1B/ScheduleGPipefor arbitrarynn.Modules.
PP is the cheapest parallelism per byte moved but the hardest to keep efficient. It is the natural cross-node axis because it tolerates higher latency.
Communication collectives
Everything above is built from five primitives:
- all-reduce — sum / average across ranks, all ranks get the result. Built as reduce-scatter + all-gather (ring or tree).
- reduce-scatter — reduce, then give each rank a distinct slice. The workhorse of FSDP’s backward.
- all-gather — each rank contributes a slice; all ranks get the concatenation. FSDP’s forward.
- broadcast — one-to-all (initial parameter sync).
- all-to-all / send-recv — used by MoE routing and by pipeline stage boundaries (p2p).
NCCL picks ring vs tree by message size (tree for small, ring for large) and topology (NVLink vs PCIe vs IB). Two facts that prevent most hangs: collectives must be issued in the same order on every rank, and the tensor shapes must match. Bucket sizes and find_unused_parameters mismatches are the usual culprits.
When to combine: 3D parallelism
Frontier training runs use all three, mapped onto the hardware hierarchy:
- TP within a node (NVLink, high bandwidth, low latency).
- PP across nodes (tolerates latency, small messages).
- DP / FSDP across replicas (the outermost axis, largest all-reduces, most bandwidth-tolerant if overlapped).
Total world size = TP × PP × DP. A 1024-GPU run of a 100B model might be TP=8, PP=8, DP=16. The artifact you tune is a process grid / device mesh (in PyTorch, DeviceMesh + DTensor), and every tensor has a placement like Shard(0), Shard(1), or Replicate.
Decision order: fit the model with FSDP first → if per-layer latency dominates, add TP inside the node → if you still need more memory or you are crossing node boundaries at scale, add PP → the rest of the cluster becomes DP.
Training vs inference multi-GPU
Training and serving share collectives but not objectives:
| Axis | Training | Inference |
|---|---|---|
| Goal | Throughput (samples/sec) | Latency (TTFT, TPOT) + $/token |
| DP meaning | Replicas that sync gradients | Independent replicas behind a router (no sync) |
| TP | Pays off only with NVLink, amortized over big batches | Very common: TP=2–8 to cut per-token latency |
| PP | Common at scale | Rare; adds first-token latency (stage serialization) |
| Memory pressure | Params + grads + optimizer + activations | Weights + KV cache + runtime workspace |
| Precision | bf16 / fp8 with fp32 master weights | fp8 / int4 weights, KV cache quantization |
| Failure mode | Divergence, NCCL hang mid-step | Tail latency, KV OOM under long context |
Serving engines and KV math: hub, TensorRT-LLM, playbook.
Flow
- 1
1 DDP replicas + AllReduce
- next2 FSDP / ZeRO shards
- 2
2 FSDP / ZeRO shards
- next3 Add TP inside the node
- 3
3 Add TP inside the node
- next4 Add PP across nodes
- 4
4 Add PP across nodes
- next5 Rest of cluster is DP
- 5
5 Rest of cluster is DP
Lesson map
Distributed PyTorch — DDP, FSDP, Tensor & Pipeline Parallelism
Past ~7B parameters, multi-GPU training failure modes change: rank-only OOM, unscaled LR, NCCL hangs from mismatched collectives, TP across a slow fabric. This lesson is the training mental model — DDP, ZeRO/FSDP, TP, PP, 3D grids — and why you should not import those intuitions into serving.
Architecture. Architecture
Select a node to see why it exists, or an edge to see the protocol, direction, effect, and consequence.
Mermaid export
flowchart TB a["1 DDP replicas + AllReduce"] b["2 FSDP / ZeRO shards"] c["3 Add TP inside the node"] d["4 Add PP across nodes"] a -->|1 DDP replicas + AllReduce| b b -->|2 FSDP / ZeRO shards| c c -->|3 Add TP inside the node| d
Sandbox: memory planner (Python)
No torch, no NCCL. Bytes-per-param stand-in for the DDP vs FSDP conversation.
Press Run. Snippets must be self-contained — no network, files, or native modules.
TypeScript: 3D grid helper
World size = TP × PP × DP. Refuse a TP that does not divide heads.
Press Run. Snippets must be self-contained — no network, files, or native modules.
Deep dive · Serving multi-GPU — pointer only
Compiled TP/PP layouts: TensorRT-LLM. KV capacity and TTFT: playbook. Do not treat serving replicas as DDP ranks.
Pitfalls
A 70B model, 16×H100 in two NVLink nodes, Ethernet between nodes. Do you start with DDP, FSDP, TP=8, or PP=2 — and what hang do you expect if TP straddles the Ethernet link?
Interview Q&A
DDP vs FSDP — one-line distinction?
Answer
DDP replicates full parameters per rank and all-reduces grads. FSDP / ZeRO shards params / grads / opt state so models larger than one GPU fit, at the cost of gather / scatter communication.
When do you add tensor parallel vs pipeline parallel?
Answer
TP shards large layers (attention / MLP matmuls) across GPUs for a single step — needs a fast interconnect. PP splits sequential stages across GPUs to fit depth — needs micro-batching to limit pipeline bubbles.
Why is training multi-GPU different from inference multi-GPU?
Answer
Training optimizes step throughput with grad sync and optimizer sharding. Inference often uses TP for latency / context of one replica, or data-parallel replicas for QPS with a serving scheduler — not the same collectives cadence.
What breaks if you just enable 3D parallel without a plan?
Answer
Over-sharding → communication-bound steps; bad PP micro-batch count → bubbles; mismatched TP size vs attention heads; checkpoint / resume complexity. Start from the memory bottleneck (params vs activations) then add parallelism axes deliberately.
Go Deeper
Public PyTorch documentation only: