A flexible model trained by gradient descent passes through a good solution on the way to an overfitted one. Early stopping simply stops training at the good point.
The idea
Split off a validation set and never train on it. After every epoch (or every few steps), compute the loss on the validation set as well as on the training set.
- The training loss keeps falling as training goes on.
- The validation loss typically falls at first, reaches a minimum, and then rises as the model starts to fit noise.
The best model is the one at the minimum of the validation curve. We keep a copy of the weights from that moment.
An example
Take the degree-12 polynomial again, now trained by gradient descent from zero weights on 15 noisy points (noise 0.3), with a separate validation set of 300 points:
import numpy as np
from numpy.polynomial import legendre
def true_f(x):
return np.sin(2 * np.pi * x)
def features(x, degree):
return legendre.legvander(2 * x - 1, degree)
r = np.random.default_rng(3)
x = r.uniform(0, 1, 15); y = true_f(x) + r.normal(0, 0.3, 15)
x_val = r.uniform(0, 1, 300); y_val = true_f(x_val) + r.normal(0, 0.3, 300)
A, A_val = features(x, 12), features(x_val, 12)
w, eta = np.zeros(13), 0.05
best_val, best_epoch, wait, patience, stop_epoch = np.inf, 0, 0, 200, None
curve = {}
for epoch in range(1, 40001):
w = w - eta * (2 / 15) * A.T @ (A @ w - y) # one gradient descent step
val = np.mean((A_val @ w - y_val) ** 2)
if epoch in (1, 10, 30, 100, 300, 1000, 3000, 10000, 40000):
curve[epoch] = (np.mean((A @ w - y) ** 2), val)
if val < best_val:
best_val, best_epoch, wait = val, epoch, 0 # new best: remember it, reset the counter
else:
wait += 1
if wait >= patience and stop_epoch is None:
stop_epoch = epoch # early stopping would trigger here
print("lowest validation error", round(best_val, 4), "at epoch", best_epoch)
print("with patience", patience, "training would stop at epoch", stop_epoch)
for e, (t, v) in curve.items():
print(f"epoch {e:6d}: training {t:.3f} validation {v:.3f}")
Output:
lowest validation error 0.1032 at epoch 171
with patience 200 training would stop at epoch 371
epoch 1: training 0.491 validation 0.556
epoch 10: training 0.317 validation 0.406
epoch 30: training 0.155 validation 0.233
epoch 100: training 0.043 validation 0.110
epoch 300: training 0.026 validation 0.106
epoch 1000: training 0.023 validation 0.116
epoch 3000: training 0.022 validation 0.126
epoch 10000: training 0.020 validation 0.210
epoch 40000: training 0.017 validation 1.037
The training error keeps improving right to the end (0.491 down to 0.017). The validation error is best at epoch 171 (0.1032) and is then ten times worse by epoch 40000 (1.037). Training to the end would have destroyed the model.
Patience
The validation curve is noisy in real problems, so stopping at the first uptick would be hasty. The standard rule has a patience parameter: keep training until the validation loss has failed to beat its best value for checks in a row, then stop and restore the weights from the best epoch. In the code, the counter wait resets every time a new best appears and triggers stopping after 200 epochs without one.
Why it works
Gradient descent started from small weights increases the model's effective complexity gradually: early on it fits the big, smooth structure, and only later the fine, noisy details. Stopping early limits how far the weights can travel from their starting point. For a quadratic loss this is closely related to an L2 penalty, which is why early stopping and weight decay often have a similar effect.
Pros and cons
- Free: no extra hyperparameter like beyond the patience, and it saves training time.
- Needs validation data, which is data not used for fitting. (Some schemes retrain on all the data afterwards for the same number of epochs.)
- It mixes two goals. Stopping early may also stop before the training loss is as low as it could be. When it is used together with other regularizers, each has to be tuned with the others in mind.
Switch to training over time and find the epoch where the held-out error is lowest.