Early Stopping

Watch the error on held-out data while training, and stop when it starts to rise. The cheapest regularizer there is.

Deep Learning- Fundamentals to Advanced Concepts

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
Training error falling steadily and validation error falling then rising sharply, with a dashed line at the minimum validation error
Training error never stops improving. Validation error bottoms out near epoch 171 and then climbs.

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 PP 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 λ\lambda 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.
Try it yourself
Overfitting lab →

Switch to training over time and find the epoch where the held-out error is lowest.

EasyEarly stopping

What is 'patience' in early stopping?

EasyEarly stopping

Why do we restore the weights from the best epoch rather than keeping the final ones?

MediumEarly stopping

Why is early stopping called a regularizer even though it changes neither the loss nor the model?