The Training Loop, Metrics & Checkpointing
The skeleton
Every line has a way to be subtly wrong across ranks. Taking them in turn.
Accumulation, correctly
Scale before, not after
Divide each micro-batch loss by $k$ before calling backward. Scaling the accumulated gradient afterwards works mathematically but lets the buffer hold $k\times$ the intended magnitude in between, which interacts badly with FP16 loss scaling and with overflow checks.
Suppress the sync
DDP all-reduces on every backward by default. Across $k$ micro-steps that is $k$ times the necessary traffic. Frameworks expose a no-sync context; use it for the first $k-1$ micro-steps. Forgetting costs throughput only, never correctness.
Clipping is a global operation
Gradient clipping rescales the whole gradient when its norm exceeds a threshold:
The norm is over all parameters, which no single rank possesses under any sharded strategy. The correct procedure is to reduce first and clip second:
Each rank computes the sum of squares of the gradient shards it owns.
All-reduce those scalars with sum, then take the square root. Every rank now holds the identical global norm.
Apply the same coefficient locally.
Mixed precision
| FP16 | BF16 | |
|---|---|---|
| Exponent bits | 5 | 8, the same range as FP32 |
| Mantissa bits | 10 | 7 |
| Underflow risk | High — small gradients flush to zero | Negligible |
| Loss scaling needed | Yes, dynamic, with skipped steps on overflow | No |
| Verdict | Legacy | Default wherever available |
BF16 trades precision for range, which is the right trade: gradients span many orders of magnitude but rarely need seven significant digits. Two things still stay in FP32 — the master weights, so tiny updates are not lost when added to large weights, and the reductions inside normalization and softmax, where accumulated error compounds.
AdamW and what to exclude
AdamW decouples weight decay from the adaptive update, applying it directly to the parameter:
Typical settings are $\beta_1=0.9$, $\beta_2=0.95$ — lower than the 0.999 default, because language-model gradients are noisy and a shorter second-moment window adapts faster — with $\epsilon=10^{-8}$ and $\lambda=0.1$.
Learning-rate schedules
Warmup exists because Adam's second-moment estimate is unreliable in the first few hundred steps: $\hat v$ is based on almost no data, so the update direction is close to noise, and a full-size step early can destabilize a run permanently. Linear warmup over 1–2% of training costs nothing and removes the risk.
Warmup then cosine
The long-standing default. Requires committing to a total step count up front: stopping early leaves the rate high and the model under-annealed, and extending distorts the shape.
Warmup, stable, decay
Hold the peak rate indefinitely, then decay over the final stretch. The stable phase is open-ended, so you can decide the length later and branch a short decay to get a deployable checkpoint at any time.
Metrics that are not lies
| Metric | How to get it wrong | How to get it right |
|---|---|---|
| Training loss | Report rank 0's local loss. | All-reduce the loss sum and the token count, then divide. |
| Tokens per second | Count padding. | Count only real tokens, and use end-to-end wall time including the data loader. |
| MFU | Include recomputation as useful work. | Use $6N$ per token. Counting recomputation gives hardware FLOP utilization, which is a different, larger number — say which one you mean. |
| Step time | Time only the compute. | Include the wait on collectives; that wait is the thing you are trying to remove. |
| Gradient norm | Log the post-clip value. | Log pre-clip. It is the best early warning of instability: a spike precedes a loss spike. |
Checkpoints that survive a resharding
A run of any length will be interrupted. What must be saved to resume exactly:
| Item | Why |
|---|---|
| Model parameters | Obvious. |
| Optimizer state | Adam moments carry substantial history; discarding them causes a visible loss spike on resume. |
| Step count and schedule position | Otherwise the learning rate restarts from warmup. |
| Data loader position | Otherwise the run silently re-trains on data it has already seen. |
| RNG states, per rank | For bitwise reproducibility of dropout and any stochastic path. |
Scale the loss before backward, suppress the sync until the last micro-step, all-reduce the norm before clipping, keep master weights in FP32, exclude norms and biases from weight decay, measure MFU against $6N$, and save checkpoints in a form that outlives the mesh you happened to use.
Check yourself
- Explain why the loss is scaled by $1/k$ before backward rather than after. reasoning
- Describe the correct sequence for gradient clipping under FSDP. design
- Explain why BF16 needs no loss scaling but FP16 does. reasoning
- Name two parameter categories excluded from weight decay and justify each. design
- Explain why warmup-stable-decay allows extending a run and cosine does not. analysis
- List the five items a resumable checkpoint must contain. recall