Week 10. Training
Learn in the app: tutor, coding problems →
Core: packing and the cross-document mask, bf16 and master weights, loss diagnostics, checkpoints · Depth: health checks and dashboards in detail · ≈ 12 h core / 22 h total
Your first real training run, and with it your first silent bugs. A model with a bug in the data or the loss usually still trains: the loss goes down, it produces text, just worse than it could. The goal of this week: learn to tell a bug from a property of the data, and read the logs well enough to see trouble coming.
Step 1. Data and loss
In plain terms. Documents of different lengths are glued into one long stream and cut into windows of length S. This is packing
(packing without padding). A window can start with the end of a recipe and continue with the start of a news story.
Without a special mask, the first word of the news story "sees" the recipe and learns to continue it. The model spends
capacity on a connection that does not exist in the data. A cross-document attention mask forbids looking across the boundary.
- Data loading: packing, document boundaries, cross-document attention masking. An
EOStoken goes between documents; a block-diagonal mask (attention only within a document) and resetting RoPE positions remain options. Llama 3 masks attention across boundaries and notes that this matters most at long context
In plain terms. A batch has two examples: the first has 2 tokens with loss 4, the second has 8 tokens with loss 1.
The per-token mean: (2·4 + 8·1)/10 = 1.6. The per-example mean: (4 + 1)/2 = 2.5. These are different objectives:
in the second, each token of the short example weighs 4 times more. The choice of normalization is not a detail, it is a decision.
- The loss is the mean over the counted tokens of the whole batch (
masked_cross_entropydivides bymask.sum()). With gradient accumulation (grad_accummicrobatches), each microbatch loss is divided by their count. If the number of tokens differs between microbatches, the mean of means ≠ the per-token mean: normalize by the total token count
Step 2. Numerical precision
In plain terms. Suppose a weight is stored with three significant digits: 1.00. Add 0.001: you get 1.001,
and rounding brings it back to 1.00. Do that a thousand times and the weight does not move, even though it should have become 2.00.
bf16 keeps only 2–3 significant digits, and an optimizer step is exactly this kind of tiny increment.
That is why updates are accumulated in an fp32 copy of the weights (master weights), and bf16 is used only for compute.
- Mixed precision: bf16 vs fp16, master weights in fp32, loss scaling (and why bf16 does not need it).
fp16: 5 exponent bits, 10 mantissa bits, a maximum of 65,504, and small gradients underflow to zero. Loss scaling
multiplies the loss by
s(for example, 2¹⁶), which shifts gradients into the representable range; before clipping and the step they are divided back; oninfthe step is skipped andsis reduced. bf16: 8 exponent bits, like fp32, and 7 mantissa bits: same range, no underflow, but coarser precision (D12) autocastkeeps weights in fp32 and picks the precision per operation: matmuls in bf16, softmax and norms in fp32. Loading the whole model in bf16 is something else: simpler, but riskier for training (see05-ГЛУБИНА)
Step 3. Diagnostics and checkpoints
- Training diagnostics: loss spikes, NaN, exploding gradient norm, dead neurons. A loss spike is usually
preceded by a grad norm spike. Common causes: the end of warmup with too high an LR, a bad batch,
growing attention logits (fixed by QK-norm, week 7). NaN:
log(0),expoverflow in fp16, a mask row made entirely of−inf(softmax divides 0 by 0) - What to watch in wandb: loss, grad norm, LR, weight norm per layer, fraction of clipped steps, tokens per second
- Checkpoints, resuming, determinism. A checkpoint holds the weights, the optimizer state (
m,vand the steptfor bias correction), the schedule step, the states of all RNGs, the position in the data, the loss scaler's scale. Bit-for-bit reproducibility on GPU requirestorch.use_deterministic_algorithms(True); MPS does not have it - Health checks before a long run: the initial loss ≈
ln V; the model can memorize a single batch. The loss has a floor, the entropy of the data: innanolm/README.en.mdperplexity levels off at 3.4, and that is not a bug
Common mistakes
- Weight decay on norms and biases: caught by
test_param_groups_exclude_norms_and_biases. Loss on the prompt or on padding: caught bytest_masked_cross_entropy_ignores_masked_positions - Resuming without the optimizer state:
tresets,mandvare zeroed, and the first steps are full-size (test_bias_correction_makes_first_step_full_sizeshows how large such a step is). The result: a loss spike - Forgetting to divide by
grad_accum, or clipping before all the backward passes are done. There is no test: compare the curves forbatch × 4and forbatchwithgrad_accum = 4, they must match - Leakage between documents with packing. nanolm has no test for it; write one modeled on
test_model_is_causal: change a token in document A, and the logits of document B must not change
Code → nanolm/train.py: train your own model on TinyStories / Shakespeare. Get meaningful generations.
This is your first real LM. Inside: get_batch (packing without regard to boundaries, read the docstring),
estimate_loss, train (cosine_lr, AdamW(get_param_groups(...)), grad_accum, clip_grad_norm_ before step).
Quick run: python -m nanolm.train --data data/toy.txt --steps 300 --vocab-size 400 --lr 3e-3.
Checks: pytest tests/test_model.py tests/test_optim.py -v, the key one is test_model_can_overfit_one_batch.
Math (Track D): D11: MLE on three distributions, and why LM cross-entropy is exactly MLE; D12: why updates get lost in bf16, format ranges, stochastic rounding.
Interview question of the week: "The loss jumped 10× at step 40,000 and is slowly coming back. What do you do?" A 3-minute structure: grad norm and LR before the spike → reproduce from a checkpoint on the same batch (data or numerics?) → attention logits and range overflow → fixes: roll back and skip the batch, lower LR, QK-norm, clipping → prevention: monitoring grad norm per layer.
Sources: Micikevicius et al., Mixed Precision Training (2018); Kalamkar et al., bfloat16 (2019); Llama 3 technical report (2024), which describes the cross-document mask.
Week outcomes
- I can train a model to coherent generation and explain every curve in the training logs.
- I can diagnose a loss spike from grad norm, LR and per-layer weight norms.
- I can implement packing with a cross-document attention mask.
- I can explain why fp16 needs loss scaling and bf16 does not.
- I can write a test that catches attention leakage between documents.
Self-check
- The loss became NaN at step 3000. What do you do, in order?
- Why keep master weights in fp32 if forward and backward run in bf16?
- What must a checkpoint contain so that resuming is indistinguishable from uninterrupted training?
✅ Checkpoint 2
- A transformer from an empty file in ≤40 minutes, passing the tests
- A trained model generates coherent text
- Out loud, in 10 minutes: a full memory and FLOPs calculation for a 7B model