= Training at Scale

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.

[#sec-memory]
== Memory accounting

Mixed-precision Adam commonly stores bf16 model weights, bf16 gradients, fp32 master weights, and two fp32 moment vectors. Per parameter, that is

[latexmath#eq-adam-memory]
++++
2 + 2 + 4 + 4 + 4 = 16 \text{ bytes} .
++++

Activations are separate: a rough saved-activation budget is stem:[B T d L] 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.

.Memory and checkpointing calculators
[source,python]
----
include::../../scratch/distributed_training.py[tag=memory]
----

Checkpointing does not change the mathematical gradient. It changes the execution schedule: instead of keeping every intermediate stem:[h_\ell], 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.

[#sec-data-zero]
== 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 stem:[(n-1)/n] of the tensor per rank, so the total traffic per rank is

[latexmath#eq-ring]
++++
2\frac{n-1}{n}\;\text{size} .
++++

ZeRO reduces memory by sharding model states across data-parallel ranks <<rajbhandari2019zero>>. With the 16-byte Adam accounting above and stem:[n] ranks, the idealized per-rank model-state memory is

[latexmath#eq-zero]
++++
\begin{aligned}
\text{DP} &= 16P, \\
\text{ZeRO-1} &= 4P + 12P/n, \\
\text{ZeRO-2} &= 2P + 14P/n, \\
\text{ZeRO-3} &= 16P/n .
\end{aligned}
++++

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.

.Ring all-reduce and ZeRO memory
[source,python]
----
include::../../scratch/distributed_training.py[tag=collectives-zero]
----

[#sec-model-parallel]
== Tensor parallel MLPs

Tensor parallelism splits one layer across ranks. For a transformer MLP stem:[\operatorname{GELU}(\mX\mW_1+\vb_1)\mW_2+\vb_2], split stem:[\mW_1] by columns and stem:[\mW_2] by matching rows <<shoeybi2019megatronlm>>. Rank stem:[r] computes
stem:[\mH_r = \operatorname{GELU}(\mX\mW_{1,r}+\vb_{1,r})] and stem:[\mO_r=\mH_r\mW_{2,r}]. Concatenating the hidden shards would reconstruct stem:[\mH]; summing the output shards gives

[latexmath#eq-tp-mlp]
++++
\sum_r \mH_r\mW_{2,r} = \mH\mW_2 .
++++

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.

.Column- then row-parallel MLP
[source,python]
----
include::../../scratch/distributed_training.py[tag=tensor-parallel]
----

[#sec-pipeline-expert]
== Pipeline and expert parallelism

Pipeline parallelism assigns consecutive layers to stages and sends activations between them. With stem:[p] stages and stem:[m] microbatches, the pipeline spends stem:[p-1] slots filling and draining. The bubble fraction is

[latexmath#eq-bubble]
++++
\frac{p-1}{m+p-1} .
++++

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.

.Pipeline bubbles and expert all-to-all counts
[source,python]
----
include::../../scratch/distributed_training.py[tag=pipeline-expert]
----

[#sec-fp8]
== 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.

.Fine-grained FP8 block scaling
[source,python]
----
include::../../scratch/distributed_training.py[tag=fp8]
----

[NOTE,caption=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>>.
====

[.key-equations#key-equations]
.Key equations
****
[latexmath]
++++
\text{Adam bytes/param}=2+2+4+4+4=16
++++
[latexmath]
++++
\text{ring bytes}=2\frac{n-1}{n}\,\text{size}
++++
[latexmath]
++++
\text{ZeRO-2}=2P+14P/n
++++
[latexmath]
++++
\sum_r \operatorname{GELU}(\mX\mW_{1,r}+\vb_{1,r})\mW_{2,r}=\mH\mW_2
++++
[latexmath]
++++
\text{bubble}=\frac{p-1}{m+p-1}
++++
****

[.teach]
[#sec-teach]
== 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?

[#sec-exercises]
== Exercises

[#ex-distributed-training-memory.exercise]
.★ Memory budget
====
Compute model-state memory for one million parameters with mixed-precision Adam. Then compute the saved-activation bytes for stem:[8] tokens, hidden size stem:[16], and stem:[12] layers at two bytes per value.
====

[#ex-distributed-training-zero.exercise]
.★★ Ring and ZeRO
====
Derive the ring all-reduce traffic stem:[2(n-1)\text{size}/n]. For stem:[P=1000] and stem:[n=4], compute the DP and ZeRO stage memory values from <<eq-zero>>.
====

[#ex-distributed-training-tensor-parallel.exercise]
.★★ Tensor-parallel equality
====
Show algebraically why column-splitting stem:[\mW_1] and row-splitting stem:[\mW_2] preserves an MLP. Then verify with the NumPy function.
====

[#ex-distributed-training-pipeline-expert-fp8.exercise]
.★★★ Pipeline, experts, and FP8
====
For stem:[p=4] pipeline stages and stem:[m=12] 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.
====

[bibliography]
[#sec-references]
== References

include::../../book/sources.adoc[tags=rajbhandari2019zero;shoeybi2019megatronlm;huang2018gpipe;lepikhin2020gshard;micikevicius2022fp8;deepseekai2024deepseekv3]
