Heads up: posts on this site are drafted by Claude and fact-checked by Codex. Both can still get things wrong — read with care and verify anything load-bearing before relying on it.
why → how

Why MLA replaced MHA

DeepSeek-V2 cut its KV cache by 93% by attacking the bottleneck differently than GQA — and in their own matched ablation it scored higher, not lower.

AI & ML intermediate Apr 30, 2026 · updated Aug 25, 2026 · 10 min read

On this page

The picture version

Five pictures for a reader who has never wondered what a chatbot keeps in memory. The prose below fills in the seams the pictures skip.

1 · The problem

The notepad grows with the chat, and eventually it outgrows the model.

what is sitting in GPU memory as your chat gets longer the model’s own weights — a fixed size, forever the notepad — one note per word, and it never stops growing somewhere around here the notepad passes the model — and the notepad, not the model, is what runs you out of memory longer chat →
To stay coherent the model keeps a small note about every word so far, and reads all of them for every new word. Push the conversation far enough and those notes take more room than the model itself — how far depends on the model, so treat this as a regime you can reach rather than a universal fact.

2 · Why the notepad is so fat

Every head files its own full page, for every word, in every layer.

one word of your chat, in one layer …128 …128 keys values 32,768 numbers per word. per layer. multiply by every layer, and by every word in a chat hundreds of messages long, and the notepad eats the machine
Standard attention gives every head its own private key and value for every token. At 128 heads that is tens of thousands of numbers per word, per layer — and all of them get re-read every time the model writes another word.

3 · The idea

File one index card per word. Give each head a stencil that turns it back into a page.

what you store now one card 512 numbers instead of 32,768 stencil every head’s page, rebuilt on demand — never actually stored The obvious catch: rebuilding pages costs arithmetic on every step, and arithmetic is not what you were short of. So the trick is to fold the stencils into the surrounding weights, once, and never rebuild a page at all.
One small shared card per word replaces every head’s private pages, and learned stencils can reconstruct any head’s page from it. Doing that reconstruction at every step would trade the memory saving for a compute bill — which is the failure the next scene fixes.

4 · The trick, and the one thing that won’t fold

Fold the stencils into the neighbours. Everything folds except the position.

what folds the stencil is a fixed matrix so multiply it into the weights next door, once, before you ever serve a request now the model reads the cards directly what doesn’t the position signal is a rotation whose angle depends on where the word is — so it can’t be baked into a fixed weight it needs a lane of its own So the notepad holds two things per word: the card, and a small shared position note. the card position small enough that the 93% saving survives
Because the reconstruction matrices are fixed, they can be pre-multiplied into the query and output weights, so the model attends against the small cards and never decompresses. The position rotation is the one step that refuses to be folded into a fixed weight, which is why it gets an inelegant side-channel of its own.

5 · Keep this card

The whole thing on one index card — appropriately enough.

the trick = store one small card per word + a separate lane for the position signal + fold the rebuilding into the weights next door so the compression costs nothing at answer time the earlier fix shared notes between heads and paid a little quality for the memory; this one compresses instead, and in its authors’ own matched comparison it scored higher rather than lower
Picture to keep: instead of every head filing its own full-page note for each message, the model files one index card per message — and each head has a stencil that turns that card back into its own page when it needs to read. The cards are what you store; the pages are never actually printed.

Why it exists

Imagine a really long ChatGPT conversation — hundreds of messages back and forth. That chat is the running example for this post. To stay coherent, the model has to “remember” every previous word in it. It does that by stashing a small note about each token in a special on-GPU scratchpad, the KV cache. Like a scratchpad, except the notes aren’t summaries anyone could read — each one is a pair of learned vectors, and there’s a separate pair for every attention head in every layer. The longer the chat, the bigger the pad. Push the context far enough and the pad outgrows the model weights — and that, not the model, is what runs out of GPU memory first. MLA is a trick that shrinks each note by roughly 14×, so the same GPU can hold a much longer conversation.

(How far is “far enough” depends on the model: layer count, head count, precision and context length all move the crossover point, so treat “cache bigger than weights” as a regime you can reach, not a universal fact.) Either way, once the cache is the thing filling your VRAM, it’s the thing you have to attack.

Two fixes predate MLA, and DeepSeek-V2’s paper names both: multi-query attention, and GQA, which makes groups of query heads share a single K/V head. GQA is the one that spread. It works, but you can feel the trade — fewer distinct K/V heads means less expressive attention, and teams generally accept a small quality dent to buy the memory back.

DeepSeek-V2 shipped in May 2024 with a different bet: Multi-head Latent Attention. Same goal — shrink the KV cache. Different tactic — don’t share heads, compress what each token stores down to a tiny shared vector and reconstitute the full keys and values at attention time. DeepSeek reports a 93.3% KV-cache reduction versus their dense 67B baseline, and a quality improvement over plain MHA, not a regression (DeepSeek-V2 paper — KV compression is §2.1.2, decoupled RoPE is §2.1.3, the throughput numbers live in §3.2.3, and the “MLA beats MHA” comparison is in Appendix D.2 / Table 9). That combination — cheaper and better — is the read I have for why this architecture spread, though no clean independent ablation isolating MLA from the rest of DeepSeek-V2’s choices has been published.

Why it matters now

