← All Posts
Deep Learning · Transformers· Attention and memory

Multi-Head Latent Attention: Derive the Compressed Cache

MLA caches a shared latent from which many heads' keys and values are derived. The important trick is to move suitable up-projections into the query and output paths, so inference need not materialize a full historical K/V tensor at every step.

One joint latent, many projected heads

Use column vectors. From residual state $h_t\in\mathbb R^d$, compute a shared content latent $c_t=W_Dh_t\in\mathbb R^{d_c}$. For head $a$, define a content key $k_{t,a}^C=U_a^Kc_t$ and value $v_{t,a}=U_a^Vc_t$.

The same $c_t$ supplies both key and value content across heads. A design with separately cached key and value latents is possible, but it is not the joint-latent cache being counted here. Nor must $d_c$ be smaller than one head width: the comparison is with the aggregate K/V width across all heads.

Move the key projection to the current query

$$q_a^{C\top}k_{t,a}^C=q_a^{C\top}U_a^Kc_t=((U_a^K)^\top q_a^C)^\top c_t.$$

Define an absorbed query $\widetilde q_a=(U_a^K)^\top q_a^C$. Historical content scores can then be computed against the cached $c_t$ directly. This is exact matrix associativity for the specified parameterization, not a post-hoc approximation to arbitrary pretrained keys.

The absorbed query has width $d_c$, which can be larger than one original key head. Saving storage does not automatically reduce every arithmetic dimension. Kernel design and the model's chosen ranks determine the practical tradeoff.

Move value reconstruction after aggregation

Let $p_{a,t}$ be attention weights for head $a$. Since the value up-projection is linear,

$$o_a=\sum_tp_{a,t}U_a^Vc_t=U_a^V\left(\sum_tp_{a,t}c_t\right).$$

Compute the weighted latent $z_a=\sum_tp_{a,t}c_t$ first. If the output projection has head block $W_O^{(a)}$, its contribution is $W_O^{(a)}U_a^Vz_a$. The matrices can be composed where appropriate. Reconstructing every old value vector just to sum them is unnecessary.

Why ordinary RoPE complicates absorption

If each content key were rotated as $R_tU_a^Kc_t$, moving its key projection into the query would introduce a transform that depends on historical position $t$. There would no longer be one absorbed query shared across all historical content latents.

A decoupled construction keeps unrotated content scores and a smaller separate rotary query/key path:

$$s_{a,t}=\frac{q_a^{C\top}k_{t,a}^C+q_a^{R\top}k_t^R}{\sqrt{d_C+d_R}}.$$

The rotary key can be shared across heads and cached separately. The two dot products are added before one softmax. They are not two independently normalized attention outputs that are concatenated afterward. This distinction changes the function.

Count every cached component

For this joint-latent, decoupled-RoPE form, each layer and token caches $d_c+d_R$ scalars: the shared content latent and the positional key. The corresponding conventional MHA cache uses $H(d_k+d_v)$ scalars.

As an illustrative configuration, take $H=32$, conventional $d_k=d_v=128$, $d_c=512$, and $d_R=64$. MHA stores 8192 scalars per token/layer; the latent design stores 576, a ratio of about 14.2. In two-byte storage that is 16,384 bytes versus 1152 bytes. This is a representation-size comparison, not a measured speedup or a promise of identical model quality.

Some architectures use positional conventions or MLA variants that change which components must be cached. Kimi K3, for example, must be read from its own configuration rather than inheriting every DeepSeek-style assumption. Quantization metadata, padding, and allocator overhead also belong in an actual memory measurement.

Test the algebra before optimizing it

# C: [tokens, dc], Uk: [dk, dc], Uv: [dv, dc]
keys = C @ Uk.T
values = C @ Uv.T
scores_expanded = keys @ q
scores_latent = C @ (Uk.T @ q)
assert np.allclose(scores_expanded, scores_latent)
# Given the same normalized attention weights p:
assert np.allclose(p @ values, Uv @ (p @ C))

This tests content-path identities. A positional score, mask, scaling convention, and final output mapping must still be included in a complete attention implementation. The lab keeps this reference separate from the optimized representation.

Try the mechanism Compare expanded and absorbed attention

A decoupled rotary path, when active, contributes to the same score before one softmax. The weighted-sum identity is exact for this parameterization. Enable JavaScript to change the example inputs; the complete calculation remains in the article.

MLA and GQA constrain different parts of the model

GQA ties K/V heads within groups. MLA ties their content to a shared latent and allows learned head-specific up-projections. Both reduce cached storage, but their parameterizations and compute paths differ. A smaller cache alone does not establish which offers the best quality/latency point.

DeepSeek-V2 is the primary reference for this MLA design. The identities and cache example here isolate the mechanism rather than repeating whole-model benchmark comparisons.

Try it: Why is “MLA stores two latents of width dc” the wrong formula for the joint-latent design above?

There is one shared content latent used for both K and V, plus a separate positional key when that path is present. The count is $d_c+d_R$, not automatically $2d_c$.

Mix addressable history with a recurrent state

Continue to hybrid attention for how MLA layers can coexist with KDA layers.