Chapter 41
Training at Scale
Memory accounting, data, tensor, pipeline, and expert parallelism, ZeRO, and FP8.
Large-model training is mostly bookkeeping: which bytes live on which accelerator, which tensors must be communicated, and which activations can be recomputed instead of stored. This chapter keeps everything single-process and NumPy-sized, but the formulas are the same ones used to plan a real training run. The purpose is to make parallel training feel like algebra over memory and communication, not magic distributed systems.
41.1 Memory accounting
Mixed-precision Adam commonly stores bf16 model weights, bf16 gradients, fp32 master weights, and two fp32 moment vectors. Per parameter, that is
Activations are separate: a rough saved-activation budget is values times the bytes per value and the number of tensors saved per layer. Unlike optimizer state, activation memory grows with batch size and sequence length. Activation checkpointing trades compute for memory by storing only selected layer boundaries during the forward pass and recomputing the missing interiors during backward.
def adam_memory_bytes(parameters, weight=2, grad=2, master=4, moments=8):
"""Mixed-precision Adam bytes: bf16 weights/grads plus fp32 state."""
return parameters * (weight + grad + master + moments)
def activation_memory_bytes(tokens, hidden, layers, bytes_per_value=2,
tensors_per_layer=1):
return tokens * hidden * layers * bytes_per_value * tensors_per_layer
def checkpointed_activation_bytes(tokens, hidden, layers, segments,
bytes_per_value=2):
"""Store segment boundaries and recompute interiors during backward."""
segment_length = int(np.ceil(layers / segments))
saved_layers = segments + segment_length
return tokens * hidden * saved_layers * bytes_per_value
Checkpointing does not change the mathematical gradient. It changes the execution schedule: instead of keeping every intermediate , backward replays a segment from its saved input until it reaches the needed intermediate. The cost is extra forward compute; the benefit is that long contexts or larger microbatches fit without changing the model.
This separation matters when comparing scaling techniques. Optimizer state is tied to parameters and is present even for batch size one. Activation memory is tied to the current batch and sequence length, so it can dominate during long-context training. Checkpointing attacks the second term only; ZeRO attacks model state; tensor, pipeline, and expert parallelism attack different forms of compute and communication. A training plan usually starts by writing these terms down before choosing any parallel layout.
41.2 Data parallelism and ZeRO
In data parallelism, each rank owns a full model replica and a different minibatch shard. Backward produces one gradient tensor per rank; all-reduce averages them so every replica applies the same update. A ring all-reduce sends and receives chunks around a ring. Reduce-scatter plus all-gather each move of the tensor per rank, so the total traffic per rank is
ZeRO reduces memory by sharding model states across data-parallel ranks [rajbhandari2019zero]. With the 16-byte Adam accounting above and ranks, the idealized per-rank model-state memory is
ZeRO-1 shards optimizer state; ZeRO-2 also shards gradients; ZeRO-3 shards parameters too, gathering them when a layer needs them. Real systems budget temporary all-gather buffers and overlap communication, but these formulas are the planning baseline.
The formulas also show why communication and memory are coupled. Ordinary data parallelism communicates full gradients but keeps full optimizer state. ZeRO-2 reduces gradient memory, yet the averaged gradient still has to be assembled logically for the update. ZeRO-3 saves the most persistent memory, but every layer now needs parameter shards to arrive before its matmul. That is why implementations try to prefetch the next layer’s parameters while the current layer computes.
def ring_all_reduce_bytes(size_bytes, ranks):
"""Bytes sent per rank by ring all-reduce."""
return 2 * (ranks - 1) / ranks * size_bytes
def zero_memory_bytes(parameters, ranks):
"""Per-rank model-state bytes for data parallel Adam and ZeRO stages."""
p = parameters
return {
"dp": 16 * p,
"zero1": 4 * p + 12 * p / ranks,
"zero2": 2 * p + 14 * p / ranks,
"zero3": 16 * p / ranks,
}
41.3 Tensor parallel MLPs
Tensor parallelism splits one layer across ranks. For a transformer MLP , split by columns and by matching rows [shoeybi2019megatronlm]. Rank computes and . Concatenating the hidden shards would reconstruct ; summing the output shards gives
Thus the sharded layer equals the unsharded layer, except that the partial outputs must be all-reduced or reduce-scattered. The test verifies equality to float64 precision.
The equality depends on splitting along the hidden dimension between the two linear maps. The activation is elementwise, so each rank can apply it to its hidden shard without seeing the others. The second projection mixes hidden features into output features; row-splitting its input dimension makes each rank produce a partial output with the full output width. Summing those partial outputs is exactly the missing matrix multiplication over the concatenated hidden dimension.
def gelu(x):
return 0.5 * x * (1.0 + np.tanh(np.sqrt(2 / np.pi) *
(x + 0.044715 * x ** 3)))
def mlp(x, w1, b1, w2, b2):
return gelu(x @ w1 + b1) @ w2 + b2
def tensor_parallel_mlp(x, w1, b1, w2, b2, ranks):
"""Column-parallel first layer, row-parallel second layer."""
w1_parts = np.array_split(w1, ranks, axis=1)
b1_parts = np.array_split(b1, ranks)
w2_parts = np.array_split(w2, ranks, axis=0)
partials = []
for w1_i, b1_i, w2_i in zip(w1_parts, b1_parts, w2_parts):
partials.append(gelu(x @ w1_i + b1_i) @ w2_i)
return np.sum(partials, axis=0) + b2
41.4 Pipeline and expert parallelism
Pipeline parallelism assigns consecutive layers to stages and sends activations between them. With stages and microbatches, the pipeline spends slots filling and draining. The bubble fraction is
More microbatches shrink the bubble, but too many microbatches can make kernels small and communication frequent. GPipe popularized this microbatch view for giant models [huang2018gpipe].
Expert parallelism shards experts instead of dense layers. A router chooses experts for each token; tokens assigned to remote experts move through an all-to-all, experts run locally, and outputs move back. This makes communication depend on the routing histogram rather than only on tensor shape. GShard-style MoE systems made this all-to-all routing a core training primitive [lepikhin2020gshard].
This is different from tensor parallelism: tensor parallel collectives are scheduled by layer shape, while expert traffic depends on the batch’s routing decisions. A balanced router sends roughly equal token counts to experts; an imbalanced router can overload one expert-owning rank while others wait. Real MoE training therefore adds routing or load-balancing rules, but the smallest useful simulator is just a count matrix from source ranks to destination ranks.
def pipeline_bubble_fraction(stages, microbatches):
return (stages - 1) / (microbatches + stages - 1)
def expert_all_to_all_counts(assignments, expert_to_rank, ranks):
"""Count tokens sent from each source rank to each expert-owning rank."""
assignments = np.asarray(assignments)
counts = np.zeros((ranks, ranks), dtype=np.int64)
for source in range(ranks):
for expert in assignments[source]:
counts[source, expert_to_rank[int(expert)]] += 1
return counts
41.5 FP8 training with fine-grained scaling
FP8 training stores or communicates selected tensors in FP8 while keeping enough higher-precision accumulation and scaling metadata to train stably [micikevicius2022fp8]. A single scale for a whole tensor is fragile: one large block forces small blocks to use a coarse step. Fine-grained scaling instead stores a scale per block, quantizes values relative to that local scale, and dequantizes before accumulation.
DeepSeek-V3 reports FP8 training with fine-grained scaling and higher-precision accumulation for stability [deepseekai2024deepseekv3]. The NumPy version below block-scales values into E4M3 range, applies the chapter’s FP8 emulator, and rescales back. It is a calculator for the error tradeoff, not a training kernel.
def fine_grained_fp8(x, block_size=16, max_value=448.0):
"""Block-scale values into E4M3 range, quantize, then dequantize."""
x = np.asarray(x, dtype=np.float32)
if x.shape[-1] % block_size:
raise ValueError("last dimension must be divisible by block_size")
blocks = x.reshape(*x.shape[:-1], x.shape[-1] // block_size, block_size)
scale = np.max(np.abs(blocks), axis=-1, keepdims=True) / max_value
scale = np.where(scale == 0, 1.0, scale).astype(np.float32)
dequant = fp8_e4m3(blocks / scale) * scale
return dequant.reshape(x.shape), scale
|
In practice
|
Large training runs combine these axes rather than choosing one. Data parallelism scales batches; ZeRO shards states inside data parallelism; tensor parallelism splits individual matrix multiplies; pipeline parallelism splits depth; expert parallelism splits sparse experts. The best layout is hardware- and model-dependent, because every saved byte can introduce a collective. FP8 training adds another axis: it reduces bandwidth and storage, but only with scaling and accumulation rules that keep optimization stable [micikevicius2022fp8] [deepseekai2024deepseekv3]. |
41.6 Teach it
One sentence: training at scale is deciding which model states, activations, and tokens are replicated, sharded, communicated, or recomputed. Analogy: a kitchen can duplicate the whole recipe at every station, split ingredients among stations, or pass dishes down an assembly line; each saves a different bottleneck. Board steps: 1. write 16 bytes per Adam parameter; 2. draw ring all-reduce traffic; 3. split an MLP’s hidden dimension; 4. draw pipeline bubbles and expert all-to-all. Misconceptions: ZeRO is data parallelism with sharded state, not tensor parallelism; checkpointing saves memory but costs compute; FP8 training is scaling policy plus accumulation, not just casting. Check: which parallelism axis creates an all-to-all over tokens?
41.7 Exercises
Compute model-state memory for one million parameters with mixed-precision Adam. Then compute the saved-activation bytes for tokens, hidden size , and layers at two bytes per value.
Derive the ring all-reduce traffic . For and , compute the DP and ZeRO stage memory values from (41.3).
Show algebraically why column-splitting and row-splitting preserves an MLP. Then verify with the NumPy function.
For pipeline stages and microbatches, compute the bubble fraction. Given two ranks with expert assignments [[0,1,3],[2,3,2]] and experts 0,1 on rank 0 and 2,3 on rank 1, compute the all-to-all count matrix. Finally, explain why block FP8 scaling helps the small block in the test.
References
-
[deepseekai2024deepseekv3] DeepSeek-AI et al. DeepSeek-V3 Technical Report. 2024. arXiv:2412.19437
-
[huang2018gpipe] Y. Huang et al. GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism. 2018. arXiv:1811.06965
-
[lepikhin2020gshard] D. Lepikhin et al. GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding. 2020. arXiv:2006.16668
-
[micikevicius2022fp8] P. Micikevicius et al. FP8 Formats for Deep Learning. 2022. arXiv:2209.05433
-
[rajbhandari2019zero] S. Rajbhandari et al. ZeRO: Memory Optimizations Toward Training Trillion Parameter Models. 2019. arXiv:1910.02054
-
[shoeybi2019megatronlm] M. Shoeybi et al. Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism. 2019. arXiv:1909.08053