If you’ve used DeepSeek-V2, V3, or R1, you’ve used MLA. The concrete present-day stake is batch size: a smaller cache per token means more concurrent conversations fit on the same GPU, which is the main lever on per-token serving cost. DeepSeek attributes its own throughput numbers to MLA together with FP8 deployment, KV-cache quantization and kernel work, so MLA is one contributor to cheap inference rather than the explanation for it. There is also now a small literature on retrofitting MLA onto already-trained MHA/GQA models (e.g. TransMLA, 2025), which is itself a tell — people want this attention shape badly enough to do surgery on existing checkpoints.

GQA is still everywhere, and nothing here says it’s going away. What I’d claim is narrower: if you’re designing new dense attention and the cache is your binding constraint, MLA is the thing you now have to argue against.

The short answer

MLA = MHA + low-rank latent K/V cache + decoupled RoPE channel

Picture to keep: instead of every head filing its own full-page note for each message in your chat, the model files one index card per message — and each head has a stencil that turns that card back into its own page when it needs to read. The cards are what you store; the pages are never actually printed.

Instead of caching every head’s full key and value per token, MLA caches one small “latent” vector per token and learns up-projection matrices that reconstruct each head’s K and V on demand. A separate small vector carries the rotary position signal, because the position rotation is the one step that refuses to be folded into a fixed weight. At inference, an algebra trick folds the up-projections into the query and output weights so the model never actually decompresses — it just attends in the compressed space.

How it works

Plain MHA caches, per token per layer, n_heads × head_dim numbers for K and the same for V. For a model with 128 heads of dimension 128, that’s 32,768 numbers (16,384 for K, 16,384 for V) per token per layer. Multiply by layers and by every message in your long chat, and the cache eats the GPU.

Follow the fix as a chain of breakages.

1. Compress to a latent. A learned matrix W_DKV projects the token’s hidden state down to a small vector c_KV — DeepSeek-V2 sets the compression dimension to 512, versus the 32,768 combined K+V elements you’d otherwise cache in the 128-head example. Only c_KV (plus a small RoPE side-channel, below) lives in the cache. That compression is the core reason MLA shrinks the cache so much.

2. Up-project per head when you need K and V. Two more learned matrices, W_UK and W_UV, take c_KV back up to per-head keys and values at attention time. Why this alone would break: you’ve turned one matmul into two, so you’d be paying extra compute on every attention step to save memory — a bad trade at decode time, where compute isn’t the thing you’re short of.

3. Absorb the up-projections into the query and output weights. This is the fix that makes MLA practical, not just memory-frugal. Because attention scores are Q · K^T and the output is (scores · V) · W_O, you can pre-multiply: fold W_UK into the query projection and W_UV into the output projection (DeepSeek-V2 paper §2.1.2 / Appendix C). At inference, the model never reconstructs full K or V — it computes attention directly against the small latent. You pay compute to train with the up-projections; you skip that compute at inference.

And then a third thing breaks. The seam the paper devotes its own subsection to is RoPE. RoPE multiplies Q and K by a position-dependent rotation matrix before the dot product. If your “K” is really W_UK · c_KV, you’d want to apply the rotation after up-projection — but that breaks the absorption trick, because the rotation depends on position and can’t be folded into a fixed weight. DeepSeek’s fix is decoupled RoPE (§2.1.3): split the key into two parts. The content part lives in the compressed latent and gets no RoPE. The position part is a small, shared-across-heads key vector (per-head dim 64 in DeepSeek-V2) that carries RoPE and is cached alongside c_KV. So the cache holds c_KV plus this small shared RoPE key per token — small enough that the 93% headline survives. It’s structurally inelegant — my read, not the paper’s framing — and it’s the price of keeping the absorption trick alive.

The honest gap: the 5.76× throughput figure is DeepSeek’s own, measured on a single node of 8 H800 GPUs and reported alongside FP8 deployment plus ~6-bit KV-cache quantization (paper §3.2.3), so MLA is doing some but not all of that work, and no independent replication has been published. The 93.3% cache reduction is the more load-bearing claim, and it’s the one stated directly in the paper; throughput in your stack will depend on batch size, context length, and serving software.

You started with MLA = MHA + low-rank latent K/V cache + decoupled RoPE channel. What did this post add that the line hides? — + an algebra trick that makes the compression free at inference. Without absorption, MLA would be a memory win paid for with extra compute on every decode step. With it, the model never decompresses at all — and the ugly-looking decoupled RoPE channel exists purely because rotation is the one operation that refuses to be folded into a fixed weight.

Check yourself

Before you go — a colleague proposes shrinking d_c from 512 to 128 to make the cache four times smaller again. What would you predict, and what would you measure?

Answer

You’d predict a quality cost that grows as d_c falls, because d_c is the bottleneck through which all per-head key and value information for a token has to pass. At 512, DeepSeek reports MLA beating plain MHA; there’s no reason to expect that to survive arbitrary compression, and the paper doesn’t claim it does. What you’d measure isn’t just average loss — it’s tasks that need fine-grained recall of specific earlier tokens, since that’s what a too-small latent should degrade first. Also note what wouldn’t shrink: the decoupled RoPE key is a separate fixed-size channel, so at some point it becomes the floor on your cache size.

And one more — someone benchmarks MLA against GQA on a workload of huge prompts with one-word answers, and finds almost no difference. Does that contradict the 93.3% number?

Answer

No — they measured the wrong regime. The cache-size win pays off when the cache is large relative to everything else and gets re-read many times: long contexts, many concurrent conversations, long generations. A huge-prompt/one-token-answer workload is dominated by prefill, which is compute-bound, and it re-reads the cache almost never. The 93.3% is a claim about how many bytes you store per token, which is real regardless — it shows up as how many of these requests you can run at once, not as latency on any single one.

Going deeper