Chapter 25
Multi-Head Latent Attention
Low-rank KV compression, weight absorption, decoupled RoPE, and sparse attention.
Multi-head latent attention, introduced in DeepSeek-V2, attacks the same bottleneck as GQA from a different direction: store a compressed memory and reconstruct keys and values only when they are needed. Instead of caching every key and every value head, MLA caches one latent vector per token. That makes long-context decoding more memory efficient while preserving head-specific queries and reconstructed heads.
25.1 Latent keys and values
In ordinary attention, each token state is projected directly to per-head keys and values. MLA first compresses the state to a latent vector of width :
The cache stores . When attention needs actual keys and values, it up-projects the latent:
The model still has multiple query heads, and the up-projections can create different key and value heads from the same cached latent. The cache, however, grows with instead of with or . The trade is compute for memory: each decode step rebuilds just enough key and value information from the compressed cache.
It helps to separate representation from storage. The attention layer can still behave as if it has head-specific keys and values after up-projection. What changes is the persistent record kept between decode steps: a compact latent replaces the full set of KV head vectors. If memory bandwidth is the limiter, reading the compact record and doing a little extra arithmetic can be cheaper than streaming a much wider cache from memory.
def compress_kv(x, w_dkv):
"""Compress token states to one latent KV vector per token."""
return x @ w_dkv
def up_project_latents(c, w_uk, w_uv):
"""Materialize per-head keys and values from cached latents."""
k = np.einsum("tc,chd->thd", c, w_uk)
v = np.einsum("tc,chd->thd", c, w_uv)
return k, v
25.2 Weight absorption
The key up-projection can be moved from the cached side to the query side. For one head, with , the score satisfies
This is only associativity of matrix multiplication. Instead of materializing every , compute an absorbed query in latent space, and then dot it with cached . Values still need up-projection to form the weighted sum, but the score matrix can be produced without storing reconstructed keys. The tests compare the materialized and absorbed score tensors exactly for random small arrays.
Absorption is an inference trick, not a new model. The learned weights are the same, and the resulting score for every query, key position, and head is the same real number up to floating-point rounding. It is useful because score computation touches every cached position, so avoiding materialized keys removes a large repeated read. The operation is easiest to see one head at a time, but the code performs it for all heads by carrying the head axis through the einsums.
def materialized_scores(q, c, w_uk):
"""Scores from explicitly materialized keys."""
k, _ = up_project_latents(c, w_uk, w_uk)
return np.einsum("thd,shd->hts", q, k)
def absorbed_scores(q, c, w_uk):
"""Scores after absorbing the key up-projection into queries."""
q_latent = np.einsum("thd,chd->thc", q, w_uk)
return np.einsum("thc,sc->hts", q_latent, c)
25.3 Why RoPE breaks simple absorption
RoPE is position-dependent: a score uses rotated vectors [su2021roformer]. Substituting the latent key gives a factor between and . Because that factor depends on the key position , it cannot be absorbed once into the query independent of which cached token is being scored.
DeepSeek-V2 handles this by decoupling the positional part: the content key can use absorbed MLA, while a smaller RoPE key is kept as a separate positional channel [shao2024deepseekv2]. The attention score is the sum of a latent-content score and a RoPE score. This preserves the useful relative-position signal without forcing the whole key cache to be materialized.
That split also clarifies what is being cached. The content memory is compressible because the same latent can feed learned up-projections. The positional memory is different: its rotation is defined by where a token sits in the sequence, so it must remain available in a form that can be combined with the current query position. Decoupling keeps the large content path absorbable and leaves only the positional channel outside that algebraic shortcut.
def rotate_pairs(x, positions, theta=10_000.0):
"""Apply a small RoPE rotation to the last axis, which must be even."""
dim = x.shape[-1]
if dim % 2 != 0:
raise ValueError("RoPE needs an even last dimension")
freqs = theta ** (-np.arange(0, dim, 2) / dim)
angles = positions[:, None] * freqs[None, :]
cos = np.cos(angles)[:, None, :]
sin = np.sin(angles)[:, None, :]
even, odd = x[..., 0::2], x[..., 1::2]
out = np.empty_like(x)
out[..., 0::2] = even * cos - odd * sin
out[..., 1::2] = even * sin + odd * cos
return out
25.4 Cache-size comparison
For the same serving example used in the previous chapter, the code below computes bytes for MHA, GQA, and MLA. MHA stores keys and values for all query heads, GQA stores them for fewer KV heads, and MLA stores only the latent vector. The tested values are:
| Attention | Cached numbers per layer and token | MiB |
|---|---|---|
MHA |
2048 |
|
GQA |
512 |
|
MLA |
128 |
def cache_size_table(layers, tokens, bytes_per_value, heads, kv_heads,
head_dim, latent_dim):
"""Return cache bytes for MHA, GQA, and MLA."""
return {
"MHA": 2 * layers * heads * head_dim * tokens * bytes_per_value,
"GQA": 2 * layers * kv_heads * head_dim * tokens * bytes_per_value,
"MLA": layers * latent_dim * tokens * bytes_per_value,
}
The table is not a universal benchmark; it isolates cache storage. Real implementations also pay for query projections, up-projections, memory layout, and the extra decoupled RoPE key if one is stored. Its purpose is to make the scaling visible: replacing cached values by can be a large win when is smaller.
25.5 Sparse attention as indexed retrieval
A separate line of work reduces the number of keys read, not their width. A cheap indexer can score candidate keys, select a top-k set, and run exact softmax attention only on that subset. This toy function is not a production sparse kernel, but it captures the idea behind DeepSeek Sparse Attention reports: use a cheaper route to decide which expensive dot products to keep [deepseek2025v32exp].
The indexer must be cheaper than the attention it saves, and its mask becomes part of the model’s behavior. If it drops a key, the softmax renormalizes over the survivors, so sparse attention is not the same as dense attention with small weights ignored afterward. When the selected set is all keys, the toy function reduces to dense attention; when it is small, the model is doing retrieval before attention.
def top_k_sparse_attention(q, k, v, index_scores, k_top):
"""Attend only to keys selected by a cheap top-k indexer."""
selected = np.argsort(index_scores, axis=-1)[:, -k_top:]
mask = np.zeros(index_scores.shape, dtype=bool)
rows = np.arange(index_scores.shape[0])[:, None]
mask[rows, selected] = True
scores = q @ k.T / np.sqrt(q.shape[-1])
scores = np.where(mask, scores, -np.inf)
weights = softmax(scores, axis=-1)
return weights @ v, mask, weights
Native Sparse Attention makes sparse patterns trainable and hardware-aligned so the model and kernel agree on which blocks are worth reading [yuan2025native].
|
In practice
|
DeepSeek-V2 presents MLA as a way to cut KV-cache size while retaining multi-head behavior [shao2024deepseekv2]. MLA and GQA are not mutually exclusive ideas: both change the memory layout behind attention, and both must be trained or converted carefully. RoPE is the main wrinkle because positional rotation ties keys to positions before the dot product. Sparse attention is another axis: it reduces how many cached tokens are read rather than how wide each cached record is. |
25.6 Teach it
The one-sentence version. MLA stores a compressed latent memory and reconstructs the keys and values that attention needs.
An analogy. Instead of storing every rendered image, keep the scene file. When a camera asks for a view, render the needed pixels from that compact scene.
At the board.
-
Draw going to through , and cache .
-
Draw two arrows from to and .
-
Move across the dot product to get absorbed queries.
-
Add RoPE rotations and point out the key-position term that blocks simple absorption.
Misconceptions to address.
-
"MLA is approximate attention." The dense version is exact for its learned projections.
-
"Absorption removes values." It removes materialized keys from score computation, not values.
-
"RoPE is just another linear layer." Its matrix changes with position.
Check for understanding. Which cached width determines MLA memory: or ?
25.7 Exercises
Describe what MLA stores during decode and what it reconstructs when a new query attends to the cache. Why does this reduce memory compared with MHA?
Starting from , derive (25.3). Which operation moves from the cached-token side to the query side?
Explain why the absorbed query cannot handle ordinary RoPE by itself. What extra object does the decoupled RoPE design keep?
Use the cache-size function to reproduce Table 25.1. Then modify the toy sparse attention so the top-k set is also causal: a query may not select a future key.
References
-
[deepseek2025v32exp] DeepSeek-AI. DeepSeek-V3.2-Exp. Source code and report, 2025. https://github.com/deepseek-ai/DeepSeek-V3.2-Exp
-
[shao2024deepseekv2] Z. Shao et al. DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model. 2024. arXiv:2405.04434
-
[su2021roformer] J. Su et al. RoFormer: Enhanced Transformer with Rotary Position Embedding. 2021. arXiv:2104.09864
-
[yuan2025native] J. Yuan et al. Native Sparse Attention: Hardware-Aligned and Natively Trainable Sparse Attention. 2025. arXiv:2502.11089