Multi-head Latent Attention
Where we left off
The previous note established that decoding is bandwidth-bound and that the KV cache is the one tensor batching cannot amortize. GQA attacks this by reducing $h_{kv}$, which works but costs quality: query heads are forced to share keys and values, so they lose the ability to attend on genuinely different criteria.
MLA asks a different question. Instead of fewer keys and values, can we store something smaller from which all the per-head keys and values can be recovered?
Step 1: a low-rank latent
Introduce a down-projection to a latent of dimension $d_c$, much smaller than the full $h\,d_h$ that keys and values would occupy:
and reconstruct per-head keys and values from it with two up-projections:
Naively this seems to just move the cost: you saved memory but now pay two up-projections on every decode step, for every cached token. That would be a bad trade. The next step is what makes it a good one.
Step 2: absorb the up-projections
Look at what the attention score actually needs. For query head $i$ at position $t$ against cached position $j$:
Matrix multiplication is associative, so regroup:
The same trick works on the output side. The attention output is $W^O\sum_j a_j\mathbf v_j=W^O W^{UV}\sum_j a_j\mathbf c_j$, so $W^{UV}$ folds into $W^O$ and values are never materialized either:
Step 3: RoPE breaks it
Now add rotary embeddings, and the whole construction falls apart. RoPE inserts a position-dependent rotation between the query and the key:
The offending object is $\mathbf R_{j-t}W^{UK}_i$. To absorb $W^{UK}$ into the query we would have to precompute this product — but it depends on the relative position $j-t$, which is different for every query-key pair and changes as generation proceeds.
Three ways out present themselves, and two are bad:
Drop RoPE
Loses the relative-position property that makes long context work. Not acceptable.
Give up absorption
Expand $\mathbf k$ from $\mathbf c$ at every step. Restores the compute cost that compression was supposed to avoid.
Split the head
Carry position in a separate small subspace that is never compressed. This is the answer.
Step 4: decoupled RoPE
Partition each head into two concatenated pieces with different jobs:
| Piece | Width | Source | RoPE? | Cached as |
|---|---|---|---|---|
| Content (“nope”) | $d_h$ | up-projected from the latent $\mathbf c_t$ | no | the latent $\mathbf c_t$, shared by all heads |
| Positional (“rope”) | $d_h^R$ | a separate small projection of $\mathbf h_t$ | yes | $\mathbf k^R_t$, a single vector shared by all heads |
Because the score is a dot product over the concatenated head, it splits cleanly into two independent terms:
The per-token, per-layer cache is therefore
and with DeepSeek's choices $d_c=512$, $d_h^R=64$ that is 576 numbers, against $2\cdot128\cdot128=32768$ for full multi-head attention with the same head count.
A footnote: the query is compressed too
The queries are not cached, so compressing them saves no memory. DeepSeek compresses them anyway, through a rank-1536 bottleneck, for a different reason: it cuts the activation memory held for the backward pass during training, and acts as a mild structural regularizer. It has no effect on inference cost.
The full accounting
| Scheme | Cached per token per layer | Numbers | Quality |
|---|---|---|---|
| MHA | $2h\,d_h$ | 32768 | baseline |
| GQA, $h_{kv}=8$ | $2h_{kv}d_h$ | 2048 | slightly below MHA |
| MQA | $2d_h$ | 256 | noticeably below MHA |
| MLA | $d_c+d_h^R$ | 576 | reported at or above MHA |
The striking part is the last column. MQA and GQA buy their savings by removing capacity — heads genuinely share keys. MLA keeps every head's key and value distinct; it merely observes that they live in a low-dimensional subspace and stores the coordinates instead of the reconstruction. Compression without sharing is why the quality cost does not appear.
Cache a latent instead of keys and values; fold the up-projections into the query and output matrices so neither is ever materialized; and route position through a small uncompressed subspace so the rotation never sits between the query and the cache. Roughly 57 times less cache than MHA, without heads sharing anything.
Check yourself
- Show that $\mathbf q^\top W^{UK}\mathbf c=(W^{UK\top}\mathbf q)^\top\mathbf c$ and say which property of matrix products you used. derivation
- Explain precisely why $\mathbf R_{j-t}W^{UK}$ cannot be precomputed. reasoning
- Compute the equivalent GQA group count for $d_c=512$, $d_h^R=64$, $d_h=128$. calculation
- Why is $\mathbf k^R$ shared across heads rather than given one vector per head? design
- Argue why MLA is a good trade for decoding but a poor one for prefill. analysis
- Implement both the absorbed and unabsorbed forms and assert they agree numerically. implementation