Optimizers: learning rate, momentum, Adam
Your training loss bounces around and won't go down. You halve the learning rate and it decreases smoothly, but painfully slowly. What exactly limits the learning rate, and what do momentum and Adam change?
Everything here happens on a loss surface you can see: a valley, \(L(w) = \tfrac12(\lambda_1 w_1^2 + \lambda_2 w_2^2)\). Near any minimum, a real network's loss looks like this locally, just in millions of directions.
What limits the learning rate?
Start with one weight and \(L = \tfrac12 \lambda w^2\). The gradient is \(\lambda w\), so one step of gradient descent is
Every step multiplies \(w\) by the same factor. \(\lambda\) is the curvature: how sharply the loss bends.
You need \(|1 - \eta\lambda| < 1\). Below \(1/\lambda\), \(w\) shrinks smoothly. At exactly \(1/\lambda\) it lands on the minimum in one step. Between \(1/\lambda\) and \(2/\lambda\) it overshoots, flipping sign each step, but still shrinks. Beyond \(2/\lambda\) every step overshoots by more than it started with, and training diverges. Slide it:
One learning rate, two curvatures
Now two weights with curvatures \(\lambda_1 = 1\) (a flat direction) and \(\lambda_2 = 20\) (a steep one). One learning rate has to serve both.
The steep direction sets the limit: \(\eta < 2/\lambda_{\max} = 0.1\). At that η the flat direction's factor is \(1 - 0.1 \times 1 = 0.9\), so it crawls. The ratio \(\kappa = \lambda_{\max}/\lambda_{\min}\), the condition number, sets how many steps you need: roughly proportional to \(\kappa\) for gradient descent. That's the bouncing-or-crawling dilemma from the opening, in two dimensions.
Descend the valley
Reach the minimum (within 0.05) from the ● start. Tune the knobs each level gives you. The path redraws instantly, so you can feel where the limits are.
A learning rate per parameter
Adam keeps two running averages per parameter: the gradient \(m\) (momentum) and the squared gradient \(v\) (its typical size). Then it divides:
With \(v_0 = 0\) and \(\beta_2 = 0.999\), after one step \(v_1 = 0.001\,g_1^2\), a thousand times too small. Dividing by \(1 - 0.999^1 = 0.001\) fixes it exactly. After correction, the very first step is about \(\eta \cdot \operatorname{sign}(g)\) for every parameter: its size no longer depends on the gradient's scale.
Dividing by \(\sqrt{\hat v}\) gives each parameter its own effective learning rate: large gradients get damped, small ones boosted. It's a diagonal rescaling, one number per coordinate. Whether that helps depends on how the valley is oriented:
● gradient descent (η = 0.095, its best) · ● Adam (η = 0.1). Same curvatures (1 and 20), 100 steps.
On the aligned valley, Adam's per-coordinate scaling matches the geometry. It takes a smooth path with no zig-zag, because each coordinate gets its own step size. On this noiseless toy it still finishes later than well-tuned gradient descent: Adam isn't magic. Rotate the same valley and the steep and flat directions mix across coordinates. The diagonal scaling no longer lines up, and Adam curls around without settling. In networks, parameters often do differ in scale by layer and type (embeddings vs norms vs attention), which is where the diagonal view helps. Adam's other job is normalizing noisy gradients, which a noiseless valley doesn't show.
Apply it somewhere else
Early on, Adam's second-moment estimates are based on very few steps, and the initial loss landscape can be sharp (large curvature). Big steps there blow up. A linear warmup gives the statistics time to settle and the network time to reach flatter regions. Then decay (cosine or linear) so the late phase can settle into a minimum, the noise-floor effect from level 3.
A 4× batch averages away 4× the gradient variance, so you can take proportionally bigger steps for similar noise per epoch. That's the linear scaling rule (Goyal et al., 2017). It breaks down beyond a critical batch size, where curvature, not noise, limits the step (section 2's \(2/\lambda_{\max}\)). For Adam, square-root scaling is often closer. Either way, re-tune.
For plain SGD, L2 in the gradient and weight decay are equivalent. With Adam the L2 gradient goes through the \(1/\sqrt{\hat v}\) rescaling, so regularization strength ends up varying per parameter. AdamW (Loshchilov & Hutter) applies \(w \leftarrow w - \eta\lambda w\) separately. It's the default for transformers.
A spike usually means one step overshot: a large gradient (bad or unusual batch, or a numerically unstable op such as log(0), softmax in fp16, or attention logits growing) times the current learning rate. Mitigations: gradient clipping by global norm, lower peak LR or longer warmup, fp32 for sensitive ops, QK-norm or logit caps, and skipping or inspecting the offending batches. Restart from a checkpoint before the spike.
Write the update rules
Plain editor, numpy, hidden tests. ⌘/Ctrl+Enter runs.
Exercise 1: momentum and Adam steps
Answer out loud, then check
- Stability: above about \(2/\lambda_{\max}\) of the local curvature, training diverges.
- Speed: progress along flat directions scales with η, so too small wastes compute.
- Noise: with stochastic gradients, the stationary noise floor scales with η. So schedules matter: warm up, hold, then decay.
- It interacts with batch size, weight decay and initialization, so re-tune when any of them changes.
- It averages successive gradients. Oscillating components (steep directions) cancel, and consistent components (flat directions) accumulate, up to a \(1/(1-\beta)\) speed-up.
- On quadratics, tuned heavy-ball momentum cuts the steps needed from about \(\kappa\) to about \(\sqrt\kappa\).
- With noisy gradients it also averages noise. The effective learning rate is \(\eta/(1-\beta)\), so re-tune η when changing β.
- Transformers have parameters with very different gradient scales (embeddings, layer norms, attention and MLP weights) and heavy-tailed gradient noise. Per-parameter normalization makes a single learning rate workable.
- For CNNs, well-tuned SGD with momentum often generalizes as well or better, and uses less memory (Adam stores two extra tensors per parameter).
- Memory-lighter variants such as Adafactor factorize the second-moment estimates.
- Try to overfit a single batch. If you can't, suspect a bug (labels, loss, masking, data pipeline) before the optimizer.
- Run a learning-rate sweep or range test. Look at per-layer gradient and update norms (updates should be about 1e-3 of the weights), dead units, and saturated activations.
- Check the schedule and warmup, the initialization scale, and normalization placement. Compare against a known-good baseline configuration.
Sources: B. Polyak, Some methods of speeding up the convergence of iteration methods (1964); Kingma & Ba, Adam (2014); Loshchilov & Hutter, Decoupled Weight Decay Regularization (2017); Goyal et al., Accurate, Large Minibatch SGD (2017).