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 xs\vx_s is projected directly to per-head keys and values. MLA first compresses the state to a latent vector cs\vc_s of width dcd_c:

cs=xsWDKV.(25.1)\vc_s = \vx_s \mW_{DKV} .\tag{25.1}

The cache stores cs\vc_s. When attention needs actual keys and values, it up-projects the latent:

ks,h=csWUK,h,vs,h=csWUV,h.(25.2)\vk_{s,h} = \vc_s \mW_{UK,h}, \qquad \vv_{s,h} = \vc_s \mW_{UV,h} .\tag{25.2}

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 dcd_c instead of with 2Hdh2H d_h or 2Gdh2G d_h. 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.

Listing 25.1 Compress once, up-project when needed
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 ks=csWUK\vk_s = \vc_s \mW_{UK}, the score satisfies

qt⊤ks=qt⊤WUK⊤cs=(qtWUK⊤)⋅cs.(25.3)\vq_t^\T \vk_s = \vq_t^\T \mW_{UK}^\T \vc_s = (\vq_t \mW_{UK}^\T) \cdot \vc_s .\tag{25.3}

This is only associativity of matrix multiplication. Instead of materializing every ks\vk_s, compute an absorbed query qt′=qtWUK⊤\vq'_t = \vq_t \mW_{UK}^\T in latent space, and then dot it with cached cs\vc_s. 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.

Listing 25.2 Materialized scores equal absorbed scores
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 (Rtqt)⊤(Rsks)(R_t\vq_t)^\T(R_s\vk_s) [su2021roformer]. Substituting the latent key gives a factor Rt⊤RsWUK⊤R_t^\T R_s \mW_{UK}^\T between qt\vq_t and cs\vc_s. Because that factor depends on the key position ss, 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.

Listing 25.3 A tiny RoPE rotation for tests
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:

Table 25.1 Cache size for one sequence in the tested example
Attention Cached numbers per layer and token MiB

MHA

2Hdh2H d_h

2048

GQA

2Gdh2G d_h

512

MLA

dcd_c

128

Listing 25.4 Cache-size rows used by the table
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 2Gdh2Gd_h cached values by dcd_c can be a large win when dcd_c 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.

Listing 25.5 Top-k sparse attention with a cheap indexer
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.

Key equations
cs=xsWDKV\vc_s = \vx_s \mW_{DKV}
ks,h=csWUK,h,vs,h=csWUV,h\vk_{s,h} = \vc_s\mW_{UK,h}, \qquad \vv_{s,h} = \vc_s\mW_{UV,h}
qt⊤ks=(qtWUK⊤)⋅cs\vq_t^\T\vk_s = (\vq_t\mW_{UK}^\T)\cdot\vc_s
MMLA=L T dc bM_{MLA} = L\,T\,d_c\,b

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.

  1. Draw xs\vx_s going to cs\vc_s through WDKV\mW_{DKV}, and cache cs\vc_s.

  2. Draw two arrows from cs\vc_s to ks,h\vk_{s,h} and vs,h\vv_{s,h}.

  3. Move WUK\mW_{UK} across the dot product to get absorbed queries.

  4. 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: HdhH d_h or dcd_c?

25.7 Exercises

Exercise 25.1 ★ What is cached?

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?

Exercise 25.2 ★★ Absorb the key up-projection

Starting from ks=csWUK\vk_s = \vc_s\mW_{UK}, derive (25.3). Which operation moves from the cached-token side to the query side?

Exercise 25.3 ★★ RoPE and position dependence

Explain why the absorbed query cannot handle ordinary RoPE by itself. What extra object does the decoupled RoPE design keep?

Exercise 25.4 ★★★ Compute cache sizes and sparsify

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