Gradient descent on a flexible model first learns the broad shape of the data and only later the noise. So one way to avoid overfitting is to watch the error on validation data during training and keep the weights from the moment it was lowest.
The model is a polynomial of degree d with weights w = [w0, ..., wd], predicting Σ_j w_j·x^j. Write train_with_watch(train, val, degree, lr, max_steps, patience), where train and val are tuples (xs, ys):
- Start from
w = [0.0] * (degree + 1); this is step 0. - One step replaces every weight at once by
w_j - lr · g_j, whereg_j = (2 / n) · Σ_i (prediction_i - y_i) · x_i^jis the gradient of the mean squared error on thentraining points. - After every step (and at step 0) compute the mean squared error on
val. A step is the new best only if its validation error is strictly smaller than the best so far. - Stop after step
max_steps, or as soon aspatiencesteps have passed since the best step (after stepbest_step + patience), whichever comes first.
Return a tuple (best_step, best_error, best_w, steps_run): the best step, its validation error, a copy of the weights after it, and the number of steps performed.
The setup provides make_wave(n, seed), which returns (xs, ys): n noisy observations of sin(3x) for x between -1 and 1.
Examples
Input: train = ([0, 1], [0, 2]), val = ([2], [3]), degree = 1, lr = 0.25, max_steps = 10, patience = 2
Output: (7, 6.689131259918213e-05, [0.4425048828125, 1.2828369140625], 9)
Explanation: step 1 gives w = [0.5, 0.5] and a validation error of (1.5 - 3)² = 2.25. The
error keeps falling until step 7; steps 8 and 9 are worse, so training stops after step 9.
Training longer would reach the line y = 2x, which predicts 4 at x = 2.
Input: the same, with max_steps = 1
Output: (1, 2.25, [0.5, 0.5], 1)
Constraints
0 <= degree <= 12,1 <= len(train xs), len(val xs) <= 60,xbetween -3 and 31 <= max_steps <= 3000,1 <= patience <= 1000, andlris small enough for the steps to converge- floats are compared with a tolerance of
1e-6
Goals
- Train a polynomial by gradient descent and watch its validation error after every step
- Keep the weights of the best step seen so far, not the last ones
- Stop after a stated number of steps without improvement