Skip to content
Snula
Curriculum
RU Open

Curriculum

Week 10. Training

Phase 2. The modern transformer · week 10 of 24

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 EOS token 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_entropy divides by mask.sum()). With gradient accumulation (grad_accum microbatches), 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; on inf the step is skipped and s is reduced. bf16: 8 exponent bits, like fp32, and 7 mantissa bits: same range, no underflow, but coarser precision (D12)
  • autocast keeps 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 (see 05-ГЛУБИНА)

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), exp overflow 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, v and the step t for 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 requires torch.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: in nanolm/README.en.md perplexity 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 by test_masked_cross_entropy_ignores_masked_positions
  • Resuming without the optimizer state: t resets, m and v are zeroed, and the first steps are full-size (test_bias_correction_makes_first_step_full_size shows 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 for batch × 4 and for batch with grad_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

  1. The loss became NaN at step 3000. What do you do, in order?
  2. Why keep master weights in fp32 if forward and backward run in bf16?
  3. What must a checkpoint contain so that resuming is indistinguishable from uninterrupted training?

✅ Checkpoint 2

  1. A transformer from an empty file in ≤40 minutes, passing the tests
  2. A trained model generates coherent text
  3. Out loud, in 10 minutes: a full memory and FLOPs calculation for a 7B model

In the app each week has skills to rate yourself on, questions with answer checking, Python coding problems and a tutor grounded in the course.

Learn in the app: tutor, coding problems
← PreviousWeek 9. Bookkeeping: parameters, FLOPs, memory Next →Week 11. Inference

Snula
Snula: LLMs from scratch

  • Home
  • Curriculum
  • App
  • Privacy
  • Terms

The course text is licensed under CC BY-NC-SA 4.0, nanolm code under Apache-2.0.