MQA/GQA attention with cache hook, window, learned sinks
Overview
Section titled “Overview”| Module | L7.5 · build · Python · Pass 5 · 4 to 5 h, plus your graded tests (rung R5) |
| You build | python/tinyllm/modern/gqa.py: GQAttention, repeat_kv, the KVCacheHook protocol and its smallest implementation ConcatKVCache; and your own oracle tests in python/tests/l7-5-gqa/ |
| Contract | course/contracts/py/tinyllm/modern/gqa.pyi |
| Tests | course/tests/L7.5/test_gqa.py (what they check: section 4), golden values from transformers 5.19.0 LlamaAttention, Qwen2Attention, and GptOssAttention in course/fixtures/L7.5/gqa_hf.npz (course/oracle/L7.5/gqa_hf.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required |
| Needs | L7.3 rope_cos_sin, apply_rope, RopeSpec · L0.4 Linear, Module · L0.2 F.matmul, F.masked_fill, F.softmax, F.concat · L0.1 Tensor · M06.3 PCG32 (default initialization) · M05.1 kv_bytes_per_token (the tests check the cache against it) · reading: M09.2 (the masked softmax), Part 5’s scaled dot-product attention (or --ref-deps) |
| Used by | L7.7 sliding-window and sink attention builds on it · L7.9 every Llama attention layer (SmolLM2: 9 query heads over 3 kv heads) · L8.2 the KVCache behind the hook · later: L9.3 and L9.4 the C attention kernels take Hkv |
| Milestone | MS-L7 (SmolLM2-135M logits match Hugging Face) |
| Optional depth | Shazeer, “Fast Transformer Decoding: One Write-Head is All You Need” (2019); Ainslie et al., “GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints” (2023); Beltagy et al., “Longformer” (2020), section 3.1; Xiao et al., “Efficient Streaming Language Models with Attention Sinks” (2023); OpenAI, “gpt-oss-120b and gpt-oss-20b model card” (2025) |
Key Takeaways
Section titled “Key Takeaways”- query heads share key/value heads; query head reads kv head , each kv head repeated times in a row (
test_repeat_kv_repeats_each_head_in_a_row,test_gqa_equals_mha_with_repeated_kv_weights). - The cache hook receives only the heads, so the KV cache shrinks by : exactly
M05.1’skv_bytes_per_token(test_cache_holds_kv_heads_only). - With a cache, the new queries are the last of keys: the causal mask uses that offset, and keys are rotated with their own positions before they are cached (
test_cache_chunks_equal_full_forward). - A sliding window of lets query read keys (
test_window_reach). - A learned sink is one extra logit per head in every softmax row, dropped afterwards: a head may attend to nothing (
test_sinks_take_weight_from_every_key).
How to work this chapter
Section titled “How to work this chapter”ol start L7.5 # stubs gqa.py; prints your test path and rung (R5)ol tests L7.5 # the course tests# write your oracle tests in python/tests/l7-5-gqa/, then:ol check L7.5 # course tests and the mutation grade of your testsol diff L7.5 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Your 2017 multi-head attention gives every query head its own key and value head. At inference the engine caches keys and values for every past token (L8.2, and the Rust engine in Pass 7), so cache memory, not compute, limits how many requests run at once. Open SmolLM2-135M’s checkpoint: k_proj.weight is , a third of q_proj’s . It has 9 query heads but only 3 key/value heads, and your MHA cannot load it. Newer models add two more things to the same layer: a sliding window (Mistral, gpt-oss) so the cache stops growing, and learned sinks (gpt-oss) so a head can decline to attend. This module builds that attention layer, with the hook through which the KV cache will plug in.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| the layer input | float32[B, T, d] | |
| , | query heads, key/value heads | int, |
| , query heads per kv head | int | |
head width (d_head; not necessarily ) | int | |
| queries after RoPE | float32[B, H, T, d_h] | |
| , | keys (after RoPE) and values | float32[B, H_{kv}, T_k, d_h] |
| , | new tokens in this call, keys visible (cached plus new) | int |
window | int or none | |
learned sink logit of head (sinks) | float32[H] | |
| attention weight of query on key | float |
2.1 Grouped-query attention
Section titled “2.1 Grouped-query attention”Each head computes scaled dot-product scores , a softmax over the visible keys, and a weighted average of the values; the heads are concatenated and projected back to by o_proj. RoPE (L7.3) rotates and (with RopeSpec’s frequencies, layout, rotary width, and attention scaling) before the scores, so a score depends on the relative position of query and key only. The scale is , the head width: a dot product of unit-variance products has variance .
2.2 Sharing kv heads
Section titled “2.2 Sharing kv heads”Multi-query attention (Shazeer 2019) keeps one key/value head for all query heads (); grouped-query attention (Ainslie et al.) keeps of them, each serving a contiguous group of query heads. repeat_kv expands to with output head reading input head : , the order the checkpoints were trained with. Tiling, , loads without error and is wrong. GQA is MHA whose kv heads come in identical groups, so an MHA layer with repeated kv weights gives the same output. Because indexing accumulates gradients, a kv head’s gradient is the sum over its copies.
2.3 Who may see whom
Section titled “2.3 Who may see whom”Query of this call sits at key index (with no cache, and ). Key is visible when:
- : causal, always;
- when a window is set: the query and the keys before it;
- the optional
maskallows it (True= may attend, broadcast to : padding).
Hidden scores are , so their weights are exactly 0, and a row with nothing visible gives zeros, not NaN (M09.2).
2.4 The cache hook
Section titled “2.4 The cache hook”Decoding one token at a time, recomputing every past key and value is wasted work. With cache given, after RoPE the layer calls cache.update(layer, k, v) with this call’s keys and values (numpy arrays, after rotation, before repetition) and attends over everything the cache returns. Two consequences: the cache stores heads, which is where GQA’s saving happens (M05.1: numbers per token per layer); and keys must be rotated with their own positions before they are stored, because a key cached at position 5 must stay rotated by 5. ConcatKVCache is the smallest implementation (it concatenates along time); L8.2’s KVCache and L8.3’s paged cache implement the same update.
2.5 Learned sinks
Section titled “2.5 Learned sinks”A softmax must put weight 1 somewhere, so a head with nothing useful to read dumps it on some token (often the first: StreamingLLM’s “attention sink”). gpt-oss gives each head a learned logit appended to every row before the softmax and dropped afterwards:
The weights now sum to : a large means the head attends to almost nothing, a very negative one recovers plain attention. Adding to every score instead does nothing at all, because softmax ignores a constant shift.
2.6 Backward
Section titled “2.6 Backward”Every step is an op of L0.2 or L7.3, so autograd reaches , the four projections, their biases, and the sinks. With a cache, the cached keys and values are constants: the cache is for inference, and no gradient flows into it.
3. Worked example by hand
Section titled “3. Worked example by hand”, , , , RoPE off (inv_freq = [0], every angle 0). Weights: q_proj makes head 0’s query and head 1’s ; ; ; o_proj keeps entry 0 of head 0 and entry 1 of head 1. Tokens , ; both heads read the single kv head: , , , .
- Token 0 sees only key 0: both heads output ; the output is .
- Token 1, head 0: , scores ; , weights ; head output .
- Token 1, head 1: , scores , weights ; head output .
o_projkeeps entry 0 of head 0 and entry 1 of head 1: .
One kv head served two query heads that attend in opposite directions. This is test_hand_example.
4. The interface
Section titled “4. The interface”class KVCacheHook(Protocol): def update(self, layer: int, k_new: NDArray, v_new: NDArray) -> tuple[NDArray, NDArray]: ...class ConcatKVCache: # update(...), seq_len(layer=0)def repeat_kv(x: Tensor, n_rep: int) -> Tensor: ... # [B, Hkv, T, dh] -> [B, Hkv n_rep, T, dh]class GQAttention(Module): def __init__(self, d, n_heads, n_kv_heads, d_head, rope: RopeSpec, qkv_bias=False, window=None, sinks=False, rng=None): ... def forward(self, x, positions, mask=None, cache=None, layer=0) -> Tensor: ... # [B, T, d] -> [B, T, d]What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example | unit | section 3 | you and the test agree on the formula |
test_golden_hf | golden | Llama GQA, an extra mask, Qwen2 biases, gpt-oss sinks and window with gapped per-row positions: outputs and every gradient | SmolLM2 in MS-L7 |
test_gradcheck_every_parameter | gradcheck | float64, biases, window, sinks | every weight trains |
test_repeat_kv_repeats_each_head_in_a_row | unit | , summed gradients | checkpoints’ head grouping |
test_gqa_equals_mha_with_repeated_kv_weights | differential | GQA = MHA with grouped kv weights | what GQA is |
test_cache_chunks_equal_full_forward | differential | prefill then decode through the cache, with and without window and sinks | L8.2’s cached generation |
test_cache_holds_kv_heads_only | property | the hook gets heads, bytes = kv_bytes_per_token | L10.2 admission by KV bytes |
test_concat_cache | unit | appends, copies, rejects bad chunks | the reference hook |
test_causality_bitwise | property | future tokens change nothing, bit for bit | a decoder |
test_window_reach | boundary | key invisible, visible | Mistral-style windows, L7.7 |
test_sinks_take_weight_from_every_key | property | large sink: output 0; very negative: plain attention | gpt-oss heads |
test_fully_masked_row_is_zero | boundary | zeros, no NaN, other rows intact | padded batches |
test_positions_shift_and_per_row | property | shift invariance; per-row positions equal separate calls | left padding, continued caches |
test_attention_scaling_squares_into_the_scores | property | YaRN scaling = on the scores | rope_scaling in L7.9 |
test_parameter_names_shapes_and_draw_order | unit | HF keys, no o_proj bias, d_head default, draw order | the safetensors keys |
test_validation | boundary | heads, d_head, window, rotary width, input, positions, mask | wiring bugs fail loudly |
Your graded tests (rung R5)
Section titled “Your graded tests (rung R5)”Your oracle is the whole layer written out in numpy float64 with the module’s own state_dict(): the projections, RoPE as complex multiplication, per head the scores of kv head , the causal and window mask, the optional sink column, the softmax, the values, and o_proj. Compare forward with and without window and sinks, decode through ConcatKVCache against the full forward, and check the head order of repeat_kv, the mask, per-row positions, the parameter names, the attention scaling, the cache’s copies, and the validation. Import only tinyllm.modern.gqa, tinyllm.modern.rope, tinyllm.autograd.tensor, and tinyllm.autograd.functional. ol check L7.5 requires 0.80 with every pitfall fault killed.
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. tiling kv heads instead of repeating them | query head 1 reads kv head 1, not 0: loads fine, wrong logits | test_repeat_kv_repeats_each_head_in_a_row, test_golden_hf (mutant s01) |
| 2. scaling by instead of | attention too flat whenever | test_golden_hf (mutant s02) |
| 3. a causal mask that ignores the cached prefix | decoded tokens see only the first keys | test_cache_chunks_equal_full_forward (mutant s03) |
| 4. adding the sink to every score | no effect at all (softmax is shift invariant) | test_sinks_take_weight_from_every_key (mutant s04) |
| 5. a window one key too wide () | one extra key per query; outputs drift from the reference | test_window_reach (mutant s05) |
| 6. rotating keys after the cache | every cached key re-rotated to the newest position: generation degrades after the prompt | test_cache_chunks_equal_full_forward (mutant s06) |
| mask polarity inverted | only padding is read | test_fully_masked_row_is_zero (mutant s07) |
| per-row positions ignored | left-padded rows rotated wrongly | test_positions_shift_and_per_row (mutant s08) |
a bias on o_proj with qkv_bias | Qwen2 checkpoints fail to load | test_parameter_names_shapes_and_draw_order (mutant s09) |
| caching the repeated heads | the cache is times too large | test_cache_holds_kv_heads_only (mutant s10) |
| softmax over the wrong axis | weights sum to 1 over queries | test_hand_example (mutant s11) |
| sinks as constants | the sinks never train | test_gradcheck_every_parameter (mutant s12) |
| ignoring the attention scaling | YaRN models lose their temperature | test_attention_scaling_squares_into_the_scores (mutant s13) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | L7.3 | rope_cos_sin, apply_rope, RopeSpec rotate and |
| Back | L0.4 | the four Linear projections |
| Back | L0.2 | F.matmul, F.masked_fill, F.softmax, F.concat give the backward |
| Back | L0.1 | Tensor and its indexing (repeat_kv) |
| Back | M06.3 | PCG32 initializes the layers when no rng is given |
| Back | M05.1 | kv_bytes_per_token is what the cache hook must receive |
| Forward | L7.9 | self_attn of every Llama layer |
| Forward | L7.7 | sliding-window caches and StreamingLLM sinks build on window and sinks |
| Forward | L8.2 | KVCache.update is the hook |
| Forward | L9.3 | the C attention kernel takes Hkv and reads kv head h / n_rep |
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
GQAttention | HF LlamaAttention, GptOssAttention | interchangeable attention backends (eager, SDPA, FlashAttention) | transformers/models/llama/modeling_llama.py |
repeat_kv | FlashAttention, vLLM paged attention | no repetition in memory: the kernel indexes kv head h / n_rep directly | FlashAttention-2 num_heads_k; vLLM paged_attention |
| the cache hook | vLLM PagedAttention, SGLang RadixAttention | block tables, prefix sharing, eviction | L8.3, L8.4 in this course |
| learned sinks | StreamingLLM | keeps the first tokens’ KV forever so a sliding window stays stable | L7.7 in this course; Xiao et al. 2023 |