← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 5 · Chapters 7–10

Multi-head Latent Attention

GQA shrinks the cache by sharing heads. MLA shrinks it by changing what you store. Cache a single low-rank latent vector per token instead of keys and values, and reconstruct everything else from it — or better, never reconstruct it at all, because the reconstruction matrices can be folded into the weights beside them.

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:

$$\mathbf c_t=W^{DKV}\mathbf h_t\in\mathbb R^{d_c},$$

and reconstruct per-head keys and values from it with two up-projections:

$$\mathbf k_t=W^{UK}\mathbf c_t,\qquad \mathbf v_t=W^{UV}\mathbf c_t.$$
Only $\mathbf c_t$ is cached. Keys and values become derived quantities. Where MHA stores $2h\,d_h$ numbers per token per layer, MLA stores $d_c$. With $h=128$, $d_h=128$ and $d_c=512$ that is 32768 versus 512 — a 64-fold reduction before we have done anything clever.

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$:

$$\text{score}_{ij}=\mathbf q_{t,i}^\top\mathbf k_{j,i}=\big(W^{UQ}_i\mathbf h_t\big)^\top\big(W^{UK}_i\mathbf c_j\big).$$

Matrix multiplication is associative, so regroup:

$$\boxed{\mathbf q_{t,i}^\top W^{UK}_i\mathbf c_j=\Big(\underbrace{W^{UK\top}_i W^{UQ}_i}_{\text{fold once, offline}}\mathbf h_t\Big)^{\!\top}\mathbf c_j.}$$
Weight absorption. $W^{UK}$ never has to be applied to the cache. Pre-multiply it into the query projection once, and the query is produced directly in latent space. Attention then becomes a dot product between an absorbed query and the raw cached latent. Keys are never materialized at all.

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:

$$\text{out}=\big(W^OW^{UV}\big)\sum_j a_j\,\mathbf c_j.$$
Naive MLAcache $\mathbf c$, expand to $\mathbf k,\mathbf v$ every step
→
Absorbed MLAfold $W^{UK}$ into $W^{UQ}$, $W^{UV}$ into $W^O$
→
Resultattend directly against $\mathbf c$
Absorption is an inference-time reassociation, not a different model. The mathematics is identical; only the order of multiplication changes. During training the unabsorbed form is usually preferred, because keeping the factors separate is what makes the projections low-rank and cheap in the first place. Absorbing produces a larger dense matrix that is only worth forming once, ahead of serving.

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:

$$\text{score}_{ij}=\big(\mathbf R_t\mathbf q_{t,i}\big)^\top\big(\mathbf R_j\mathbf k_{j,i}\big)=\mathbf q_{t,i}^\top\,\mathbf R_{j-t}\,W^{UK}_i\mathbf c_j.$$

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.

The conflict, stated plainly. Absorption needs a position-independent matrix sitting between $\mathbf q$ and $\mathbf c$. RoPE deliberately puts a position-dependent one there. You can have low-rank caching or you can have rotary positions applied to the compressed keys, but not both.

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:

PieceWidthSourceRoPE?Cached as
Content (“nope”)$d_h$up-projected from the latent $\mathbf c_t$nothe 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:

$$\text{score}_{ij}=\underbrace{\big(W^{UK\top}_iW^{UQ}_i\mathbf h_t\big)^{\!\top}\mathbf c_j}_{\text{content: absorbed, no RoPE}}\;+\;\underbrace{\big(\mathbf R_t\mathbf q^R_{t,i}\big)^{\!\top}\big(\mathbf R_j\mathbf k^R_j\big)}_{\text{positional: RoPE, uncompressed}}.$$
Each half keeps what it needs. The content half never meets a rotation, so absorption is valid. The positional half is never compressed, so RoPE applies normally. And because $\mathbf k^R$ is shared across all heads — exactly the MQA idea, applied to a $d_h^R=64$ slice rather than a full head — the extra cache is negligible.

The per-token, per-layer cache is therefore

$$\boxed{d_c+d_h^R}$$

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.

MLA numbers / token / layer
smaller than MHA
equivalent GQA groups
cache at this context
128 heads of width 128, 61 layers, BF16, one sequence. “Equivalent GQA groups” solves $2gd_h=d_c+d_h^R$ for $g$: the number of grouped-query groups that would cost the same cache.

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

SchemeCached per token per layerNumbersQuality
MHA$2h\,d_h$32768baseline
GQA, $h_{kv}=8$$2h_{kv}d_h$2048slightly below MHA
MQA$2d_h$256noticeably below MHA
MLA$d_c+d_h^R$576reported 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.

MLA is not free. The absorbed matrices are larger and denser than the factors they replace, so parameter count and per-step FLOPs rise. That is a deliberate trade: decoding is bandwidth-bound, so spending FLOPs to save bytes moves you along the roofline in the direction you want. On a compute-bound workload the same trade would be a loss.
Takeaway

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