Skip to content

Multi-head latent attention (DeepSeek-V2/V3), weight absorption

ModuleL7.6 · build · Python · Pass 5 · 3 to 4 h, plus your graded tests (rung R5)
You buildpython/tinyllm/modern/mla.py: MLAttention, ConcatLatentCache, AbsorbedMLA; and your own oracle tests in python/tests/l7-6-mla/
Contractcourse/contracts/py/tinyllm/modern/mla.pyi
Testscourse/tests/L7.6/test_mla.py (what they check: section 4), golden values from transformers 5.19.0 DeepseekV3Attention in course/fixtures/L7.6/mla_hf.npz (course/oracle/L7.6/mla_hf.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsL7.3 apply_rope, rope_cos_sin, RopeSpec · L7.1 RMSNorm · L0.4 Linear, Module · L0.2 the op library · L0.1 Tensor · M06.3 PCG32 · M05.1 kv_bytes_per_token (the tests compare with it) · reading: M03.5 low rank, L7.5 the cache hook, M09.2 (or --ref-deps)
Used byL7.9 builds it for tl_attention = "mla" · later: L8.2 LatentCache, C1’s MLA-vs-GQA ablation at equal KV bytes
MilestoneMS-L7 (your decoder loads and matches Hugging Face checkpoints)
Optional depthDeepSeek-AI, “DeepSeek-V2” (2024), section 2.1; “DeepSeek-V3 Technical Report” (2024), section 2.1.1
  • MLA caches one latent vector cc of width rr plus one shared rope key of width drd_r per token, whatever the number of heads: r+drr + d_r numbers, M05.1’s kv_bytes_per_token for mla (test_cache_bytes_match_m05_1).
  • Every head’s key and value are linear in cc, so a decode step folds WUKW_{UK} into the query and WUVW_{UV} into o_proj and never rebuilds per-head keys (test_absorbed_decode_equals_naive).
  • RoPE cannot be absorbed (it depends on position), so MLA splits each query and key into a position-free “nope” part and a small rotated part (test_scores_see_relative_position_only).
  • The cached latent is the normalized one, and the layout of every projection is HF’s, so DeepSeek checkpoints load by name (test_mla_golden).
Terminal window
ol start L7.6 # stubs mla.py; prints your test path and rung (R5)
ol tests L7.6 # the course tests
# write your oracle tests in python/tests/l7-6-mla/, then:
ol check L7.6 # course tests and the mutation grade of your tests
ol diff L7.6 # after passing: your code against the reference

Your GQA layer (L7.5) already shrank SmolLM2’s cache by three: 3 kv heads instead of 9. The cache still grows by 2Hkvdh2 H_{kv} d_h numbers per token per layer, and at serving time (L10.2 admits requests by KV bytes) that is what limits how many conversations fit in memory. DeepSeek-V2 cut it much further with a different idea: cache a low-rank summary of the token and rebuild keys and values from it. Built naively that costs compute at every step; built with weight absorption it costs almost nothing. This module builds both forms and proves they agree, so C1 can compare MLA and GQA at equal KV bytes.

SymbolMeaningType / shape
xt∈Rdx_t \in \mathbb{R}^dthe input of token ttfloat32[d]
HHquery headsint
rrkv_lora_rank: width of the latentint
rqr_qq_lora_rank: width of the query bottleneck (None: none)int
dnd_n, drd_r, dvd_vqk_nope_dim, qk_rope_dim, v_dim per headint
ct∈Rrc_t \in \mathbb{R}^rthe latent of token tt: what the cache holdsfloat32[r]
ktR∈Rdrk^R_t \in \mathbb{R}^{d_r}the rotated rope key, shared by every headfloat32[d_r]
WUKh∈Rdn×rW_{UK}^h \in \mathbb{R}^{d_n \times r}, WUVh∈Rdv×rW_{UV}^h \in \mathbb{R}^{d_v \times r}head hh‘s rows of kv_b_projmatrices
WOh∈Rd×dvW_O^h \in \mathbb{R}^{d \times d_v}head hh‘s columns of o_projmatrix
RpR_pthe RoPE rotation at position pp (L7.3)dr×drd_r \times d_r
sssoftmax_scale =(dn+dr)−1/2= (d_n + d_r)^{-1/2}float

M03.5 showed that a matrix of rank rr factors as a product of an m×rm \times r and an r×nr \times n matrix, and that truncating the SVD keeps the best rank-rr approximation. MLA bets that the keys and values of all heads, stacked as one H(dn+dv)H(d_n + d_v)-wide vector per token, live near an rr-dimensional subspace. It learns the factorization directly: a down projection to ctc_t (kv_a_proj_with_mqa, the first rr outputs), an RMSNorm (kv_a_layernorm, L7.1), and an up projection kv_b_proj of shape [H(dn+dv),r][H(d_n + d_v), r]. SmolLM2-sized MHA caches 2Hdh2 H d_h numbers per token; DeepSeek-V3 uses r=512r = 512 and dr=64d_r = 64 for 128 heads.

Steps of the naive form, for each token:

ct=RMSNorm(WDKVxt),ktN,h=WUKhct,vth=WUVhct,c_t = \mathrm{RMSNorm}(W_{DKV} x_t), \qquad k^{N,h}_t = W_{UK}^h c_t, \qquad v^h_t = W_{UV}^h c_t, ktR=Rt WKRxt,qth=[ qtN,h∣Rt qtR,h ],k^R_t = R_t\, W_{KR} x_t, \qquad q^h_t = [\,q^{N,h}_t \mid R_t\, q^{R,h}_t\,], scoreh(t,j)=s (qtN,h⋅kjN,h+(RtqtR,h)⋅kjR).\mathrm{score}^h(t, j) = s\,\big(q^{N,h}_t \cdot k^{N,h}_j + (R_t q^{R,h}_t) \cdot k^R_j\big).

Each head’s query is [nope | rope] with the nope part first, and kv_b_proj splits per head into [k_nope | v], nope first. Positions enter only through RR, so a shift of every position changes nothing (L7.3’s law). The rope key is one vector shared by all heads, like MQA, which keeps it to drd_r cached numbers.

The cache receives ctc_t after the norm and ktRk^R_t after the rotation, r+drr + d_r numbers per token per layer: ConcatLatentCache.update(layer, c, k_rope) appends a chunk and returns everything held. A query of a chunk of TT new tokens sits at key index Tk−T+tT_k - T + t, exactly as in L7.5, so prefilling in chunks equals one full forward.

Rewrite the nope score and the output with matrix products moved around:

qtN,h⋅(WUKhcj)=((WUKh)⊤qtN,h)⋅cj,WOh∑jpjWUVhcj=(WOhWUVh)∑jpjcj.q^{N,h}_t \cdot (W_{UK}^h c_j) = \big((W_{UK}^h)^\top q^{N,h}_t\big) \cdot c_j, \qquad W_O^h \sum_j p_j W_{UV}^h c_j = (W_O^h W_{UV}^h) \sum_j p_j c_j .

The query is mapped into latent space once (qtlat,h=(WUKh)⊤qtN,hq^{lat,h}_t = (W_{UK}^h)^\top q^{N,h}_t, an rr-vector), the scores are dot products with the cached cjc_j, the weighted latent ∑jpjcj\sum_j p_j c_j is formed once per head, and WOhWUVhW_O^h W_{UV}^h (precomputed, d×rd \times r) maps it out. No per-head key or value is ever built: decoding reads r+drr + d_r numbers per cached token. RoPE’s part cannot be absorbed (it depends on t−jt - j), which is the reason for the decoupled rope key.

Absorption is an inference rewrite. Training uses the naive form built from the op library, so gradients reach every projection through the latent, its norm, and the rope key.

d=2d = 2, one head, r=1r = 1, dn=1d_n = 1, dr=2d_r = 2, dv=1d_v = 1, inv_freq =(π/2)= (\pi/2), so position 1 rotates the rope pair by 90 degrees: (a,b)↦(−b,a)(a, b) \mapsto (-b, a). Weights: q=[x0∣x0,x1]q = [x_0 \mid x_0, x_1]; craw=x0+x1c_{raw} = x_0 + x_1 and kR=(x1,x0)k^R = (x_1, x_0); kv_b_proj gives kN=2ck^N = 2c, v=3cv = 3c; o_proj outputs (o,0)(o, 0). The latent norm of a single number is c/∣c∣=±1c / |c| = \pm 1 (the ϵ=10−6\epsilon = 10^{-6} changes the 7th digit).

  • Token 0, x=(1,1)x = (1, 1), position 0: craw=2c_{raw} = 2, so c=1c = 1; kR=(1,1)k^R = (1, 1), not rotated; v=3v = 3. It sees only itself: output (3,0)(3, 0).
  • Token 1, x=(0,−1)x = (0, -1), position 1: craw=−1c_{raw} = -1, so c=−1c = -1, v=−3v = -3; kR=(−1,0)k^R = (-1, 0) rotated to (0,−1)(0, -1). Query qN=0q^N = 0, qR=(0,−1)q^R = (0, -1) rotated to (1,0)(1, 0).
  • Scores of token 1: key 0 is 0⋅2+(1,0)⋅(1,1)=10 \cdot 2 + (1, 0)\cdot(1, 1) = 1; key 1 is 0⋅(−2)+(1,0)⋅(0,−1)=00 \cdot (-2) + (1, 0) \cdot (0, -1) = 0. Scaled by s=3−1/2=0.577350s = 3^{-1/2} = 0.577350: softmax gives p=(0.640457,0.359543)p = (0.640457, 0.359543).
  • Output: 3⋅0.640457−3⋅0.359543=0.8427453 \cdot 0.640457 - 3 \cdot 0.359543 = 0.842745, so (0.842745,0)(0.842745, 0).

Absorbed: WUK=2W_{UK} = 2, so qlat=0q^{lat} = 0 and the scores are the rope terms alone, the same two numbers. The weighted latent is 0.640457−0.359543=0.2809150.640457 - 0.359543 = 0.280915 and WOWUV=1⋅3=3W_O W_{UV} = 1 \cdot 3 = 3: again 0.8427450.842745. The cache holds latents (1,−1)(1, -1) and rope keys (1,1)(1, 1), (0,−1)(0, -1): three numbers per token. This is test_hand_example.

class LatentCacheHook(Protocol):
def update(self, layer: int, c_new: NDArray, k_rope_new: NDArray) -> tuple[NDArray, NDArray]: ...
class ConcatLatentCache: # update, seq_len
class MLAttention(Module):
def __init__(self, d, n_heads, q_lora_rank, kv_lora_rank, qk_nope_dim, qk_rope_dim, v_dim, rope: RopeSpec,
attention_bias=False, softmax_scale=None, rng=None)
def forward(self, x, positions, mask=None, cache=None, layer=0) -> Tensor # [B, T, d] -> [B, T, d]
def absorb_weights(self) -> AbsorbedMLA
class AbsorbedMLA: # w_uk [H, dn, r], w_ov [H, d, r], o_bias
def forward(self, x, positions, cache, layer=0, mask=None) -> NDArray
TestKINDChecksWhy it matters downstream
test_hand_exampleunitsection 3: outputs, weights, cache contents, absorbed formyou and the test agree on every split
test_state_dict_keys_and_shapes_are_hfunitHF’s parameter names, order, shapesDeepSeek checkpoints load by name in L7.9
test_mla_goldengoldenDeepseekV3Attention outputs and cached latents, plain and low-rank queriesthe model in MS-L7 and C1
test_absorbed_decode_equals_naivedifferentialprefill then decode, absorbed vs naive, float64L8.2 decodes with the cheap form
test_cache_chunks_equal_full_forwarddifferentialchunked prefill through the cachechunked prefill in L10.3
test_cache_bytes_match_m05_1propertycached arrays are kv_bytes_per_token for mlaadmission by KV bytes, C1’s equal-bytes ablation
test_scores_see_relative_position_onlypropertyshifting every position changes nothingthe decoupled rope key is RoPE
test_padding_mask_and_fully_masked_rowsboundarypadded keys hidden in both forms; an empty row gives 0batched prompts of unequal length
test_gradcheck_through_the_latentgradcheckfloat64 central differences through the latentC1 trains MLA
test_concat_latent_cacheunitappend, copy, seq_len, chunk checksthe hook L8.2 implements
test_validationboundaryrope width, sizes, scale, positions, mask shapewiring bugs fail loudly

Your oracle is MLA written out in numpy float64 from the equations of section 2 (RoPE as complex multiplication), compared with MLAttention on random weights for both query forms and both layouts; the absorbed form against the naive one while decoding; a padding mask in the absorbed form; a central-difference gradient through kv_a_proj_with_mqa; the cache’s copy; the rope-width check; the state-dict names. Import only tinyllm.modern.mla, tinyllm.modern.rope, and tinyllm.autograd. ol check L7.6 requires 0.80 with every pitfall fault killed.

PitfallSymptomCaught by
1. reading kv_b_proj per head as v first, then k_nopea DeepSeek checkpoint loads and attends to garbagetest_mla_golden, test_absorbed_decode_equals_naive (mutant s01)
2. splitting the query rope-firstscores mix position-free and rotated entriestest_hand_example, test_mla_golden (mutant s02)
3. forgetting to rotate the shared rope keyrelative position is lost; outputs drift with absolute positiontest_scores_see_relative_position_only (mutant s03)
4. caching the latent before kv_a_layernormthe cached keys and values are scaled wrongtest_hand_example, test_mla_golden (mutant s04)
5. scaling scores by dn−1/2d_n^{-1/2} instead of (dn+dr)−1/2(d_n + d_r)^{-1/2}softmax too sharp; HF disagreestest_mla_golden, test_validation (mutant s05)
the rope key in a different layout from the queryrelative scores break only for interleaved checkpointstest_mla_golden (mutant s06)
absorbing with the value rows of kv_b_projthe absorbed decode diverges from the naive onetest_absorbed_decode_equals_naive (mutant s07)
a cache that returns only the new chunkeach decode step forgets the pasttest_cache_chunks_equal_full_forward (mutant s08)
the absorbed form dropping the padding maskpadded batches decode wrong only in productiontest_padding_mask_and_fully_masked_rows (mutant s09)
pairing o_proj columns with the wrong head in WOWUVW_O W_{UV}absorbed outputs wrong for H>1H > 1 onlytest_absorbed_decode_equals_naive (mutant s10)
a detached latentkv_a_proj_with_mqa never learnstest_gradcheck_through_the_latent (mutant s11)
DirectionModuleHow it uses this
BackL7.3rotates the decoupled rope parts with a RopeSpec
BackL7.1RMSNorm on the latent and the query bottleneck
BackL0.4Linear projections, Module registration
BackL0.2the ops that make the naive form trainable
BackL0.1Tensor slicing and products
BackM06.3the default initialization stream
BackM05.1kv_bytes_per_token for mla, checked against the cache
ForwardL7.9tl_attention = "mla" builds this layer in every block
ForwardL8.2LatentCache implements the hook and decodes with AbsorbedMLA
ForwardC1the MLA-vs-GQA ablation at equal KV bytes
Your pieceProduction equivalentWhat it addsWhere to look
MLAttentionHF DeepseekV3AttentionYaRN mscale folded into the softmax scale, interleaved rope weightstransformers/models/deepseek_v3/modeling_deepseek_v3.py
AbsorbedMLAvLLM MLA backend, FlashMLAabsorbed decode kernels over a paged latent cachevLLM vllm/attention/backends/mla/, DeepSeek FlashMLA
ConcatLatentCacheSGLang MLATokenToKVPoolone latent pool per layer, radix-tree prefix sharingSGLang mem_cache/memory_pool.py
the factorizationTransMLA, MHA2MLAconverting a trained GQA model to MLA with an SVD and fine-tuningMeng et al. 2025, Ji et al. 2025