Course Content
Deep Learning Essentials
13 sections · 61 lessons
Does overfitting happen in neural networks, and how would you treat it?
What you need to know
Diagnose before you treat
Plot training and validation loss per epoch. Three patterns matter:
- Both high and flat: underfitting. The model is too small or training is broken. Regularisation would make it worse.
- Training falls, validation falls then rises: overfitting. The epoch where validation is lowest is where you should have stopped.
- Validation suspiciously close to training, then poor in production: leakage. The validation set shares information with training.
Check for leakage first
Leakage makes overfitting invisible. Typical causes: the same customer, patient or product appearing in both splits; near-duplicate images; random splits of time-series data, which let the model peek at the future; or preprocessing statistics computed on the full dataset. Split by group (patient, user, store) or by time when those exist.
The treatments, in order of impact
- More data or augmentation — the most reliable fix. If labelling more is too expensive, augment realistically.
- Transfer learning — start from pretrained weights so the model needs far fewer examples to learn your task.
- Early stopping — keep the checkpoint with the best validation loss and stop when it has not improved for a few epochs.
- Weight decay — penalise large weights; use AdamW so the decay is applied correctly.
- Dropout — mainly in dense layers and transformer blocks.
- Smaller model — fewer layers or units if the model is much larger than the data can support.
- Label smoothing, MixUp, ensembling — further gains once the basics are in place.
A minimal early-stopping loop:
1best, patience, bad = float("inf"), 5, 02for epoch in range(100):3 train_one_epoch(model, train_loader, optimizer)4 val_loss = evaluate(model, val_loader)5 if val_loss < best:6 best, bad = val_loss, 07 torch.save(model.state_dict(), "best.pt") # keep the best weights8 else:9 bad += 110 if bad >= patience:11 break12model.load_state_dict(torch.load("best.pt"))The loop saves the model whenever validation loss improves and stops after 5 epochs without improvement, then restores the best weights, not the last ones.
A real-life example
A regional hospital builds a chest X-ray triage model from 6,000 images of 2,100 patients. The first run shows 99% training accuracy and 95% validation accuracy, which looks great, but on a new month of X-rays it scores 78%.
The engineer checks the split: images were split randomly, so the same patient's follow-up X-rays appear in both training and validation. The model had learned to recognise patients, not disease. After splitting by patient ID, validation accuracy drops to an honest 81%, and the learning curves now show clear overfitting after epoch 6.
She then switches from training from scratch to a pretrained ResNet-50, adds small rotations and contrast jitter, uses AdamW with weight decay 0.05, and early stopping with patience 5. Validation accuracy rises to 89%, and next month's X-rays score 88%. The validation number now predicts production.
Follow-up questions to expect
- "How do you know it is overfitting and not a distribution shift?" — Overfitting shows as a gap between training and validation from the same source. If validation is fine but production is poor, the production data differs, which is a shift, not overfitting.
- "Can a network overfit even with millions of examples?" — Yes, especially large models trained for many epochs, but it is much less likely. With huge data, the usual problem becomes underfitting or compute.
- "Why might validation loss rise while validation accuracy still improves?" — The model becomes over-confident on the examples it gets wrong, which raises cross-entropy loss even as more predictions become correct. It is an early sign of overfitting.