Overfitting and regularization
Spot overfitting and fight it with more data, weight decay, dropout, augmentation and early stopping.
- Recognize overfitting from training and validation loss
- Apply weight decay, dropout and data augmentation
- Choose when to stop training with early stopping
Deep networks have so many parameters that they can memorize their training set - noise, typos and all. That’s overfitting: the training loss keeps dropping while the loss on held-out validation data starts rising. The network is getting better at its homework and worse at the exam.
Always split your data three ways: train (to learn from), validation (to make decisions like when to stop), and test (touched once, at the very end, for an honest final score).
Try it
Watch a model overfit
Raise the model’s capacity (here, a polynomial’s degree) and compare the training error with the validation error. Find where validation error is lowest - beyond that, extra capacity just memorizes noise.
Training error
0.378
Validation error
0.277
Underfitting: too simple to follow the curve.
The regularization toolbox:
| Technique | Idea |
|---|---|
| More data | the best regularizer of all - real or synthetic |
| Data augmentation | make new training examples by flipping, cropping, rotating, adding noise; the label stays the same |
| Weight decay (L2) | add to the loss, so the gradient gains and weights stay small |
| Dropout | during training, randomly zero a fraction p of activations, so no neuron can rely on any other |
| Early stopping | keep the weights from the epoch with the best validation loss |
| Smaller model | fewer parameters, less room to memorize |
Dropout, carefully
Dropout (2014) is wonderfully simple, with one subtlety. If you zero half the activations during training, the next layer sees inputs half as big as it will at test time, when nothing is dropped. Inverted dropout fixes that by scaling the survivors up by during training, so at test time you just do nothing:
1import numpy as np
2
3def dropout(x, p, training, rng):
4 if not training or p == 0:
5 return x
6 keep = rng.random(x.shape) >= p
7 return x * keep / (1 - p)
8
9rng = np.random.default_rng(0)
10activations = np.ones((1000, 100))
11out = dropout(activations, 0.5, training=True, rng=rng)
12print(f"zeroed: {np.mean(out == 0):.3f} mean: {out.mean():.3f}")
13print(np.array_equal(dropout(activations, 0.5, training=False, rng=rng), activations))zeroed: 0.501 mean: 0.998 True
Key takeaways
Overfitting: training loss falls while validation loss rises. Keep train, validation and test sets separate.
More data and augmentation help most; weight decay keeps weights small.
Inverted dropout zeroes activations with probability p and scales the rest by 1/(1−p), only in training.
Early stopping keeps the weights from the best validation epoch.
Lesson quiz
7 questions · pass with 5 correct · up to 50 XP
Passing this quiz completes the lesson and keeps your streak going. Questions you miss come back in review sessions later.
Practice: write Python
Write Python in the editor and run it against sample inputs. Python runs locally in your browser using a WebAssembly runtime.
Inverted dropout
The input is p rows cols. Implement dropout(x, p, training, rng) (inverted dropout) and apply it to an all-ones array of that shape with rng = np.random.default_rng(42). Print the fraction of zeros and the mean to 3 decimals for training mode, then whether evaluation mode returns the input unchanged: train: zeros=0.299 mean=1.001 / eval unchanged: True. With p = 0, nothing may change even in training mode.
- p = 0.3
- p = 0
Python runs in a sandboxed browser worker with a 60 second time limit. Its runtime loads from the Pyodide CDN; your code stays in this browser.
Early stopping with patience
The input has the patience on the first line and the validation loss after each epoch on the second. Training stops once the validation loss hasn’t improved on its best for patience epochs in a row. Print the epoch where training stops and the best epoch (whose weights you’d keep), counting epochs from 1, and the best loss: stopped after epoch 7, best epoch 4 (loss 0.410). If it never stops, print ran all N epochs, best epoch ....
- Overfits after epoch 4
- Keeps improving
Python runs in a sandboxed browser worker with a 60 second time limit. Its runtime loads from the Pyodide CDN; your code stays in this browser.
Questions about this lesson
Stuck? Ask. Figured something out? Share it. Explaining is one of the best ways to learn.
Loading posts…