Why Adam beat plain SGD for LLMs
Vision models are mostly trained with SGD + momentum. Transformers are almost always trained with Adam or AdamW. Why did one optimizer win one regime and lose the other?
On this page
The picture version
Six pictures for a reader who has never trained anything. The prose below fills in the seams the pictures skip.
1 · The problem
The recipe that trains your image model stalls on your text model.
2 · Where the mismatch lives
One row of the model sees a word constantly. Another sees it twice.
3 · Why there is no good setting
The two demands don’t overlap anywhere on the dial.
4 · The trick
Stop setting one dial. Give every number its own, and let it set itself.
5 · What it costs
The running record is bigger than the model it’s about.
6 · Keep this card
The whole thing on one index card.
Why it exists
You have a training recipe that works. A small image classifier, SGD with momentum, a learning-rate schedule you copied from a tutorial. The loss goes down, the accuracy goes up. So you point the same recipe at your first small transformer over text — same optimizer, same schedule, different data. The loss drops for a few hundred steps and then crawls. You lower the learning rate and it crawls slower. You raise it and it blows up.
You probably assume that’s a bug in your code, or that the model needs
more data. It’s usually neither. Swap SGD for
Adam,
change nothing else, and the same model trains. That small transformer over
text is the running example for the rest of this post — in particular, one
row of its embedding table: the row for a rare token like aardvark.
The split shows up across the literature too. CNN-era vision papers — ResNets, the original ImageNet results, the workhorse 2010s recipes — are SGD + momentum with a hand-tuned schedule. Transformer papers are Adam or its cousin AdamW: Llama, GPT-style pretraining, fine-tuning, the open-source recipes, all Adam-flavored. Almost nobody trains a serious LLM with plain SGD — recent work has started revisiting that, but it’s the exception, not the default. (Modern ViT-style vision recipes have largely moved to AdamW too, which is a good hint that this is an architecture story more than an images-vs-text one.)
That should feel weird. SGD is the simpler algorithm. Adam stores two extra fp32 buffers per parameter — for a 70B-parameter model, ~560 GB of optimizer state, roughly 2× the size of the fp32 weights themselves. Why is the expensive, more complicated optimizer the one that won the regime where compute is most precious?
Not because Adam is a strictly better optimizer. Adam is better at the specific shape of the loss landscape transformers produce on text — and “one global step size for every parameter” is exactly the assumption that shape punishes.
Why it matters now
Try to full-fine-tune a 7B model on a 24 GB consumer GPU and you meet this directly: the weights are only ~14 GB in bf16, but Adam’s two fp32 moment buffers add ~56 GB on their own. The optimizer, not the model, is what doesn’t fit. That arithmetic is why three separate things exist:
- VRAM
budgets are dominated by optimizer state. During training you store
weights, gradients, and Adam’s two moment buffers
mandv. In fp32 that’s 8 bytes per parameter just for the optimizer state, on top of the 4 bytes for the weights and 4 for the gradients — so the optimizer state alone is about twice the size of the fp32 weights. This is a big part of why memory-efficient optimizers like Adafactor and 8-bit Adam exist (and why systems-level tricks like ZeRO/FSDP shard the state across devices). - Recipes don’t transfer. A schedule that works on a CNN does not trivially port to a transformer. If you don’t know why the field switched, you’ll waste runs trying to make SGD “just work” on a language model and conclude your code is broken.
- The frontier is moving. Recent work — Kunstner et al. (2024) on why Adam wins, and Srećković, Geiping & Orvieto (2025) on when it stops winning — has pushed back on the idea that Adam is necessarily better for LLMs: at small batch sizes, with the right tweaks, plain SGD can keep up. So the “why” isn’t settled folklore; it’s an active research question, and the answer matters for anyone designing the next generation of optimizers.
The short answer
Adam = SGD + per-parameter learning rate that adapts to each parameter's gradient history
Picture to keep: SGD is one master volume knob for the whole orchestra;
Adam gives every player their own knob, set automatically by how loud that
player has been recently — so the quiet oboe (your aardvark embedding row)
stops getting drowned out by the brass. The analogy covers step size only:
unlike a mixing desk, Adam’s knobs move every step and are set by the
gradients themselves, and nothing in it decides which direction a parameter
should move.
SGD uses one global learning rate for every parameter. Adam keeps a running estimate of how big each parameter’s gradient typically is, and shrinks the step for parameters whose gradients are large or noisy while letting parameters with small gradients take bigger steps. That single change — making the effective step size per-parameter and adaptive — is the whole difference. For loss landscapes where different parameters live on wildly different scales (like transformers), it turns out to matter a lot.
How it works
Naive attempt: one step size for everything. Plain gradient descent is one rule:
w ← w − η · g
where g is the gradient of the loss with respect to w and η is the
learning rate. Same η for every parameter in the model — the attention
weights, the LayerNorm scales, and the aardvark embedding row all move by
η times whatever gradient they happened to receive.
Why it breaks. Those parameters do not receive comparably-sized
gradients. Because aardvark is rare, the gradient reaching its embedding
row is tiny compared to the one reaching the. With a single η, you are
forced into a compromise: pick η large enough that the rare rows learn
anything and the frequent, sharply-curved directions overshoot and diverge;
pick η small enough to keep those stable and the rare rows barely move.
That is exactly the “lower it and it crawls, raise it and it blows up”
symptom from the hook. Momentum doesn’t fix it — averaging gradients over
time smooths the trajectory, but a running average of a tiny gradient is
still a tiny gradient.
The fix: normalize each parameter by its own gradient scale. Adam tracks
two extra running averages per parameter — m (the gradient’s running mean)
and v (the gradient’s running mean squared) — and applies:
w ← w − η · m̂ / (√v̂ + ε)
(the hats mean bias-corrected — early in training the running averages start at zero and need rescaling, which is a startup detail, not the idea.)
The two halves do different jobs. m̂ smooths the direction, the way
momentum does. √v̂ in the denominator is the new part: divide each
parameter’s step by the typical size of its own recent gradients, and the
step becomes approximately invariant to that parameter’s gradient scale —
a parameter that always sees small gradients now takes a step comparable to
one that always sees large gradients. The aardvark row stops being drowned
out, and the sharp directions stop dictating the global η. It’s a
per-parameter rescaling that emerges from the training data — no human has
to set it. (“Approximately” is doing real work: ε, the bias correction,
weight decay, and gradient statistics that shift during training all break
exact invariance.)
That sounds like a small detail. For transformers on text, it isn’t — and the “why” splits into three explanations that don’t rule each other out.
Why this matters more for transformers than for CNNs
The remaining question is why the single-η compromise hurts a transformer
so much more than a convnet. Three lines of explanation in the literature
each name a different way that compromise fails; they aren’t mutually
exclusive, and none is fully settled. Treat this section as a working
synthesis, not a verdict.
1. The Hessian is “block heterogeneous.”
The Hessian of a transformer’s loss has very different curvature in different parameter blocks: the attention weights, the MLP weights, the embedding table, and the LayerNorm scales all live on different scales. A single global learning rate is forced to compromise: small enough for the sharpest block, which is wasteful for the flat ones. Adam is coordinate-wise adaptive, but in practice that ends up giving different parameter blocks usefully different effective scales. “Why Transformers Need Adam: A Hessian Perspective” (Zhang et al., NeurIPS 2024, arXiv 2402.16788) argues that block-wise Hessian heterogeneity is a key reason SGD struggles on transformers and that this heterogeneity is much milder in CNNs. They argue a cause; I wouldn’t read it as the settled cause.
2. Token frequencies are heavy-tailed.
Natural-language tokens follow something Zipf-shaped: a small set of very
common tokens (the, ,, ., of) and a long tail of rare ones like
aardvark. If you train with SGD’s single global learning rate, the gradient
signal at the output layer is dominated by frequent tokens. Rare tokens make
tiny contributions to the average gradient, so under plain SGD the loss on
rare-token classes goes down much more slowly than on frequent ones. Adam
divides by the per-parameter √v, which tends to be smaller for coordinates
that repeatedly receive smaller gradients — so rare-token directions get
amplified instead of drowned out. Kunstner et al.’s “Heavy-Tailed Class
Imbalance and Why Adam Outperforms Gradient Descent on Language Models”
(NeurIPS 2024, arXiv 2402.19449) makes this argument with a deliberately
designed empirical study.
3. SGD’s update directions are too sharp.
A complementary line of work looks at the directional sharpness of the update — how much the loss curves along the direction you’re stepping in. Pan & Li’s “Toward Understanding Why Adam Converges Faster Than SGD for Transformers” (arXiv 2306.00204) argues that SGD’s update steps land in much sharper directions than Adam’s on transformers, that this is driven by a few coordinates with poorly behaved curvature, and that Adam’s coordinate-wise scaling is essentially a directional-sharpness reduction. Their coordinate-wise-clipping experiments support the idea that a small set of badly scaled coordinates drives much of the gap — which is a claim about where the gap lives, not a drop-in replacement for Adam.
The seams
Where the textbook story gets less clean:
- Recent results challenge “Adam is necessary.” Srećković, Geiping & Orvieto’s “Is your batch size the problem? Revisiting the Adam–SGD gap in language modeling” (arXiv 2506.12543, 2025) argue the gap shrinks dramatically — in some tuned settings, nearly to nothing — when you use small batch sizes, proper gradient clipping, and momentum. Their reading: a lot of what we attributed to Adam might really be about Adam’s interaction with the large-batch regime that LLM training happens to live in. This is recent work rather than settled folklore, and it has not yet been absorbed into standard practice — but it’s a good reason to treat any one-line explanation skeptically.
- AdamW vs Adam. AdamW (Loshchilov & Hutter, ICLR 2019, arXiv 1711.05101) is just Adam
with weight decay applied to the weights directly, not via the
gradient. With Adam, “L2 regularization” and “weight decay” stop being
equivalent (because the gradient is rescaled by
√v), and AdamW fixes it. Many modern LLM recipes explicitly specify AdamW, so when a post says “trained with Adam,” it’s worth checking which one the config actually names. - Adam costs memory. Two extra fp32 buffers per parameter. For a 100B-parameter model that’s roughly 800 GB of optimizer state, which is why Adafactor and 8-bit Adam got invented and why ZeRO/FSDP-style systems shard the state across devices. The cost is real; the field pays it because in practice plain SGD on a transformer at modern LLM batch sizes converges materially slower for the same compute — but see the seam above; “in practice” is doing work in that sentence.
- It is not a free lunch on generalization. A long-running empirical thread reports that Adam can generalize worse than tuned SGD in some image-classification setups, often described in terms of the sharpness of the minima it finds. There is no one canonical result to pin that on — it’s a broad pattern across many papers rather than a single finding — so treat it as folklore with real evidence behind it, not a theorem. It is a big part of why CNNs stuck with SGD + momentum, and why the transformers-use-Adam pattern was non-obvious in advance.
- Why the field standardized so fast (my read, not consensus). The Adam paper (Kingma & Ba, ICLR 2015, arXiv 1412.6980) shipped well before transformers existed. Vaswani et al.’s “Attention Is All You Need” (2017) used Adam as the default optimizer, and the recipe worked well enough at scale that it became the obvious starting point for everything that followed. Some of what looks like a principled choice today is plausibly path dependence — which is another reason the recent “actually, SGD can work too” papers are worth paying attention to.
You started with Adam = SGD + per-parameter learning rate. What did this
post add to the why? — + a loss landscape whose parameters live on wildly different scales, which is the half that isn’t in the algorithm at all. The
best current working explanation is not that Adam is a better optimizer, but
that transformer training on text produces gradients spread across orders of
magnitude, and a single global step size cannot serve them all at once.
Change the landscape — a CNN on images, or, per the 2025 work, a much
smaller batch size — and the advantage shrinks.
So the answer to the fear you opened with: your code probably wasn’t broken,
and the model probably didn’t need more data. You were asking one number to
be the right step size for a parameter that sees the ten thousand times an
hour and a parameter that sees aardvark twice. It can’t be.
The transfer rule worth keeping: expect Adam-style methods to help whenever gradient scales are badly mismatched across parameters, and expect the gap to narrow when the training regime changes enough that a well-tuned SGD can cope.
Check yourself
Before you go — imagine two blocks of your small transformer whose gradients are consistently 1000× apart in magnitude, step after step. Which optimizer treats them more differently, SGD or Adam?
Answer
SGD. Its update is η · g, so the block with the larger gradients takes
steps 1000× bigger — it either diverges or forces you to drop η for the
whole model, which starves the other block. Adam divides each block’s step by
its own √v̂, which is also about 1000× apart, so the two ratios m̂/√v̂
come out comparable and both blocks move at a similar rate. That approximate
per-coordinate scale-invariance is the mechanism, and it’s the
block-heterogeneity story above stated in one sentence.
And one trade-off: your 7B model trains fine with AdamW but you’re out of VRAM. Someone suggests switching to SGD + momentum to free the optimizer state. What do you actually save, and what do you risk?
Answer
Momentum keeps one buffer per parameter instead of Adam’s two, so it halves the optimizer state — in the fp32 accounting above (4 bytes weights + 4 gradients + 8 state) that’s 4 bytes of 16, about a quarter of the total. Real, but not dramatic. The risk is the whole post: you also give up the per-parameter scaling, and on a transformer that usually shows up as slower convergence per unit of compute, so you may spend more GPU-hours than you saved GPU-bytes. The cheaper trades are the ones designed for exactly this — 8-bit Adam or Adafactor keep the adaptivity and shrink the state, and ZeRO/FSDP shard it across devices instead of dropping it.
Famous related terms
- SGD —
SGD = "step against the gradient" + a single global learning rate. The baseline. Cheap, generalizes well in vision; struggles to learn rare-token directions in language. - Momentum —
momentum = SGD + a running average of past gradients. Helps SGD push through flat regions and dampens noise. Doesn’t, on its own, fix the per-parameter scale problem in transformers. - Adam —
Adam = SGD + momentum on the gradient + momentum on the squared gradient + divide by √(squared gradient). The default LLM optimizer. - AdamW —
AdamW = Adam + weight decay applied to weights, not gradients. Restores the regularization Adam accidentally breaks. What most modern LLM recipes actually use. - KV cache — different memory bottleneck (inference, not training), but the same theme: the cost of the algorithm shows up as VRAM.
- Adafactor / 8-bit Adam —
Adafactor ≈ Adam with the squared-gradient state factorized to save memory. Lives where you can’t afford the two extra fp32 buffers per parameter. - Scaling laws — every optimizer choice is implicitly inside the scaling-law constants. Switching optimizers can change them.
Going deeper
- Adam: A Method for Stochastic Optimization (Kingma & Ba, ICLR 2015,
arXiv 1412.6980) — the primary source, and
the place to go if you want to know exactly what
m,v, and the bias correction do, because the pseudocode is clearer than any prose about it. - Sebastian Ruder, An overview of gradient descent optimization algorithms (ruder.io, also arXiv 1609.04747) — the explainer to read if you want the whole family, from momentum through Adagrad and RMSProp to Adam, and to see that Adam is the end of a chain rather than a bolt from the blue.
- Heavy-Tailed Class Imbalance and Why Adam Outperforms Gradient Descent on Language Models (Kunstner et al., NeurIPS 2024, arXiv 2402.19449) — the rabbit hole, and the best answer to “is the rare-token story actually testable, or just a nice narrative?”