Multi-head latent attention (DeepSeek-V2/V3), weight absorption
Overview
Section titled “Overview”| Module | L7.6 · build · Python · Pass 5 · 3 to 4 h, plus your graded tests (rung R5) |
| You build | python/tinyllm/modern/mla.py: MLAttention, ConcatLatentCache, AbsorbedMLA; and your own oracle tests in python/tests/l7-6-mla/ |
| Contract | course/contracts/py/tinyllm/modern/mla.pyi |
| Tests | course/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 |
| Needs | L7.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 by | L7.9 builds it for tl_attention = "mla" · later: L8.2 LatentCache, C1’s MLA-vs-GQA ablation at equal KV bytes |
| Milestone | MS-L7 (your decoder loads and matches Hugging Face checkpoints) |
| Optional depth | DeepSeek-AI, “DeepSeek-V2” (2024), section 2.1; “DeepSeek-V3 Technical Report” (2024), section 2.1.1 |
Key Takeaways
Section titled “Key Takeaways”- MLA caches one latent vector of width plus one shared rope key of width per token, whatever the number of heads: numbers, M05.1’s
kv_bytes_per_tokenformla(test_cache_bytes_match_m05_1). - Every head’s key and value are linear in , so a decode step folds into the query and into
o_projand 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).
How to work this chapter
Section titled “How to work this chapter”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 testsol diff L7.6 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Your GQA layer (L7.5) already shrank SmolLM2’s cache by three: 3 kv heads instead of 9. The cache still grows by 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.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| the input of token | float32[d] | |
| query heads | int | |
kv_lora_rank: width of the latent | int | |
q_lora_rank: width of the query bottleneck (None: none) | int | |
| , , | qk_nope_dim, qk_rope_dim, v_dim per head | int |
| the latent of token : what the cache holds | float32[r] | |
| the rotated rope key, shared by every head | float32[d_r] | |
| , | head ‘s rows of kv_b_proj | matrices |
head ‘s columns of o_proj | matrix | |
the RoPE rotation at position (L7.3) | ||
softmax_scale | float |
2.1 Low rank, again
Section titled “2.1 Low rank, again”M03.5 showed that a matrix of rank factors as a product of an and an matrix, and that truncating the SVD keeps the best rank- approximation. MLA bets that the keys and values of all heads, stacked as one -wide vector per token, live near an -dimensional subspace. It learns the factorization directly: a down projection to (kv_a_proj_with_mqa, the first outputs), an RMSNorm (kv_a_layernorm, L7.1), and an up projection kv_b_proj of shape . SmolLM2-sized MHA caches numbers per token; DeepSeek-V3 uses and for 128 heads.
2.2 The latent and the decoupled rope key
Section titled “2.2 The latent and the decoupled rope key”Steps of the naive form, for each token:
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 , 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 cached numbers.
2.3 What the cache holds
Section titled “2.3 What the cache holds”The cache receives after the norm and after the rotation, numbers per token per layer: ConcatLatentCache.update(layer, c, k_rope) appends a chunk and returns everything held. A query of a chunk of new tokens sits at key index , exactly as in L7.5, so prefilling in chunks equals one full forward.
2.4 Weight absorption
Section titled “2.4 Weight absorption”Rewrite the nope score and the output with matrix products moved around:
The query is mapped into latent space once (, an -vector), the scores are dot products with the cached , the weighted latent is formed once per head, and (precomputed, ) maps it out. No per-head key or value is ever built: decoding reads numbers per cached token. RoPE’s part cannot be absorbed (it depends on ), which is the reason for the decoupled rope key.
2.5 Training the naive form
Section titled “2.5 Training the naive form”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.
3. Worked example by hand
Section titled “3. Worked example by hand”, one head, , , , , inv_freq , so position 1 rotates the rope pair by 90 degrees: . Weights: ; and ; kv_b_proj gives , ; o_proj outputs . The latent norm of a single number is (the changes the 7th digit).
- Token 0, , position 0: , so ; , not rotated; . It sees only itself: output .
- Token 1, , position 1: , so , ; rotated to . Query , rotated to .
- Scores of token 1: key 0 is ; key 1 is . Scaled by : softmax gives .
- Output: , so .
Absorbed: , so and the scores are the rope terms alone, the same two numbers. The weighted latent is and : again . The cache holds latents and rope keys , : three numbers per token. This is test_hand_example.
4. The interface
Section titled “4. The interface”class LatentCacheHook(Protocol): def update(self, layer: int, c_new: NDArray, k_rope_new: NDArray) -> tuple[NDArray, NDArray]: ...class ConcatLatentCache: # update, seq_lenclass 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) -> AbsorbedMLAclass AbsorbedMLA: # w_uk [H, dn, r], w_ov [H, d, r], o_bias def forward(self, x, positions, cache, layer=0, mask=None) -> NDArrayWhat the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example | unit | section 3: outputs, weights, cache contents, absorbed form | you and the test agree on every split |
test_state_dict_keys_and_shapes_are_hf | unit | HF’s parameter names, order, shapes | DeepSeek checkpoints load by name in L7.9 |
test_mla_golden | golden | DeepseekV3Attention outputs and cached latents, plain and low-rank queries | the model in MS-L7 and C1 |
test_absorbed_decode_equals_naive | differential | prefill then decode, absorbed vs naive, float64 | L8.2 decodes with the cheap form |
test_cache_chunks_equal_full_forward | differential | chunked prefill through the cache | chunked prefill in L10.3 |
test_cache_bytes_match_m05_1 | property | cached arrays are kv_bytes_per_token for mla | admission by KV bytes, C1’s equal-bytes ablation |
test_scores_see_relative_position_only | property | shifting every position changes nothing | the decoupled rope key is RoPE |
test_padding_mask_and_fully_masked_rows | boundary | padded keys hidden in both forms; an empty row gives 0 | batched prompts of unequal length |
test_gradcheck_through_the_latent | gradcheck | float64 central differences through the latent | C1 trains MLA |
test_concat_latent_cache | unit | append, copy, seq_len, chunk checks | the hook L8.2 implements |
test_validation | boundary | rope width, sizes, scale, positions, mask shape | wiring bugs fail loudly |
Your graded tests (rung R5)
Section titled “Your graded tests (rung R5)”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.
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
1. reading kv_b_proj per head as v first, then k_nope | a DeepSeek checkpoint loads and attends to garbage | test_mla_golden, test_absorbed_decode_equals_naive (mutant s01) |
| 2. splitting the query rope-first | scores mix position-free and rotated entries | test_hand_example, test_mla_golden (mutant s02) |
| 3. forgetting to rotate the shared rope key | relative position is lost; outputs drift with absolute position | test_scores_see_relative_position_only (mutant s03) |
4. caching the latent before kv_a_layernorm | the cached keys and values are scaled wrong | test_hand_example, test_mla_golden (mutant s04) |
| 5. scaling scores by instead of | softmax too sharp; HF disagrees | test_mla_golden, test_validation (mutant s05) |
| the rope key in a different layout from the query | relative scores break only for interleaved checkpoints | test_mla_golden (mutant s06) |
absorbing with the value rows of kv_b_proj | the absorbed decode diverges from the naive one | test_absorbed_decode_equals_naive (mutant s07) |
| a cache that returns only the new chunk | each decode step forgets the past | test_cache_chunks_equal_full_forward (mutant s08) |
| the absorbed form dropping the padding mask | padded batches decode wrong only in production | test_padding_mask_and_fully_masked_rows (mutant s09) |
pairing o_proj columns with the wrong head in | absorbed outputs wrong for only | test_absorbed_decode_equals_naive (mutant s10) |
| a detached latent | kv_a_proj_with_mqa never learns | test_gradcheck_through_the_latent (mutant s11) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | L7.3 | rotates the decoupled rope parts with a RopeSpec |
| Back | L7.1 | RMSNorm on the latent and the query bottleneck |
| Back | L0.4 | Linear projections, Module registration |
| Back | L0.2 | the ops that make the naive form trainable |
| Back | L0.1 | Tensor slicing and products |
| Back | M06.3 | the default initialization stream |
| Back | M05.1 | kv_bytes_per_token for mla, checked against the cache |
| Forward | L7.9 | tl_attention = "mla" builds this layer in every block |
| Forward | L8.2 | LatentCache implements the hook and decodes with AbsorbedMLA |
| Forward | C1 | the MLA-vs-GQA ablation at equal KV bytes |
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
MLAttention | HF DeepseekV3Attention | YaRN mscale folded into the softmax scale, interleaved rope weights | transformers/models/deepseek_v3/modeling_deepseek_v3.py |
AbsorbedMLA | vLLM MLA backend, FlashMLA | absorbed decode kernels over a paged latent cache | vLLM vllm/attention/backends/mla/, DeepSeek FlashMLA |
ConcatLatentCache | SGLang MLATokenToKVPool | one latent pool per layer, radix-tree prefix sharing | SGLang mem_cache/memory_pool.py |
| the factorization | TransMLA, MHA2MLA | converting a trained GQA model to MLA with an SVD and fine-tuning | Meng et al. 2025, Ji et al. 2025 |