Optimisers, Schedulers and the Training Loop
Early stopping and keeping the best model
Watch the validation loss, stop when it has not improved for a while, and rewind to a saved copy of the best weights — a three-part habit that beats many fancy regularisers.
- 7 min read
- 3 reading levels
- Published
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Early stopping ends training when the model stops getting better on unseen data — and restores the best version it ever was.
Think of reducing milk on the stove for kheer. There is a perfect stage: thick, golden, fragrant. Cook past it and the bottom starts to catch — more time on the flame now makes it worse, not better.
A good cook tastes as they go, and takes the pot off at the best taste — not at some fixed clock time.
Why it exists
Trained long enough, most models begin overfitting: memorising their training examples instead of learning patterns that transfer. From that point, training scores keep improving while scores on fresh data quietly decay. If that idea is new, overfitting and underfitting covers it.
The tricky part is that fresh-data scores wobble. One bad reading does not mean the peak has passed. So the rule includes patience: keep cooking until there has been no new best taste for several checks in a row. Then stop, and serve the best batch you saved.
How it works
score on
fresh data ▲ best — a copy of the model is saved here
██ ██ █
██ ███ █
█ █ ██ ███ ← no new best for 5 checks
█ → stop, hand back the saved copy
█
└────────────────────────────── training timeThree parts, all essential: watch fresh-data score, stop after enough patience, return the saved best — not whatever the model became afterwards.
Where you have seen this
Think of any app that "trains on your data": a spam filter adapting to your inbox, a keyboard learning your typing. None of them can afford to train forever, or to ship an overcooked model. Early stopping is the standard guard, precisely because it needs no tuning knowledge from the user.
Remember this
- Longer training stops helping at some point and starts hurting.
- Patience absorbs the wobble: stop only after several checks with no new best.
- Always restore the saved best weights — the final weights are past their peak.
What to learn next
- PyTorch Lightning — this pattern as two ready-made callbacks.
- Overfitting and underfitting — the disease early stopping treats.
- Experiment tracking — keeping the curves that justify your stopping decisions.
Developer — Code and libraries.
Setup
pip install torchOutputs captured with torch 2.5.1 on CPU; the run is seeded and deterministic.
The full pattern in one file
A small noisy training set against a roomy model — overfitting bait, on purpose.
import copy
import torch
import torch.nn as nn
torch.manual_seed(0)
Xtr, Xval = torch.randn(32, 20), torch.randn(200, 20)
w = torch.randn(20, 1)
ytr = Xtr @ w + 1.5 * torch.randn(32, 1) # small, noisy train set: overfitting bait
yval = Xval @ w + 1.5 * torch.randn(200, 1)
model = nn.Sequential(nn.Linear(20, 64), nn.ReLU(), nn.Linear(64, 1))
opt = torch.optim.Adam(model.parameters(), lr=0.01)
loss_fn = nn.MSELoss()
best_val, best_state, patience, bad_epochs = float("inf"), None, 5, 0
for epoch in range(1, 101):
opt.zero_grad()
train_loss = loss_fn(model(Xtr), ytr)
train_loss.backward()
opt.step()
with torch.no_grad():
val_loss = loss_fn(model(Xval), yval).item()
if val_loss < best_val:
best_val, bad_epochs = val_loss, 0
best_state = copy.deepcopy(model.state_dict()) # a snapshot, not a reference
else:
bad_epochs += 1
if bad_epochs == patience:
print(f"epoch {epoch:3d} stopping: no improvement for {patience} epochs")
break
if epoch % 10 == 0:
print(f"epoch {epoch:3d} train {train_loss.item():7.3f} val {val_loss:7.3f}")
model.load_state_dict(best_state) # rewind to the best epoch
with torch.no_grad():
print(f"restored model val loss: {loss_fn(model(Xval), yval).item():.3f} (best was {best_val:.3f})")epoch 10 train 7.262 val 12.250 epoch 20 train 1.371 val 6.248 epoch 30 train 0.323 val 5.472 epoch 35 stopping: no improvement for 5 epochs restored model val loss: 5.472 (best was 5.472)
Train loss is racing toward zero — memorising 32 noisy points — while validation loss bottoms out near epoch 30 and stalls. Five patience epochs later, training ends and the model rewinds to its best self. The budget said 100 epochs; the data said 35.
Walkthrough
copy.deepcopy(model.state_dict()) — the single most important line, and the most commonly botched one. state_dict() returns references to the live weight tensors. Store it without the deepcopy and every later opt.step() mutates your "snapshot" in place: at the end you restore the best epoch's name wearing the final epoch's weights, and the whole exercise silently does nothing.
patience=5 — how many consecutive non-improving checks to tolerate. Small patience quits on noise; large patience wastes compute. Values of 5–20 validation checks are typical; noisier validation deserves more patience.
bad_epochs = 0 on every new best — patience measures a streak. Any improvement resets it.
A refinement worth knowing: min_delta. Counting a 0.00001 improvement as "progress" lets a flatlined run limp on forever. Requiring val_loss < best_val - min_delta (say min_delta=0.01) treats trivial wiggles as no improvement.
Saving to disk instead of memory — for big models, replace the deepcopy with torch.save(model.state_dict(), "best.pt") and restore via torch.load with map_location — the checkpoint habits from moving tensors between devices apply.
Common mistakes
Storing state_dict() without deepcopy. Described above; it is the classic. The symptom is suspicious: "restored" performance exactly equals final performance, every run.
Watching train loss. Train loss almost never rises; a stopper watching it never fires. The signal must come from data the optimiser cannot touch.
Stopping on the very first bad epoch. Patience of zero, in effect. Validation wobbles; the demo's val loss would have stopped this run around epoch 12, well before its true best near epoch 30.
Tuning against the same set you stop on. The stopping epoch was chosen by the validation set, so the reported best-val score is slightly flattering. When the honest number matters, keep a third, untouched test set — the discipline behind train/test splits.
Try it yourself
Set patience=2 and rerun: note the earlier, worse stop. Then patience=20: note the wasted epochs beyond the best. Finally add min_delta=0.05 and watch the stopping epoch shift again. Three knobs, one experiment each.
What to learn next
- PyTorch Lightning — this pattern as two ready-made callbacks.
- Overfitting and underfitting — the disease early stopping treats.
- Experiment tracking — keeping the curves that justify your stopping decisions.
Researcher — Mathematics and papers.
Early stopping as regularisation
For linear models with quadratic loss, gradient flow from $\theta_0 = 0$ yields, in the Hessian eigenbasis with eigenvalues $\lambda_i$:
$$\theta(t)_i = \left(1 - e^{-\eta \lambda_i t}\right) \theta^*_i$$
- $\theta^*$ — the least-squares solution; $\eta$ — learning rate; $t$ — training time.
Stopping at time $T$ filters directions by curvature: high-$\lambda$ (signal-dominated) directions converge first, low-$\lambda$ (noise-dominated) directions remain shrunk — formally analogous to ridge regression with $\lambda_{\text{ridge}} \approx 1/(\eta T)$. Goodfellow, Bengio and Courville (2016), Deep Learning, §7.8, give the classical treatment; the ridge correspondence goes back to Sjöberg and Ljung (1995).
Choosing the criterion
Prechelt (1998), Early Stopping — But When? (in Neural Networks: Tricks of the Trade), catalogues stopping criteria — generalisation loss thresholds, progress-quotient rules, streak rules — and finds slower criteria buy small accuracy gains at large compute cost; the streak ("UP") rule used here is the standard compromise.
When the U-shape lies
Epoch-wise double descent (Nakkiran et al., 2020, Deep Double Descent) documents validation error curves that worsen and then improve again with continued training, especially near the interpolation threshold with label noise. A patience-based stopper can quit inside the temporary bump. Grokking (Power et al., 2022) is the extreme case: validation improvement arriving long after training loss saturates. For most practical fine-tuning these regimes are the exception, but they are the honest caveat against treating the U-shape as a law.
Selection bias of the reported score
The best-of-$T$ validation score is a maximum over correlated random variables, hence upward-biased as an estimate of generalisation; bias grows with checkpoint frequency and validation noise. Cawley and Talbot (2010) quantify the analogous model-selection bias; the operational rule — report from a held-out test set evaluated once — is unchanged.
In frameworks
Lightning's EarlyStopping(monitor=..., patience=..., min_delta=...) and ModelCheckpoint(save_top_k=1, monitor=...) callbacks implement this lesson verbatim, including the restore step — see the next lesson.
What to learn next
- PyTorch Lightning — this pattern as two ready-made callbacks.
- Overfitting and underfitting — the disease early stopping treats.
- Experiment tracking — keeping the curves that justify your stopping decisions.