← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 16 · Chapters 13, 24

The Training Loop, Metrics & Checkpointing

The loop is where the abstractions leak. Every strategy in this series assumed a step that computes a gradient and applies it. Doing that correctly across shards — accumulating, clipping, scaling, scheduling and saving — is where the subtle, silent bugs live.

The skeleton

One optimizer step with gradient accumulation
1for micro-step $j=1,\dots,k$:
2fetch micro-batch, forward, compute loss
3scale the loss by $1/k$ before backward
4backward, with gradient sync suppressed for $j<k$
5all-reduce the gradients (happens implicitly on micro-step $k$)
6all-reduce the squared gradient norm, then clip with the global coefficient
7set the learning rate from the schedule
8optimizer step, then zero the gradients

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.

The token-weighting trap, once more. A language-model loss is a mean over tokens. If micro-batches contain different numbers of real tokens, dividing each by $k$ weights them equally and is therefore wrong. Accumulate the loss sum and the token count separately, all-reduce both, and divide at the end. With fixed-length packed sequences the issue disappears, which is one more reason packing is standard.

Clipping is a global operation

Gradient clipping rescales the whole gradient when its norm exceeds a threshold:

$$\mathbf g\leftarrow\mathbf g\cdot\min\!\left(1,\frac{\tau}{\|\mathbf g\|_2}\right).$$

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:

1

Each rank computes the sum of squares of the gradient shards it owns.

2

All-reduce those scalars with sum, then take the square root. Every rank now holds the identical global norm.

3

Apply the same coefficient locally.

Clipping per rank is a real bug, not an approximation. Different ranks compute different coefficients, apply different updates, and replicas that were supposed to stay identical drift apart. Since the loss still decreases, this can survive an entire run undetected.

Mixed precision

FP16BF16
Exponent bits58, the same range as FP32
Mantissa bits107
Underflow riskHigh — small gradients flush to zeroNegligible
Loss scaling neededYes, dynamic, with skipped steps on overflowNo
VerdictLegacyDefault 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:

$$\theta\leftarrow\theta-\eta\left(\frac{\hat m}{\sqrt{\hat v}+\epsilon}+\lambda\theta\right).$$
Do not decay everything. Weight decay should apply to matrices that do real linear work. Exclude normalization gains — decaying them toward zero shrinks the activations they are supposed to rescale — and biases, which carry no capacity worth regularizing. Building two parameter groups, decayed and not, is a five-line change that noticeably affects final loss.

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.

steps at peak rate
mean rate / peak
final rate / peak
Cosine must know the total step count in advance. Warmup-stable-decay does not, so a run can be extended, and a usable checkpoint can be produced at any point by branching off a short decay.

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

MetricHow to get it wrongHow to get it right
Training lossReport rank 0's local loss.All-reduce the loss sum and the token count, then divide.
Tokens per secondCount padding.Count only real tokens, and use end-to-end wall time including the data loader.
MFUInclude 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 timeTime only the compute.Include the wait on collectives; that wait is the thing you are trying to remove.
Gradient normLog 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:

ItemWhy
Model parametersObvious.
Optimizer stateAdam moments carry substantial history; discarding them causes a visible loss spike on resume.
Step count and schedule positionOtherwise the learning rate restarts from warmup.
Data loader positionOtherwise the run silently re-trains on data it has already seen.
RNG states, per rankFor bitwise reproducibility of dropout and any stochastic path.
Save in a layout-independent form. Writing each rank's raw shard bakes the parallelism layout into the file, so a job resumed on a different mesh cannot load it. Distributed checkpointing instead records each shard with metadata describing which slice of which logical tensor it holds, so any layout can reassemble it. Given how often you will change the mesh mid-project, this is worth doing from the start.
Write asynchronously, and keep more than one. A synchronous multi-terabyte save stalls every rank. Copy to host memory, then write in the background. And never overwrite the only checkpoint in place: a crash mid-write leaves nothing to resume from.
Takeaway

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