Building a Distributed Training Framework
These are my own study notes, written while working through Umar Jamil's marathon walkthrough Building a distributed training framework from first principles (about 19.5 hours, 30 chapters). The notes are original write-ups of the underlying ideas, with my own derivations, figures and worked numbers; they are not a transcript and they contain no material copied from the video. Watch the original for the live coding, which is the part notes cannot replace: youtube.com/watch?v=XoGvCBRnwLs.
The premise
A single accelerator gives you a fixed amount of memory and a fixed number of FLOP/s. A modern training run needs more of both than any one device has. Every parallelism strategy is an answer to the question which axis of the computation do we split, and what does that force us to communicate?
| Strategy | Split along | Replicated | Communication per step |
|---|---|---|---|
| Data parallel (DDP) | batch | all parameters | all-reduce of the gradient |
| ZeRO / FSDP | batch, plus optimizer and parameter sharding | nothing permanently | all-gather of parameters, reduce-scatter of gradients |
| Pipeline parallel | layers | nothing | point-to-point activations at stage boundaries |
| Tensor parallel | inside each matmul | activations at block boundaries | all-reduce (or reduce-scatter and all-gather) per block |
| Context parallel | sequence positions | parameters | ring exchange of keys and values inside attention |
| Expert parallel | MoE experts | non-expert weights | all-to-all dispatch and combine |
Everything in these notes elaborates one row of that table, or explains a prerequisite you need before a row makes sense.
The prerequisites are not optional
Three ideas do most of the work throughout the series, and they are worth stating up front.
Gradients are sums
The loss over a batch is a sum over examples, and differentiation is linear. That single fact is what makes data parallelism correct rather than an approximation.
Arithmetic intensity decides
Whether a kernel is compute-bound or memory-bound follows from FLOPs per byte moved. It explains the KV cache, MLA, and why some communication can hide behind compute.
Collectives are the vocabulary
All-reduce, all-gather, reduce-scatter and all-to-all are the only primitives you need. Every strategy above is a choice of which one to call and where.
The roadmap
The video's 30 chapters group naturally into five arcs. Each note below is self-contained, but they are ordered so that nothing depends on something later.
Arc 1 · The model you are going to train
Count every parameter in a decoder-only transformer, derive the $6ND$ training-FLOP rule from the forward and backward passes, and turn it into wall-clock estimates and model FLOP utilization.
Chapter 2Derive rotary position embeddings from the requirement that attention scores depend only on relative position, then extend the context window with position interpolation, NTK-aware scaling and YaRN.
Chapters 3–4Pre-norm versus post-norm, RMSNorm, SwiGLU, tied embeddings, and why residual-branch initialization has to be scaled by depth.
Chapter 5Why autoregressive decoding is memory-bound, how to size the KV cache, the roofline model, and how MQA and GQA trade quality for bandwidth.
Chapter 6Compress the KV cache into a low-rank latent, discover why RoPE breaks the trick, fix it with a decoupled positional path, and absorb the projections into the query and output matrices.
Chapters 7–10Arc 2 · Why distributed training is correct
Reverse-mode differentiation as vector-Jacobian products, the two rules that govern split computation graphs (sum at forks, route at concatenations), and the proof that averaging gradients across shards is exact.
Chapters 11–12Broadcast, reduce, all-reduce, all-gather, reduce-scatter and all-to-all, with ring algorithms, bandwidth-optimal cost models, and how to predict what a collective will cost you.
Chapter 19Arc 3 · Splitting the batch
Gradient bucketing and backward overlap in DDP, the memory ledger that motivates ZeRO, and how FSDP turns a replicated model into a sharded one with all-gather and reduce-scatter.
Chapters 16, 20Arc 4 · Splitting the model
Cut the model by layer, discover the bubble, fix most of it with microbatching, and work out the activation-memory bill that pipelining creates.
Chapters 14, 17The same bubble fraction, wildly different memory. Interleaved stages, and how splitting the backward pass into input-gradient and weight-gradient halves buys you a near-zero bubble.
Chapter 15Column-parallel and row-parallel linear layers, the conjugate pair of identity/all-reduce operators, how to shard attention head-wise for free, and sequence parallelism for the normalization layers.
Chapters 21–22Shard the sequence itself. Online softmax makes attention decomposable over key blocks, a ring passes those blocks around, and causal masking creates a load-balancing problem with a neat fix.
Chapter 23Treat the world as an n-dimensional grid of ranks. Which axis should be innermost, why the ordering follows communication volume, and how 3D and 4D parallelism are assembled.
Chapters 18, 25Arc 5 · Sparsity and the training loop
Decouple parameter count from FLOPs per token. Top-$k$ routing, the capacity factor, load-balancing losses, shared experts, and why the router is the hard part.
Chapter 26Place experts on different devices, ship each token to the ranks that own its experts with an all-to-all, and combine the results on the way back. Plus tensor parallelism for MoE and the two combined.
Chapters 27–30Gradient accumulation done correctly under DDP, clipping across shards, mixed precision, learning-rate schedules, throughput metrics that are not lies, and distributed checkpoints that survive a resharding.
Chapters 13, 24Notation used throughout
| Symbol | Meaning |
|---|---|
| $B$, $S$ | batch size (sequences) and sequence length (tokens per sequence) |
| $d$ or $d_{\text{model}}$ | model hidden size |
| $L$ | number of transformer blocks |
| $h$, $d_h$ | number of attention heads and per-head dimension, $d=h\,d_h$ |
| $d_{\text{ff}}$ | feed-forward inner dimension |
| $V$ | vocabulary size |
| $N$ | total non-embedding parameter count |
| $P$ | number of ranks (devices) in a group |
| $\alpha$, $\beta$ | link latency and inverse bandwidth in communication cost models |
If you only read two notes, read the autograd one and the collectives one. Every parallelism strategy in the rest of the series is a combination of those two ideas: what has to be summed, and which primitive sums it.