Optimization & Training Dynamics
Why plain gradient descent isn't what actually trains modern models — Adam's mechanics, why warmup exists, how learning-rate schedules are shaped, and what batch size actually trades off at scale
ML Fundamentals covered gradient descent as the core idea: compute a gradient, step against it, repeat. That's correct, and it's also not what actually runs when a real model trains. Every practical training run wraps that core idea in a stack of decisions — which optimizer, how the learning rate changes over time, how large a batch — each of which has its own failure modes and its own reasons for existing. This lesson is that stack, one layer at a time.
Adam: momentum and per-parameter step size, combined
Plain SGDGradient DescentGradient descent is the optimization algorithm that trains models by repeatedly stepping parameters in the opposite direction of the loss function's gradient. takes the same step size, in the raw gradient direction, for every parameter. AdamGradient DescentGradient descent is the optimization algorithm that trains models by repeatedly stepping parameters in the opposite direction of the loss function's gradient. changes both parts, by tracking two running averages of the gradient for every parameter individually:
- A running average of the gradient itself (momentum, the first moment) — this damps oscillation: a parameter whose gradient keeps flipping sign gets its updates partially cancelled out, while a parameter with a gradient that's consistently pointing one way keeps accelerating in that direction.
- A running average of the squared gradient (the second moment) — this estimates how large and noisy a parameter's gradients typically are, and Adam divides the step by its square root. A parameter with small, consistent gradients gets a relatively larger effective step; a parameter with huge, erratic gradients gets damped.
Adam maintains two exponential moving averages per parameter, with decay rates (typically 0.9) and (typically 0.999):
Both start at zero, which biases early estimates toward zero too — so Adam bias-corrects them before using them:
The parameter update is then:
The bias correction matters most in exactly the early steps where training is already fragile — without it, starts near zero, is tiny, and the first few updates would be huge.
The practical upshot: Adam converges faster and more reliably than plain SGD on the loss landscapes deep networks actually have, at the cost of tracking two extra numbers per parameter — for a billion-parameter model, that's two billion extra floats of optimizer state, on top of the model itself.
| Optimizer | Tracks per parameter | Effective step size | Typical use |
|---|---|---|---|
| SGD | Nothing extra | Fixed, set by the learning rate alone | Simple or small models; a well-understood baseline |
| SGD + Momentum | Running average of the gradient | Same for every parameter, oscillation damped | Classic computer vision training recipes |
| Adam | 1st and 2nd moment of the gradient | Adapted per parameter | Fast, reliable convergence out of the box |
| AdamW | Same as Adam, with decoupled weight decay | Adapted per parameter | Current default for training LLMsLLM (Large Language Model)An LLM is a large transformer trained to predict the next token on massive text corpora, then fine-tuned to follow instructions — the architecture behind GPT, Claude, Gemini, and Llama. from scratch |
Why warmup exists
Cold-starting training at the full target learning rate is one of the most common ways a run destabilizes in the first few hundred steps — and the reason is specific to what's happening in those first steps: a randomly initialized model's early gradients are large and poorly informative about the actual loss landscape, and Adam's own second-moment estimate is itself unreliable before it's seen enough gradients to average over. Warmup starts the learning rate near zero and ramps it up linearly over the first few hundred to few thousand steps, giving both the model's weights and Adam's internal running averages time to settle into a reasonable regime before the full step size is applied.
A transformerTransformerThe transformer is the neural network architecture built around self-attention, introduced in 2017, underlying essentially all modern LLMs. trained from scratch with warmup skipped and the target
learning rate applied from step one will typically show a specific
signature: loss drops normally for the first handful of steps, then
spikes sharply — sometimes to NaN — within the first few hundred. The
model wasn't diverging on the actual task; it diverged because the very
first updates, taken before the optimizer's internal state had any
signal to work with, were simply too large for weights that were still
close to their random initialization. Warmup is cheap insurance against
exactly that failure mode, which is why it's standard on nearly every
transformer training recipe rather than an optional tweak.
Shaping the rest of the schedule
Warmup only covers the first small slice of training. What happens after it is the learning rate schedule, and the standard shape is warmup, then a long decay for the rest of the run:
- Cosine decay — the learning rate follows a cosine curve down to (near) zero by the end of training. This is the most common choice for large transformer runs: it decays slowly at first, drops faster through the middle, then flattens out again near the end, spending more time than linear decay in the range that tends to work well.
- Linear decay — a straight line down to zero. Simpler, and common in fine-tuningFine-tuningFine-tuning continues training a pretrained model on a smaller, curated dataset to teach it a specific behavior, such as following instructions. runs that are short enough that the exact decay shape matters less.
- Constant with warmup — the rate holds flat after warmup, with no decay at all. Rare for training from scratch, since a loss that stops decreasing on a constant rate usually means the rate is now too large for the current region of the loss surface — but sometimes deliberately used for short fine-tuning runs.
None of these shapes is "correct" in isolation — they're all ways of answering the same question (how much should the step size shrink, and on what timeline) with different tradeoffs between simplicity and how precisely the decay is tuned to the run's length.
What batch size actually trades off
ML Fundamentals already covered the headline tradeoff: smaller batches are noisier but buy more steps per unit of compute. Two things that follow from it are easy to get wrong in practice:
- Doubling the batch size does not mean you can just double the learning rate and get the same result for free. The common heuristic — the linear scaling rule — says that a larger batch's lower-variance gradient estimate can tolerate a proportionally larger step, and this holds up reasonably well in practice up to a point. Past that point (the critical batch size, which varies by model and dataset), the gradient estimate is already so close to the true full-dataset gradient that adding more examples per step barely changes it — you're spending compute on redundant samples within the same step rather than getting a meaningfully better step, and large batches beyond this point tend to hurt final model quality even when training remains numerically stable.
- Batch size is often a memory constraint before it's an optimization choice. A batch that would be optimal for training dynamics often doesn't fit in GPUGPU (Graphics Processing Unit)GPUs, originally built for rendering graphics, turned out to be extremely well-suited to the parallel matrix multiplications deep learning requires. memory. Gradient accumulation — running several smaller forward/backward passes, summing their gradients, and only then taking one optimizer step — simulates a larger batch size without needing to hold it all in memory at once, at the cost of that many times more forward/backward passes for the same number of optimizer steps.
Reading a training run
Put together, most of what looks like "training instability" traces back to one of these three knobs:
| Symptom | Usual cause |
|---|---|
Loss spikes early, sometimes to NaN | Missing or too-short warmup |
| Loss plateaus, then diverges later | Learning rate too high for the current region of the loss surface — schedule needs more decay |
| Loss decreases but final model underperforms | Batch size past the critical point, or a schedule that decayed too early relative to total training length |
None of these three — the optimizer, the schedule, the batch size — is tuned in isolation. They interact: a larger batch generally wants a larger (but not proportionally unlimited) learning rate under the linear scaling rule, a longer warmup tends to matter more the larger that peak rate is, and the whole schedule has to be shaped for the total number of steps the run will actually take. Getting this stack right is most of what separates a training run that works from one that doesn't, on exactly the same model architecture and the same data.