If you tune only one hyperparameter, tune the learning rate. And once you have a good value, recognise that the best learning rate changes during training: large steps early to make rapid progress and escape poor regions, small steps late to settle into a good minimum. A learning rate schedule encodes this. For large models, a short warm-up at the start is also essential.
Why decay the learning rate?#
With stochastic gradients, SGD with a constant learning rate $\eta$ does not converge to a point; it bounces around the minimum in a region whose size scales with $\eta$ times the gradient noise. Reducing $\eta$ shrinks that region. Empirically, loss curves often show a sharp drop right after each learning-rate decrease.
Common schedules#
Step decay: multiply by $\gamma$ (e.g. 0.1) at fixed epochs (e.g. 30, 60, 90). Classic for ResNets on ImageNet.
Exponential decay: $\eta_t = \eta_0\gamma^t$.
Inverse square root: $\eta_t \propto 1/\sqrt{t}$ โ used in the original Transformer paper after warm-up.
Cosine annealing (Loshchilov & Hutter, 2016):
Smooth, only one real hyperparameter (the total length $T$), and very popular for both vision and language models. Cosine with warm restarts periodically resets to $\eta_{\max}$.
One-cycle policy (Smith, 2018): increase the learning rate linearly from low to a high maximum over the first part of training, then decrease it to a very low value, while momentum moves inversely. It often trains faster ("super-convergence").
Warm-upโstableโdecay (WSD): warm up, hold a constant rate for most of training, then decay quickly at the end. It is convenient for large-model training because you can branch off decayed checkpoints at different points without committing to a total length in advance.
Reduce on plateau: decrease when validation loss stops improving โ adaptive but reactive.
Warm-up#
For the first few hundred or thousand steps, increase the learning rate linearly from near zero to its target:
Why does this help?
- At initialisation, gradients can be large and poorly aligned; big steps can push the network into regions it never recovers from (divergence, dead units, loss spikes).
- Adam's second-moment estimate $\hat{\mathbf{v}}$ is unreliable for the first steps (few samples), so update magnitudes are noisy; warm-up keeps them small until statistics stabilise.
- For transformers with post-layer normalisation, early training is especially unstable; warm-up was critical in the original Transformer.
Typical warm-up: 1โ5% of total steps, or a few thousand steps for large models.
Implementing a schedule#
import math
import torch
model = torch.nn.Linear(10, 1)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
total_steps, warmup = 10_000, 500
def lr_lambda(step):
if step < warmup:
return (step + 1) / warmup # linear warm-up
progress = (step - warmup) / max(1, total_steps - warmup)
return 0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress)) # cosine to 10% of peak
sched = torch.optim.lr_scheduler.LambdaLR(opt, lr_lambda)
for step in range(total_steps):
# ... forward, loss.backward(), opt.step(), opt.zero_grad()
sched.step()
if step in (0, 250, 500, 5000, 9999):
print(step, f"{sched.get_last_lr()[0]:.2e}")Call scheduler.step() in the right place โ per step for step-based schedules, per epoch for epoch-based ones. Mixing these up is a common silent bug.
Finding the learning rate: the range test#
Leslie Smith's LR range test: train for a few hundred steps while increasing the learning rate exponentially from very small (e.g. $10^{-7}$) to large (e.g. 10), recording the loss. Plot loss against learning rate:
- the loss first stays flat (too small), then decreases (good range), then explodes (too large);
- pick a maximum learning rate somewhat below the point where the loss is lowest โ often about a tenth of the value where it starts to blow up.
import numpy as np
def lr_range_test(train_step, lr_min=1e-7, lr_max=10, steps=200):
lrs = np.geomspace(lr_min, lr_max, steps)
losses = []
for lr in lrs:
loss = train_step(lr) # performs one update with this lr, returns loss
losses.append(loss)
if not np.isfinite(loss) or loss > 4 * min(losses):
break
return lrs[:len(losses)], np.array(losses)Batch size and learning rate#
Larger batches give less noisy gradients and allow larger learning rates. The linear scaling rule โ multiply the learning rate by $k$ when the batch is multiplied by $k$, with warm-up โ works well up to a point (Goyal et al. trained ResNet-50 on ImageNet in one hour with this rule). For Adam, a square-root scaling is sometimes used. Beyond the "critical batch size", gains diminish.