Training diagnostics
Be able to read loss curves, gradient norms and learning rate sweeps and decide what is wrong.
Prerequisites
- DTrain a neural network in PyTorchrequired
- DVisualisation with Matplotlibrequired
Intuition
When the training does not work there is a handful of causes, and they look different in the curves. Being able to read them saves days.
| The pattern | The diagnosis | The measure |
|---|---|---|
| Both curves fall and flatten out together | healthy | done |
| Training falls, validation turns upwards | overfitting | early stopping, more data, regularisation |
| Both flatten out high | underfitting | a bigger model, longer training, better features |
| The curve jumps wildly | too high a learning rate or too small a batch | lower lr, increase the batch |
| The curve is completely flat from the start | nothing is learning at all | see below |
| The curve goes to NaN | an explosion or a division by zero | gradient clipping, a lower lr, check the data |
The flat curve is the most confusing one. It is nearly always due to something trivial: a forgotten optimizer.step(), requires_grad=False somewhere, a learning rate of 0, or the loss being computed on the wrong tensor.
Rule number one in all debugging: get the model to overfit 10 examples first. If it cannot memorise ten data points the fault is in the code, not in the hyperparameters.
Formal
Quantities to log from the start, not when it has already gone wrong:
| The quantity | A healthy level | What a deviation means |
|---|---|---|
| The training loss | falling | flat = the bug is in the code |
| The validation loss | falling, then flattening | turning = overfitting |
| The gradient norm | 0.1–10 | → 0: vanishing; > 100: an explosion |
| The share of active ReLUs | 0.2–0.7 | near 0: dead neurons |
| The norm of the weights | growing slowly | racing away: too weak a weight decay |
| The update-to-weight ratio | ~1e-3 | > 1e-2: the lr is too high; < 1e-4: too low |
The last row is the most underrated measure. Log per layer: it says straight away whether the learning rate is reasonable, independently of the scale of the model.
A learning rate sweep (an LR range test) finds the right order of magnitude in a few minutes:
- Start at a very small lr (1e-7).
- Increase it exponentially every batch up to ~1.
- Plot the loss against the lr on a log scale.
- Choose an lr just before the minimum — typically a tenth of the lr where the loss is lowest, since the curve is already on its way to becoming unstable there.
The debugging order, one after the other:
| # | The check |
|---|---|
| 1 | Can the model overfit 10 examples to nearly zero loss? |
| 2 | Is the loss at the start reasonable? (ln(the number of classes) for cross-entropy) |
| 3 | Are the labels correctly tied to the input? |
| 4 | Is the normalisation the same in training and in evaluation? |
| 5 | Run the LR sweep |
| 6 | Only now: the other hyperparameters |
Check 2 is free and reveals a lot. With 10 classes the first loss should be roughly . If it is 7 something is wrong with the initialisation or the labels; if it is 0.1 the answer key is leaking into the model.
Code
import math, torch, torch.nn as nn
# 1. Can the model overfit 10 examples? ALWAYS do this first.
def can_overfit(model, x, y, steps=300, lr=1e-3):
x, y = x[:10], y[:10]
opt = torch.optim.Adam(model.parameters(), lr=lr)
for _ in range(steps):
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad(); loss.backward(); opt.step()
return float(loss) # should become < 0.01; otherwise the fault is in the code
# 2. Is the starting loss reasonable?
def check_start(model, x, y, n_classes):
with torch.no_grad():
f = float(nn.functional.cross_entropy(model(x), y))
expected = math.log(n_classes)
print(f"start {f:.3f}, expected ~{expected:.3f}")
if abs(f - expected) > 0.5:
print(" ⚠ check the initialisation and the labels")
# 3. Log what reveals problems
def diagnostics(model, lr):
norms = {n: float(p.grad.norm()) for n, p in model.named_parameters()
if p.grad is not None}
total = sum(v**2 for v in norms.values()) ** 0.5
ratio = {n: lr * float(p.grad.norm()) / (float(p.norm()) + 1e-12)
for n, p in model.named_parameters() if p.grad is not None}
return {"grad_norm": round(total, 4),
"smallest_layer": min(norms, key=norms.get),
"update_per_weight": round(sum(ratio.values()) / len(ratio), 6)}
# 4. The learning rate sweep
def lr_sweep(model, dataloader, lr_min=1e-7, lr_max=1.0, steps=100):
opt = torch.optim.Adam(model.parameters(), lr=lr_min)
factor = (lr_max / lr_min) ** (1 / steps)
lrs, losses = [], []
it = iter(dataloader)
for i in range(steps):
try:
x, y = next(it)
except StopIteration:
it = iter(dataloader); x, y = next(it)
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad(); loss.backward(); opt.step()
lrs.append(opt.param_groups[0]["lr"]); losses.append(float(loss))
if float(loss) > 4 * losses[0]: # derailed, abort
break
for g in opt.param_groups:
g["lr"] *= factor
best = lrs[int(min(range(len(losses)), key=lambda i: losses[i]))]
return {"lr_at_minimum": best, "recommended": best / 10}
recommended = best / 10 is the rule of thumb from Smith's LR range test: at the lr where the loss is lowest the training is already on its way to becoming unstable, so you take an order of magnitude below it.
Mastery means
- Diagnoses from the loss curve
- Logs and interprets gradient norms
- Runs a learning rate sweep
Sign in to do the exercises and build your mastery up.
Sources
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- arXiv — Cyclical Learning Rates for Training Neural Networks — arXiv (open access; licence per article)
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0