Foundations

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 SGD takes the same step size, in the raw gradient direction, for every parameter. Adam 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.

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.

OptimizerTracks per parameterEffective step sizeTypical use
SGDNothing extraFixed, set by the learning rate aloneSimple or small models; a well-understood baseline
SGD + MomentumRunning average of the gradientSame for every parameter, oscillation dampedClassic computer vision training recipes
Adam1st and 2nd moment of the gradientAdapted per parameterFast, reliable convergence out of the box
AdamWSame as Adam, with decoupled weight decayAdapted per parameterCurrent default for training LLMs 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.

What skipping warmup actually looks like

A transformer 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-tuning 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 GPU 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:

SymptomUsual cause
Loss spikes early, sometimes to NaNMissing or too-short warmup
Loss plateaus, then diverges laterLearning rate too high for the current region of the loss surface — schedule needs more decay
Loss decreases but final model underperformsBatch 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.

On this page