Optimisers, Schedulers and the Training Loop
Writing a validation loop that reports the truth
A validation loop needs eval mode, no gradients, and correct averaging across uneven batches — miss any one and it reports numbers that flatter or slander your model.
- 7 min read
- 3 reading levels
- Published
Read these first
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A validation loop measures your model on data it never trained on — and small mistakes in that loop produce confidently wrong numbers.
Here is the classic trap, with school marks. Class A has 40 students averaging 60. Class B has 4 students averaging 90. Averaging the two class averages gives 75 — but the true average of all 44 students is much closer to 63, because Class B is tiny.
Averaging averages of different-sized groups misleads. Validation loops process data in batches, the last batch is usually smaller, and this exact trap sits waiting inside nearly every hand-written loop.
Why it exists
Training loss cannot be trusted as a report card. A model can score perfectly on material it has memorised — like a student re-taking last year's paper. Validation data is the fresh exam.
But the fresh exam only works if graded honestly. Three habits break the grading: leaving the model in training behaviour, measuring while gradient machinery runs, and the batch-averaging trap above.
How it works
for each batch of held-out data:
model in "exam mode" → predictions → add results to running totals
after ALL batches: divide totals once → the honest numberThe safe pattern: accumulate raw totals — points scored, samples seen — across all batches, and divide only once, at the end.
Where you have seen this
Every leaderboard number, every "our model reaches 95%" claim, every A/B test of two models is a validation loop's output. When a published number later proves wrong, the story is often not a bad model but a flattering measurement.
Remember this
- Validate on data the model never trained on, in exam mode.
- Accumulate totals, divide once at the end.
- Batch averages of uneven batches are a trap — weight by batch size.
What to learn next
- TorchMetrics — the accumulate-then-compute pattern, packaged and battle-tested.
- Train and eval mode — what
model.eval()actually switches. - Model evaluation — which metric to report in the first place.
Developer — Code and libraries.
Setup
pip install torchOutputs captured with torch 2.5.1 on CPU; deterministic under the seed shown.
A correct loop, with the trap measured
100 samples in batches of 32 make batches of 32, 32, 32 and 4 — deliberately uneven.
import torch
import torch.nn as nn
from torch.utils.data import DataLoader, TensorDataset
torch.manual_seed(0)
X = torch.randn(100, 8)
y = torch.randint(0, 2, (100,))
loader = DataLoader(TensorDataset(X, y), batch_size=32) # batch sizes: 32, 32, 32, 4
model = nn.Linear(8, 2)
loss_fn = nn.CrossEntropyLoss()
model.eval() # dropout off, batchnorm frozen
naive_sum, loss_sum, correct, n = 0.0, 0.0, 0, 0
with torch.no_grad(): # no graph: honest speed, no memory growth
for xb, yb in loader:
out = model(xb)
batch_loss = loss_fn(out, yb)
naive_sum += batch_loss.item() # wrong: a 4-sample batch counts as much as a 32-sample one
loss_sum += batch_loss.item() * len(xb) # right: weight every batch by its size
correct += (out.argmax(dim=1) == yb).sum().item()
n += len(xb)
print(f"mean of batch means: {naive_sum / len(loader):.4f}")
print(f"true per-sample loss: {loss_sum / n:.4f}")
print(f"accuracy: {correct}/{n} = {correct / n:.1%}")mean of batch means: 0.7542 true per-sample loss: 0.8033 accuracy: 53/100 = 53.0%
The two loss numbers disagree by 6% — from nothing but arithmetic. Four samples got the voting power of thirty-two. With an unlucky last batch the gap grows, and it changes run to run as the data order changes, adding noise exactly where you want stability.
The three guards, one by one
model.eval() — switches layers with two personalities into exam behaviour: dropout stops dropping, batchnorm uses its stored statistics instead of the current batch's. Skipping this makes validation noisy and wrong. Details live in train and eval mode.
torch.no_grad() — tells autograd not to build the graph for anything inside. Validation needs no gradients; without this guard it pays graph-building time and holds activations in memory for a backward pass that never comes.
Totals, then one division — the fix for the class-average trap. Accuracy above never touches batch averages at all: it counts correct predictions and samples, then divides once.
.item() — converts the loss tensor to a plain float on the spot. Accumulating the tensor itself drags its history along and leaks memory — the mechanism is dissected in memory that grows every epoch.
Common mistakes
Forgetting model.train() afterwards. The loop above leaves the model in eval mode. Resume training without switching back and dropout stays dead while batchnorm stops updating — training continues, quietly wrong. Pair every eval() with a train() when the loop returns.
Metrics that cannot be averaged at all. Loss and accuracy accumulate as totals. Precision, recall, F1 and AUC do not: the F1 of the whole set is not any weighted average of per-batch F1s. Accumulate raw counts (or all predictions) and compute once — or use TorchMetrics, which exists for exactly this.
Validating with augmented data. Random crops and flips belong to training. A validation set that changes every epoch measures a moving target; keep its transforms deterministic.
Peeking every few batches with a partial validation. Validating on the first two batches "for speed" reports a different population each time the loader order changes. Validate on the full set, less often, instead.
Try it yourself
Set batch_size=30 (batches of 30, 30, 30, 10) and rerun. Predict first whether the naive number moves closer to the truth or further away, then check. Then add drop_last=True to the loader and explain what the naive method now silently ignores.
What to learn next
- TorchMetrics — the accumulate-then-compute pattern, packaged and battle-tested.
- Train and eval mode — what
model.eval()actually switches. - Model evaluation — which metric to report in the first place.
Researcher — Mathematics and papers.
Decomposable vs non-decomposable metrics
A metric $M$ is decomposable when the dataset value is a function of additive per-batch sufficient statistics. Mean loss qualifies:
$$\bar{\mathcal{L}} = \frac{1}{n} \sum_{b} n_b\, \bar{\mathcal{L}}_b$$
- $n_b$ — batch size; $\bar{\mathcal{L}}_b$ — batch mean loss; $n = \sum_b n_b$.
The naive estimator $\frac{1}{B}\sum_b \bar{\mathcal{L}}_b$ equals $\bar{\mathcal{L}}$ only when all $n_b$ are equal — the demo's discrepancy is exactly the last-batch reweighting.
F1 is a ratio of sums, $F_1 = \frac{2\,TP}{2\,TP + FP + FN}$, so the correct accumulants are the confusion-matrix counts, not per-batch F1 values (a ratio of averages is not an average of ratios). AUC is worse: it is a U-statistic over pairs of samples spanning batches, so it requires retaining all scores or a rank-sketch. This taxonomy is the design basis of TorchMetrics' update/compute split.
The validation number is an estimate
Validation accuracy on $n$ samples has standard error $\sqrt{\hat{p}(1-\hat{p})/n}$ — at $\hat{p}=0.9$, $n=1000$, that is roughly $\pm 1$ point. Differences inside one standard error are noise; treating them as model-selection signal is how spurious "improvements" get shipped. Repeated model selection against one validation set also biases the selected score upward — Cawley and Talbot (2010), On Over-fitting in Model Selection and Subsequent Selection Bias in Performance Evaluation, is the standard reference, and nested cross-validation the standard remedy at small scale.
no_grad vs inference_mode
torch.inference_mode() (PyTorch 1.9+) is a stricter, faster sibling of no_grad: tensors created inside skip version-counter tracking entirely. The cost is that such tensors cannot later join any autograd computation. For a pure validation loop it is a drop-in and marginally faster; no_grad remains the safer default when results feed back into training-adjacent code.
Distributed caveat
Under DDP, each rank sees a shard of the validation set, and batch-mean averaging across ranks reintroduces the uneven-group trap at rank level — DistributedSampler pads the dataset to divide evenly, silently duplicating a few samples. Correct treatment accumulates raw counts per rank and all_reduces the sums, which is what TorchMetrics' dist_reduce_fx machinery automates.
What to learn next
- TorchMetrics — the accumulate-then-compute pattern, packaged and battle-tested.
- Train and eval mode — what
model.eval()actually switches. - Model evaluation — which metric to report in the first place.