Optimisers: momentum, Adam, scheduling
Be able to explain momentum and Adam, choose a learning rate and a schedule, and debug a training run that does not converge.
Prerequisites
- DGradient descentrequired
- DTrain a neural network in PyTorchrequired
Intuition
SGD takes a step towards minus the gradient. The problem: the gradient jumps between batches, and in a long narrow valley it bounces between the walls instead of going forward.
Momentum adds inertia: the step is a moving average of the earlier gradients. The bounces cancel each other out, the forward direction is reinforced.
Adam adds one more thing: every parameter gets its own step length, based on how large its gradients usually are. Parameters with small gradients get larger steps. That makes Adam robust against badly scaled problems — which is why it is the default choice.
The learning rate is still the most important thing. Too large: the loss becomes NaN or bounces. Too small: nothing happens. Typical values: 3e-4 for transformers, 1e-3 for small networks, 2e-5 for fine-tuning.
Formal
Momentum: , with .
Adam: keeps two moving averages — the first moment (the direction) and the second moment (the size):
,
Bias correction (important at the start, when the averages begin at zero): , .
The update: with , , .
AdamW separates the weight decay from the gradient ( separately) — that is correct regularisation and the standard for transformers.
The schedule: a linear warm-up over a few hundred steps (otherwise the first Adam steps are unstable) followed by a cosine decay towards zero.
Code
import torch
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.01, betas=(0.9, 0.999))
sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=3e-4, total_steps=total_steps, pct_start=0.05)
for xb, yb in loader:
loss = criterion(model(xb), yb)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # against exploding gradients
opt.step(); sched.step(); opt.zero_grad()
A debugging scheme for when the training does not converge:
| The symptom | The likely cause | The measure |
|---|---|---|
| loss = NaN immediately | lr too high, or a division by zero in the data | lower lr 10×, check the data |
| the loss bounces | lr too high, the batch too small | lower lr, increase the batch |
| the loss stands still | lr too low, a dead ReLU, the wrong loss | raise lr, check the norm of the gradients |
| the loss falls but the validation rises | overfitting | regularise, early stopping |
| the loss does not fall even on 10 examples | a bug, not a hyperparameter | deliberately overfit a small batch first |
The last row is the best test there is: a correct implementation should be able to get the loss close to zero on ten examples. If it cannot, the fault is in the code.
Mastery means
- Explains momentum and Adam's two moments
- Chooses a learning rate and a schedule
- Debugs a training run that does not converge
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Adam: A Method for Stochastic Optimization — arXiv (open access; licence per article)
- arXiv — Decoupled Weight Decay Regularization (AdamW) — arXiv (open access; licence per article)
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0