How Models Are Actually Trained
Reading a training run while it happens
The loss curve is the last thing to move when a training run goes wrong, so serious runs watch five or six other numbers that move first.
- 14 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
A training run is watched like a patient on a hospital monitor. The loss curve is one line among several, and it is the slowest to react.
The analogy you have already lived
Stand next to a hospital bed and look at the monitor. Heart rate, breathing, oxygen level, blood pressure — four or five traces, all at once.
The nurse does not watch the patient's face and wait for them to look unwell. By then it is late. She watches the traces, because they move first.
A training run gets the same treatment. Nobody stares at the loss curve. They stare at the panel of numbers that change before the loss does.
Why it exists
A large training run costs a fortune and takes weeks. If something breaks on day three and you notice on day nine, you have burnt six days.
Worse, some failures never show up in the loss at all. A run can look perfectly healthy and still be wasting half its hardware, or quietly training on the wrong data.
So you build a dashboard, and you learn what healthy looks like before you need to spot unhealthy.
The numbers people actually watch
Loss. How wrong the model is. Should fall fast, then slowly. Everybody watches this one.
Gradient size. How large a correction each step wants. Should be steady and slowly shrinking. A jump here comes before the loss reacts.
Learning rate. How big the steps are. You set this yourself, so it should match your plan exactly. Plot it anyway — misconfigured schedules are one of the most common bugs.
Speed. Tokens processed per second, and what fraction of the hardware's theoretical peak that represents. A drop means a slow disk, a network problem, or a stalled machine.
Update size relative to weight size. How far the weights moved this step, compared with how big they are. Too large and the model is thrashing. Too small and it has stopped learning.
Confidence. How spread out the model's guesses are. Falls steadily as it learns. A sudden collapse to total confidence usually means something is broken.
How it looks
healthy something is wrong
loss \___ loss \____/‾‾‾
grad ~~~~~~~~___ grad ~~~~^~~~~~ <- moved first
speed --------- speed ------____ <- machine dropped out
conf \_____ conf \____| <- collapsedThe point of the picture is the second column. The gradient line spiked, and the speed line fell, before the loss line did anything visible.
What is honestly hard here
Knowing what "normal" looks like is not something you can read. You have to watch a few healthy runs first.
That is genuinely frustrating, and everyone goes through it. The practical shortcut is a tiny model on a tiny dataset, run for a few minutes. Look hard at every number. Keep that as your reference picture.
Where you have already seen this
- A hospital monitor beside a patient.
- The temperature gauge on a long drive, which you check more than the speedometer.
- A cricket scoreboard showing run rate and wickets, not only the total.
Remember this
- The loss curve reacts last. Watch gradient size and speed to see trouble early.
- Plot the learning rate even though you chose it. Schedule bugs are common.
- Learn what healthy looks like on a tiny run before you need to spot unhealthy.
What to learn next
- Reading a loss curve — the shapes and what causes them.
- Experiment tracking — where all of this gets logged.
- Instruction tuning — what happens to the model after pretraining ends.
Developer — Code and libraries.
Setup
pip install torchRuns on a CPU in about fifteen seconds.
A training loop with a real dashboard
import time
import torch
import torch.nn as nn
import torch.nn.functional as F
torch.manual_seed(0)
TEXT = "the cat sat on the mat. the cat ate the rat. the rat sat on the mat. " * 60
chars = sorted(set(TEXT))
stoi = {c: i for i, c in enumerate(chars)}
V, CTX = len(chars), 32
data = torch.tensor([stoi[c] for c in TEXT])
x = torch.stack([data[i:i + CTX] for i in range(0, len(data) - CTX - 1, 4)])
y = torch.stack([data[i + 1:i + CTX + 1] for i in range(0, len(data) - CTX - 1, 4)])
layer = nn.TransformerEncoderLayer(64, 4, 128, batch_first=True, dropout=0.0)
model = nn.ModuleDict({"emb": nn.Embedding(V, 64), "pos": nn.Embedding(CTX, 64),
"blocks": nn.TransformerEncoder(layer, 2), "head": nn.Linear(64, V)})
opt = torch.optim.AdamW(model.parameters(), lr=1e-3)
N = sum(p.numel() for p in model.parameters())
WARM, TOTAL = 20, 120
print(f"model: {N:,} parameters, vocab {V}, context {CTX}, batch {len(x)}")
print(f"{'step':>5} {'lr':>8} {'loss':>7} {'ppl':>7} {'|g|':>7} "
f"{'upd/w':>8} {'H(p)':>6} {'tok/s':>9}")
prev = {k: p.detach().clone() for k, p in model.named_parameters()}
for step in range(1, TOTAL + 1):
lr = 1e-3 * min(step / WARM, 1.0)
for g in opt.param_groups:
g["lr"] = lr
t0 = time.perf_counter()
h = model["emb"](x) + model["pos"](torch.arange(CTX))
mask = nn.Transformer.generate_square_subsequent_mask(CTX)
logits = model["head"](model["blocks"](h, mask=mask, is_causal=True))
loss = F.cross_entropy(logits.reshape(-1, V), y.reshape(-1))
opt.zero_grad()
loss.backward()
gnorm = nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
dt = time.perf_counter() - t0
if step % 20 == 0 or step == 1:
with torch.no_grad():
# update-to-weight ratio: how far the weights moved, relative to their size
num = sum((p - prev[k]).pow(2).sum() for k, p in model.named_parameters())
den = sum(p.pow(2).sum() for _, p in model.named_parameters())
ratio = (num.sqrt() / den.sqrt()).item()
probs = logits.softmax(-1)
entropy = -(probs * probs.clamp_min(1e-9).log()).sum(-1).mean().item()
print(f"{step:>5} {lr:>8.2e} {loss.item():>7.3f} {loss.exp().item():>7.2f} "
f"{gnorm.item():>7.3f} {ratio:>8.5f} {entropy:>6.3f} "
f"{x.numel() / dt:>9,.0f}")
prev = {k: p.detach().clone() for k, p in model.named_parameters()}
print("\nper-parameter-group gradient norms at the last step:")
for name, p in model.named_parameters():
if p.grad is not None and p.dim() > 1:
print(f" {name:<40} {p.grad.norm().item():>8.4f}")model: 70,540 parameters, vocab 12, context 32, batch 1027
step lr loss ppl |g| upd/w H(p) tok/s
1 5.00e-05 2.559 12.92 1.677 0.00023 2.327 285,185
20 1.00e-03 1.541 4.67 0.851 0.04094 1.997 383,277
40 1.00e-03 0.786 2.19 0.265 0.05926 1.089 396,927
60 1.00e-03 0.602 1.83 0.188 0.03788 0.817 421,555
80 1.00e-03 0.373 1.45 0.175 0.04221 0.619 433,929
100 1.00e-03 0.200 1.22 0.200 0.03983 0.418 375,345
120 1.00e-03 0.127 1.14 0.222 0.02762 0.271 351,806
per-parameter-group gradient norms at the last step:
emb.weight 0.0226
pos.weight 0.0072
blocks.layers.0.self_attn.in_proj_weight 0.1658
blocks.layers.0.self_attn.out_proj.weight 0.0555
blocks.layers.0.linear1.weight 0.0407
blocks.layers.0.linear2.weight 0.0679
blocks.layers.1.self_attn.in_proj_weight 0.0715
blocks.layers.1.self_attn.out_proj.weight 0.0217
blocks.layers.1.linear1.weight 0.0138
blocks.layers.1.linear2.weight 0.0185
head.weight 0.0627Written against PyTorch 2.5.1, CPU. The tok/s column is timing-based and will not match on your machine — it depends on your CPU, its load, and thermal state. Every other column is seeded and reproducible.
What each column tells you
lr rises from 5.00e-05 to 1.00e-03 over 20 steps and then holds. That is the warmup you configured, confirmed rather than assumed. Log it every run.
|g| is the gradient norm before clipping, from clip_grad_norm_'s return value. It falls 1.677 → 0.175 and then ticks back up. Falling is healthy. The uptick at the end is the model fitting harder on a nearly-solved task. A sudden 100× jump is a loss spike arriving.
upd/w is the ratio of the update norm to the weight norm, aggregated over 20 steps. It should sit in a band; here it hovers around 0.03–0.06 for a 20-step window, which is roughly 0.002 per step. The classic single-step rule of thumb is around 1e-3. Two orders of magnitude either side is a signal: too big means the learning rate is too high, too small means that part of the model has stopped moving.
H(p) is the mean entropy of the output distribution, in nats. It falls 2.327 → 0.271 as the model becomes confident. A sudden collapse to near zero, rather than a smooth fall, usually means the model found a degenerate shortcut — predicting one token everywhere, or leaking labels.
Per-parameter gradient norms are the diagnostic people reach for last and should reach for first. Layer 0's attention projection has an 8× larger gradient than layer 1's linear1. That spread is normal. A layer whose norm is exactly 0.0000 is receiving no signal — a detached tensor, a frozen module, or a dead branch.
Adding MFU on a GPU
Model FLOPs Utilisation answers "what fraction of the hardware am I actually using?"
def mfu(params, tokens_per_sec, peak_flops):
"""params = non-embedding parameter count; peak_flops = device bf16 peak."""
return 6 * params * tokens_per_sec / peak_flops
# example: a 1.3B model at 40,000 tokens/s on one H100 (989e12 bf16 peak, dense)
print(f"{mfu(1.3e9, 40_000, 989e12):.1%}")31.5%
That figure is arithmetic from the numbers you pass in, so it reproduces exactly — but the 40_000 is an example, not a measurement from your hardware. Substitute your own measured throughput.
Healthy dense pretraining sits at 35–55% MFU. Below 20% and something is wrong: input pipeline starvation, too-small batches, unfused kernels, or communication overhead. See is the GPU waiting for data.
The minimum dashboard for a real run
| Metric | Healthy | What a change means |
|---|---|---|
| train loss | smooth fall | spike → see clipping lesson |
| val loss | tracks train | diverges → overfitting or data bug |
| grad norm (pre-clip) | steady, slow fall | jump → instability incoming |
| learning rate | exactly your schedule | mismatch → config bug |
| tokens/sec, MFU | flat | drop → a node, disk or network problem |
| update/weight ratio | ~1e-3 per step | too high → LR too large |
| output entropy | smooth fall | collapse → degenerate solution |
| GPU memory | flat | creeping up → a leak |
| fraction of tokens that are padding | near zero | high → packing broken |
Log every one of these to MLflow or an equivalent, at a fixed step interval, from step 1. Adding a metric halfway through a run makes the two halves incomparable.
Common mistakes
Logging only at evaluation intervals. Spikes last a few steps. Log the cheap scalars every step, or at worst every ten.
Logging after clipping. The clipped norm is capped at your threshold and carries no information. Log the return value of clip_grad_norm_.
Smoothing the chart before looking at it. Exponential smoothing on a dashboard hides exactly the single-step events you are hunting. Keep an unsmoothed view. See smoothing and lying charts.
Reporting throughput including padding. Count real tokens.
Not saving the config alongside the metrics. Six weeks later, "which data mixture was run 47?" has to be answerable. See config files, not arguments.
Checkpointing too rarely. The recovery procedure for a bad spike is to rewind. If your checkpoints are 12 hours apart, you lose 12 hours.
Try it yourself
Add model["emb"].weight.requires_grad_(False) before the loop, and re-run. Watch emb.weight vanish from the gradient report while the loss keeps falling. That silent disappearance is what a real freezing bug looks like.
What to learn next
- Reading a loss curve — the shapes and what causes them.
- Experiment tracking — where all of this gets logged.
- Instruction tuning — what happens to the model after pretraining ends.
Researcher — Mathematics and papers.
What the update-to-weight ratio measures
For parameter tensor $W$ with update $\Delta W$ at step $t$, the diagnostic is
$$ \rho_t = \frac{\lVert \Delta W_t \rVert_F}{\lVert W_t \rVert_F} $$
Under Adam, $\lVert \Delta W \rVert \approx \eta \sqrt{n}$ for an $n$-element tensor, because each coordinate's update has magnitude close to $\eta$. So $\rho$ is roughly $\eta \sqrt{n} / \lVert W \rVert$, and for a tensor at a standard initialisation scale this lands near $10^{-3}$ at typical learning rates. The empirical rule "$\rho \approx 10^{-3}$" is therefore a statement about $\eta$ relative to the initialisation scale, not a universal constant.
The useful version is per-tensor. A global $\rho$ hides the case where the embedding matrix is moving 100× faster than the attention projections, which is exactly the pathology that layer-wise learning rates and µP were designed to remove.
MFU and HFU
Model FLOPs Utilisation (Chowdhery et al., 2022, PaLM):
$$ \mathrm{MFU} = \frac{6 N \cdot \text{tokens/sec}}{\text{device peak FLOP/s}} $$
It counts only the FLOPs the model definition requires. Hardware FLOPs Utilisation counts the FLOPs actually executed, which includes recomputation from gradient checkpointing. HFU exceeds MFU whenever activations are recomputed, and MFU is the metric that matters for cost, because recomputed FLOPs are overhead.
A more complete accounting adds the attention term, which the $6N$ approximation omits:
$$ C_{\text{step}} \approx 6ND + 12\,L\,H\,d\,s\,D $$
for $L$ layers, $H$ heads, head dimension $d$, sequence length $s$. At $s = 2048$ and a few billion parameters the attention term is a few percent; at $s = 128{,}000$ it dominates, and MFU computed with $6N$ alone becomes meaningless.
Reference points on H100 SXM: 989 TFLOP/s dense bf16, roughly 1979 with 2:4 structured sparsity (rarely achieved). Published dense pretraining MFU sits at 35–55%; MoE runs report lower, because expert routing is communication-bound.
Signals that lead the loss
Ordered by how early they move:
- Maximum attention logit. Divergence here precedes everything else. Wortsman et al., 2023 track $\max_{ij} q_i \cdot k_j / \sqrt{d}$ per layer and show it growing steadily before instability; qk-layernorm bounds it.
- Output logit magnitude / softmax normaliser $Z$. A z-loss $\lambda \log^2 Z$ both monitors and controls this.
- Gradient-norm distribution. Not the mean — the tail. Track a high percentile.
- Per-tensor update ratios. Catches a layer going dead or running away.
- Loss. Last.
To this, production runs add infrastructure signals that have nothing to do with the model: per-rank step time (to find the straggler), all-reduce time as a fraction of step time, dataloader wait time, and ECC error counts. At thousand-GPU scale, hardware failure is a routine event rather than an exception, and the monitoring must distinguish "the model is diverging" from "rank 412 fell off the network".
Evaluation during training
Validation loss on a held-out slice of the same distribution measures optimisation. It does not measure capability, and it is uninformative about downstream behaviour — see perplexity.
Practical compromise used by most open pretraining efforts:
- Every N steps: validation loss on several held-out domains separately, never pooled. A rise on code with a fall on web text is a mixture problem, and pooling hides it.
- Every M × N steps: a small fast benchmark suite (a few thousand examples of multiple-choice tasks) scored by likelihood rather than generation, so it needs no sampling.
- At cooldown only: the expensive generative evaluations.
Benchmark scores during pretraining are noisy at small scale and near-flat until surprisingly late. Treat a flat benchmark line at 10% of the run as uninformative rather than as bad news.
Papers and tools
- Chowdhery et al., PaLM, 2022 — arxiv.org/abs/2204.02311 (defines MFU)
- Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models, 2022 — arxiv.org/abs/2205.05198 (MFU vs HFU)
- Wortsman et al., Small-scale proxies for large-scale Transformer training instabilities, 2023 — arxiv.org/abs/2309.14322
- Groeneveld et al., OLMo: Accelerating the Science of Language Models, 2024 — arxiv.org/abs/2402.00838 — releases full training logs, the best public reference for what a healthy run looks like.
What to learn next
- Reading a loss curve — the shapes and what causes them.
- Experiment tracking — where all of this gets logged.
- Instruction tuning — what happens to the model after pretraining ends.