🧩 Distributed Training & Parallelism
Work out what lives on each GPU, which tensors cross the network, and why a model can fit yet still train slowly.
On this page
Before you start
You need the training loop from Deep Learning Basics: forward pass, loss, backward pass and optimizer update. Linear Algebra helps with tensor shapes; Numerical Computing explains precision and memory units.
By the end, you should be able to draw where weights, gradients and optimizer state live, estimate persistent memory per device, and explain the communication your plan requires.
Start with the problem: the training state does not fit
Suppose a model has N = 7 × 10⁹ parameters. Consider one specific mixed-precision Adam layout:
| Stored item | Bytes per parameter | Purpose |
|---|---|---|
| BF16 or FP16 weights | 2 | Forward/backward computation |
| BF16 or FP16 gradients | 2 | Accumulated update signal |
| FP32 master weights | 4 | Higher-precision optimizer updates |
| FP32 first and second moments | 8 | Adam's running gradient statistics |
| Total | 16 | Persistent state in this example |
Adam's second moment is an average of squared gradients, not a centered variance. Some training stacks keep FP32 gradients or omit a separate master copy, so inspect the actual implementation before using 16 bytes/parameter.
The calculation is 7 × 10⁹ × 16 = 112 × 10⁹ bytes: 112 GB, or about 104.31 GiB. Here, 1 GB = 10⁹ bytes and 1 GiB = 2³⁰ bytes. That already exceeds a 64 GB device, before saved activations, temporary buffers and allocator overhead.
There are two different jobs: partition state so it fits, and process more examples per second. Adding replicas only addresses the second job.
Follow one update across two replicas
In data parallelism, both GPUs start with the same weights but receive different examples. Suppose equally sized local batches produce mean gradients [2, 4] and [6, 0]. Sum and divide by two to get [4, 2]. With learning rate 0.1 and starting weights [1, 1], an SGD update gives [0.6, 0.8] on both replicas.
An all-reduce combines values and makes the result available to every participating rank. A rank is one process in the distributed job. Unequal local batch sizes need example-count weighting; simply averaging rank means can optimize the wrong objective. A sum collective also needs the intended normalization applied somewhere.
Plain data parallelism keeps the full model state on each GPU. More devices can increase throughput, but synchronization, input loading, stragglers and the chosen global batch size limit scaling.
Three ways to split the work
| Strategy | What is split? | What crosses devices? | Main limitation |
|---|---|---|---|
| Data parallelism, DP | Independent examples | Gradients or their shards | Replicated state unless it is sharded |
| Tensor parallelism, TP | Tensors within a layer | Partial activations/results via collectives | Frequent communication |
| Pipeline parallelism, PP | Consecutive groups of layers | Boundary activations and backward gradients | Imbalanced stages and idle bubbles |
For a linear layer Y = XW, splitting columns of W lets ranks compute different output features. How those features are combined depends on the next operation and sharding layout; not every layer uses the same collective.
TP often benefits from the fastest links within a node, but topology and measured communication determine placement. PP exchanges fewer kinds of tensors at stage boundaries, yet a large activation or slow network can still make that exchange expensive.
Use the walkthrough to follow which state moves. Its memory and activation estimates are a teaching model; a bar below the device budget is not a measured peak-memory guarantee.
Shard duplicated state: ZeRO and FSDP
Let D be the number of data-parallel ranks. In the 16-byte layout above, the optimizer-related bucket includes the master weights and two moments: 12 bytes/parameter.
| Layout | Persistent bytes per rank | At N = 7B and D = 8 |
|---|---|---|
| Replicated DP | 16N |
112 GB |
| ZeRO-1: shard optimizer state | (4 + 12/D)N |
38.5 GB |
| ZeRO-2: also shard gradients | (2 + 14/D)N |
26.25 GB |
| ZeRO-3: also shard parameters | 16N/D |
14 GB |
ZeRO-3 gathers needed parameters for computation and releases or reuses them according to its schedule. FSDP full sharding follows a related idea; its grouping, prefetch and resharding behavior depend on configuration and version. A gathered layer still occupies memory temporarily. Persistent shards are a lower-level accounting step, not the whole peak.
An all-gather assembles shards for each rank; a reduce-scatter combines gradients and leaves each rank its assigned shard. Overlapping collectives with computation helps only when there is enough independent work and memory headroom.
Keep the pipeline busy
For a simple balanced pipeline schedule with P stages and m microbatches, an idealized bubble fraction is (P − 1)/(m + P − 1). This assumes comparable stage timing and no extra network stalls.
With P = 4, eight microbatches give 3/11 ≈ 27.3% idle fraction. Thirty-two give 3/35 ≈ 8.6%. More microbatches can reduce bubbles, but smaller matrix operations may run less efficiently. Different schedules, imbalance and activation lifetimes change the result.
Keep batch accounting separate: effective batch = examples per microbatch × accumulation steps × DP degree. TP and PP cooperate on the same examples. For 2 × 16 × 4, the effective batch is 128 examples, regardless of how each replica is split internally.
Build a mesh, then measure it
When axes are independent, total devices are TP × PP × DP. A 32-device candidate with TP=4, PP=2, DP=4 gives eight devices per model replica. A 70B model has 70/8 = 8.75B parameters per TP/PP shard under ideal balance. ZeRO-3 makes its persistent 16-byte state 8.75 × 16 / 4 = 35 GB per device. Activations, gathered parameters and communication workspaces still need room.
Long sequences may need sequence/context parallelism to distribute activation or attention work. MoE models may place experts on different devices and route tokens with all-to-all communication. These axes can reuse device groups: define the mesh explicitly before multiplying degrees.
PyTorch provides DDP, FSDP and tensor-sharding APIs; JAX represents sharded arrays over device meshes and supports compiler-assisted and explicit distributed programs. Both require correct partitioning and communication reasoning. Compare a pinned workload, not an assumed framework speed hierarchy.
Try the device-mesh playground, then ask what its model omits: sequence length, activation checkpointing, layer balance, collective time and peak allocations. In a real run, measure step time, throughput, peak memory and time spent waiting at collectives. Verify loss and gradient scaling before benchmarking speed.
Check yourself
An eight-rank job uses the example 7B model with ZeRO-2. A colleague says its persistent state is 112/8 = 14 GB per GPU. Is that right?
Solution: No. ZeRO-2 leaves weights replicated. Weights take 14 GB; gradients plus optimizer state take 7B × 14/8 = 12.25 GB. The total is 26.25 GB, before activations and temporary storage. The 14 GB result belongs to idealized ZeRO-3 persistent state.
Where to go next
Read Quantization to separate precision savings from sharding, then Model Serving to see why inference has a different memory budget.
References
- ZeRO paper: the state-sharding stages and memory accounting.
- PyTorch FSDP tutorial: current gather/reshard behavior and implementation guidance.
- Megatron-LM paper: tensor parallelism within transformer layers.