Testing a training loop
A training-loop test runs the whole loop for a few fast steps on fake data, so a broken loss function or a crashed checkpoint is caught in seconds instead of after a multi-hour run.
- 10 min read
- 3 reading levels
- Published
Read these first
On this page 8
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Testing a training loop means running the whole loop for a few fast steps on fake data. This checks it works before you spend hours on a real run.
The analogy you have already lived
You have watched a school play's dress rehearsal, run at full speed, days before opening night. Nobody waits until a paying audience is in the seats to find out an actor forgot an entrance. The whole show runs once, fast, so problems show up while they are cheap to fix.
A training-loop test is that dress rehearsal for your code. It runs the entire loop — load data, forward pass, compute loss, backward pass, save a checkpoint. But it only runs two steps on four fake rows, not two hundred steps on two million real ones.
Why it exists
Unit tests check small functions on their own. A training loop is the thing that wires all of them together — and wiring is exactly where bugs hide.
A real training run can take hours. The checkpoint-saving code might have a typo, or the loss function might get the wrong shape of tensor. You often do not find out until the run finishes, or crashes near the end. Hours of compute time are gone, and you are back to square one having learned nothing except that something, somewhere, was wrong.
A training-loop test exercises the exact same code path, every line the real run would touch. It finishes in seconds because the data and the step count are both tiny.
How it works
real training run training-loop test
------------------ ------------------
millions of real rows 4 fake rows
thousands of steps 2 steps
hours on a GPU seconds on a CPU
------------------ ------------------
same code: load -> forward -> loss -> backward -> save checkpointThe right-hand column runs the identical code. It is not a simplified copy — it is the real loop, given less to do.
A real example you have seen
Every video game has a "quick save" you can reload instantly, instead of restarting the level from zero. A training-loop test is the same idea applied to code: instead of restarting a two-hour run to find one bug, you restart a two-second one.
The honest part
A training-loop test cannot tell you the model will be accurate. It only tells you the machinery does not crash and the loss is a real number that moves in the right direction. That is a smaller promise than it sounds, and it is worth having anyway — most training failures are machinery failures, not accuracy failures.
Remember this
- A training-loop test runs the real loop, on fake, tiny data, for a handful of steps.
- It exists to catch broken machinery in seconds, before it wastes hours of real training time.
- It checks that things run and numbers stay sane — it does not check that the model will be good.
What to learn next
- Behavioural tests for models — checking what a trained model does, not only whether training ran.
- Golden outputs and regression tests — freezing known-good behaviour so a retrain cannot quietly break it.
- Experiment tracking — recording every real run this smoke test is protecting.
Developer — Code and libraries.
Setup
pip install torch pytestThe loop being tested
A small PyTorch model and a training function, written exactly the way you would write it for a real run — nothing simplified for the test.
"""A tiny training loop, small enough to run in full inside a test."""
import torch
from torch import nn
def make_model() -> nn.Module:
return nn.Sequential(nn.Linear(3, 8), nn.ReLU(), nn.Linear(8, 1))
def make_fake_batch(n: int = 16, seed: int = 0):
g = torch.Generator().manual_seed(seed)
X = torch.randn(n, 3, generator=g)
y = (X[:, 0] * 2 - X[:, 1] + 0.5).unsqueeze(1)
return X, y
def train(model, X, y, steps: int, lr: float = 0.05):
"""Runs `steps` gradient updates and returns the loss after each one."""
opt = torch.optim.SGD(model.parameters(), lr=lr)
loss_fn = nn.MSELoss()
history = []
for _ in range(steps):
opt.zero_grad()
pred = model(X)
loss = loss_fn(pred, y)
loss.backward()
opt.step()
history.append(loss.item())
return history
if __name__ == "__main__":
torch.manual_seed(0)
model = make_model()
X, y = make_fake_batch()
history = train(model, X, y, steps=200)
print("first loss:", round(history[0], 4), " last loss:", round(history[-1], 4))python loop.pyfirst loss: 5.1238 last loss: 0.0146
Those two numbers are a real measurement from one run with a fixed seed, so they will reproduce on your machine too — torch.manual_seed(0) is what makes that possible. A 200-step run like this takes a couple of seconds on a laptop CPU. The point of the test below is that you never have to wait even that long.
The test
Four checks. None of them needs 200 steps or a GPU.
import math
import torch
from loop import make_fake_batch, make_model, train
def test_a_full_run_shrinks_the_loss():
torch.manual_seed(0)
model = make_model()
X, y = make_fake_batch()
history = train(model, X, y, steps=5)
assert history[-1] < history[0]
def test_no_step_produces_nan_or_inf():
torch.manual_seed(0)
model = make_model()
X, y = make_fake_batch()
history = train(model, X, y, steps=5)
assert all(math.isfinite(loss) for loss in history)
def test_output_shape_matches_the_target():
torch.manual_seed(0)
model = make_model()
X, y = make_fake_batch(n=4)
pred = model(X)
assert pred.shape == y.shape
def test_checkpoint_round_trip_gives_identical_predictions():
torch.manual_seed(0)
model = make_model()
X, _ = make_fake_batch(n=4)
before = model(X)
torch.save(model.state_dict(), "ckpt.pt")
reloaded = make_model()
reloaded.load_state_dict(torch.load("ckpt.pt", weights_only=True))
after = reloaded(X)
assert torch.allclose(before, after)pytest test_loop.py -q.... [100%] 4 passed in 4.59s
That 4.59s figure is a real measurement on one laptop, and most of it is PyTorch's own import and startup cost, not the training. It will vary by machine — treat the pass count, not the timing, as the thing to trust.
Line-by-line: the checkpoint test
test_checkpoint_round_trip_gives_identical_predictions saves the model's weights, loads them into a fresh model object, and checks both give the same prediction on the same input. This catches a specific, common bug: saving the weights but forgetting to also record something the model needs at load time (like a custom layer's configuration). A model that "loads without error" is not the same as a model that loads correctly — this test checks the second, stronger claim.
Now break something on purpose
Raise the learning rate from 0.05 to 50.0 — a typo that is easy to make and does not raise an exception:
# in loop.py
def train(model, X, y, steps: int, lr: float = 50.0):FF.. [100%]
================================== FAILURES ===================================
______________________ test_a_full_run_shrinks_the_loss _______________________
def test_a_full_run_shrinks_the_loss():
torch.manual_seed(0)
model = make_model()
X, y = make_fake_batch()
history = train(model, X, y, steps=5)
> assert history[-1] < history[0]
E assert nan < 5.123818397521973
test_loop.py:13: AssertionError
______________________ test_no_step_produces_nan_or_inf _______________________
def test_no_step_produces_nan_or_inf():
torch.manual_seed(0)
model = make_model()
X, y = make_fake_batch()
history = train(model, X, y, steps=5)
> assert all(math.isfinite(loss) for loss in history)
E assert False
=========================== short test summary info ===========================
FAILED test_loop.py::test_a_full_run_shrinks_the_loss - assert nan < 5.123818...
FAILED test_loop.py::test_no_step_produces_nan_or_inf - assert False
2 failed, 2 passed in 3.76sA learning rate that is far too high pushes the weights to a huge value on step one, and the loss becomes nan (not-a-number) by step two. This is one of the single most common real training failures, and the test catches it in under four seconds instead of at the end of an overnight run.
Common mistakes
Testing on the real dataset "to be safe". That defeats the purpose — the test becomes slow, and slow tests get skipped. Four fake rows exercise the same code paths as four million real ones.
Only checking that the loop runs without an exception. A loop can run to completion while quietly returning nan every step, exactly as shown above. Always assert the loss is finite and moving in the right direction, not only that no exception was raised.
Forgetting torch.manual_seed. Without it, test_a_full_run_shrinks_the_loss is comparing two runs starting from different random weights, and will occasionally fail for no real reason. A flaky test that passes four times out of five gets ignored, which is worse than no test — always fix the seed for anything used in an assertion.
Testing training but never testing resuming. A separate, equally common bug is a checkpoint that saves correctly but cannot be resumed from — the optimiser state, epoch counter, or learning-rate schedule silently resets. The checkpoint test above only checks the model weights; a full suite should also resume training from the checkpoint and confirm the loss continues falling, not restarts from its initial value.
Try it yourself
Add a fifth test, test_loss_actually_uses_the_labels. Hint: train the same model twice with the labels y and with y shuffled into a random order, and assert the two runs produce different final losses. This catches a real bug — a loss function accidentally computed against the wrong tensor — that none of the four tests above would catch.
What to learn next
- Behavioural tests for models — checking what a trained model does, not only whether training ran.
- Golden outputs and regression tests — freezing known-good behaviour so a retrain cannot quietly break it.
- Experiment tracking — recording every real run this smoke test is protecting.
Researcher — Mathematics and papers.
What this class of test can and cannot prove
A training-loop smoke test is a liveness and machinery check, not a correctness check in the statistical sense. Formally: it establishes that the composed map (data loader, forward pass, loss, backward pass, optimiser step, checkpoint I/O) is well-defined and numerically stable on a small input, for a small number of iterations. It says nothing about generalisation, and nothing about whether the loss function encodes the objective you actually want.
This distinction matters because the two failure modes have very different costs. A machinery failure (nan, shape mismatch, a crash) is caught by a five-second test. A correctness failure (a loss that trains stably toward the wrong objective) requires the behavioural tests and golden-output regression tests covered next in this section.
Numerical stability as a testable property
Divergence to nan or inf is usually explained by one of a small number of mechanisms, each independently testable:
- Learning rate above the stability threshold of the local loss curvature — for a quadratic loss with Hessian eigenvalue $\lambda_{\max}$, gradient descent diverges once $\text{lr} > 2/\lambda_{\max}$.
- Unclipped gradient norms during early training, common in recurrent and deep architectures — mitigated by gradient clipping, itself worth a unit test (
assert grad_norm <= clip_valueafter a step). - Loss-scale mismatch under mixed-precision training, where
float16underflows or overflows before the optimiser step — the reasontorch.cuda.amp.GradScalerexists, and worth a dedicated test on any pipeline that uses it.
Determinism, and its limits
torch.manual_seed controls the CPU random-number generator that make_fake_batch and make_model's initialisation both draw from, which is sufficient for the reproducibility this lesson relies on. Full bit-for-bit determinism, including on GPU, additionally requires torch.use_deterministic_algorithms(True) and disabling non-deterministic cuDNN kernels, at a real throughput cost — appropriate for a CI smoke test, rarely appropriate for a production training run.
Where this sits relative to the literature
Sculley et al. (2015), Hidden Technical Debt in Machine Learning Systems, names exactly this gap: conventional software testing verifies code behaviour, while ML systems additionally accumulate debt from untested "glue code" between pipeline stages — the loading, shaping, and checkpointing logic a training-loop test is designed to exercise. Breck et al.'s ML Test Score (2017) later operationalises this as Infra Test 1: "training is reproducible," which this lesson's four tests jointly approximate for the CPU case.
Papers
- Sculley et al., Hidden Technical Debt in Machine Learning Systems, NeurIPS 2015 — papers.nips.cc/paper/5656
- Breck et al., The ML Test Score, IEEE Big Data 2017
- Micikevicius et al., Mixed Precision Training, ICLR 2018 — arxiv.org/abs/1710.03740
What to learn next
- Behavioural tests for models — checking what a trained model does, not only whether training ran.
- Golden outputs and regression tests — freezing known-good behaviour so a retrain cannot quietly break it.
- Experiment tracking — recording every real run this smoke test is protecting.