Skip to content

MQA/GQA attention with cache hook, window, learned sinks

ModuleL7.5 · build · Python · Pass 5 · 4 to 5 h, plus your graded tests (rung R5)
You buildpython/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/
Contractcourse/contracts/py/tinyllm/modern/gqa.pyi
Testscourse/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
NeedsL7.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 byL7.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
MilestoneMS-L7 (SmolLM2-135M logits match Hugging Face)
Optional depthShazeer, “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)
  • HH query heads share HkvH_{kv} key/value heads; query head hh reads kv head ⌊h/nrep⌋\lfloor h / n_\text{rep} \rfloor, each kv head repeated nrepn_\text{rep} 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 HkvH_{kv} heads, so the KV cache shrinks by H/HkvH / H_{kv}: exactly M05.1’s kv_bytes_per_token (test_cache_holds_kv_heads_only).
  • With a cache, the new queries are the last TT of TkT_k 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 WW lets query tt read keys t−W+1..tt - W + 1 .. t (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).
Terminal window
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 tests
ol diff L7.5 # after passing: your code against the reference

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 192×576192 \times 576, a third of q_proj’s 576×576576 \times 576. 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.

SymbolMeaningType / shape
xxthe layer inputfloat32[B, T, d]
HH, HkvH_{kv}query heads, key/value headsint, Hkv∣HH_{kv} \mid H
nrepn_\text{rep}H/HkvH / H_{kv}, query heads per kv headint
dhd_hhead width (d_head; not necessarily d/Hd / H)int
qqqueries after RoPEfloat32[B, H, T, d_h]
kk, vvkeys (after RoPE) and valuesfloat32[B, H_{kv}, T_k, d_h]
TT, TkT_knew tokens in this call, keys visible (cached plus new)int
WWwindowint or none
σh\sigma_hlearned sink logit of head hh (sinks)float32[H]
atja_{tj}attention weight of query tt on key jjfloat

Each head computes scaled dot-product scores qt⋅kj/dhq_t \cdot k_j / \sqrt{d_h}, a softmax over the visible keys, and a weighted average of the values; the heads are concatenated and projected back to dd by o_proj. RoPE (L7.3) rotates qq and kk (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 1/dh1/\sqrt{d_h}, the head width: a dot product of dhd_h unit-variance products has variance dhd_h.

Multi-query attention (Shazeer 2019) keeps one key/value head for all query heads (Hkv=1H_{kv} = 1); grouped-query attention (Ainslie et al.) keeps HkvH_{kv} of them, each serving a contiguous group of nrepn_\text{rep} query heads. repeat_kv expands [B,Hkv,T,dh][B, H_{kv}, T, d_h] to [B,H,T,dh][B, H, T, d_h] with output head hh reading input head ⌊h/nrep⌋\lfloor h / n_\text{rep} \rfloor: (a,b)→(a,a,a,b,b,b)(a, b) \to (a, a, a, b, b, b), the order the checkpoints were trained with. Tiling, (a,b,a,b,a,b)(a, b, a, b, a, b), 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.

Query tt of this call sits at key index i=Tk−T+ti = T_k - T + t (with no cache, Tk=TT_k = T and i=ti = t). Key jj is visible when:

  • j≤ij \le i: causal, always;
  • i−j<Wi - j < W when a window is set: the query and the W−1W - 1 keys before it;
  • the optional mask allows it (True = may attend, broadcast to [B,H,T,Tk][B, H, T, T_k]: padding).

Hidden scores are −∞-\infty, so their weights are exactly 0, and a row with nothing visible gives zeros, not NaN (M09.2).

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 [B,Hkv,T,dh][B, H_{kv}, T, d_h] keys and values (numpy arrays, after rotation, before repetition) and attends over everything the cache returns. Two consequences: the cache stores HkvH_{kv} heads, which is where GQA’s saving happens (M05.1: 2Hkvdh2 H_{kv} d_h 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.

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 σh\sigma_h appended to every row before the softmax and dropped afterwards:

atj=estjeσh+∑j′estj′.a_{tj} = \frac{e^{s_{tj}}}{e^{\sigma_h} + \sum_{j'} e^{s_{tj'}}}.

The weights now sum to 1−psink≤11 - p_\text{sink} \le 1: a large σh\sigma_h means the head attends to almost nothing, a very negative one recovers plain attention. Adding σh\sigma_h to every score instead does nothing at all, because softmax ignores a constant shift.

Every step is an op of L0.2 or L7.3, so autograd reaches xx, 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.

d=2d = 2, H=2H = 2, Hkv=1H_{kv} = 1, dh=2d_h = 2, RoPE off (inv_freq = [0], every angle 0). Weights: q_proj makes head 0’s query xx and head 1’s −x-x; k=xk = x; v=(x0,2x1)v = (x_0, 2 x_1); o_proj keeps entry 0 of head 0 and entry 1 of head 1. Tokens x(0)=(1,0)x^{(0)} = (1, 0), x(1)=(0,1)x^{(1)} = (0, 1); both heads read the single kv head: k(0)=(1,0)k^{(0)} = (1, 0), k(1)=(0,1)k^{(1)} = (0, 1), v(0)=(1,0)v^{(0)} = (1, 0), v(1)=(0,2)v^{(1)} = (0, 2).

  1. Token 0 sees only key 0: both heads output v(0)=(1,0)v^{(0)} = (1, 0); the output is (1,0)(1, 0).
  2. Token 1, head 0: q=(0,1)q = (0, 1), scores (0,1)/2=(0,0.707107)(0, 1) / \sqrt 2 = (0, 0.707107); e0.707107=2.028115e^{0.707107} = 2.028115, weights (1,2.028115)/3.028115=(0.330237,0.669763)(1, 2.028115) / 3.028115 = (0.330237, 0.669763); head output 0.330237 v(0)+0.669763 v(1)=(0.330237,1.339527)0.330237\, v^{(0)} + 0.669763\, v^{(1)} = (0.330237, 1.339527).
  3. Token 1, head 1: q=(0,−1)q = (0, -1), scores (0,−0.707107)(0, -0.707107), weights (0.669763,0.330237)(0.669763, 0.330237); head output (0.669763,0.660474)(0.669763, 0.660474).
  4. o_proj keeps entry 0 of head 0 and entry 1 of head 1: (0.330237,0.660474)(0.330237, 0.660474).

One kv head served two query heads that attend in opposite directions. This is test_hand_example.

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]
TestKINDChecksWhy it matters downstream
test_hand_exampleunitsection 3you and the test agree on the formula
test_golden_hfgoldenLlama GQA, an extra mask, Qwen2 biases, gpt-oss sinks and window with gapped per-row positions: outputs and every gradientSmolLM2 in MS-L7
test_gradcheck_every_parametergradcheckfloat64, biases, window, sinksevery weight trains
test_repeat_kv_repeats_each_head_in_a_rowunit(a,b)→(a,a,a,b,b,b)(a, b) \to (a, a, a, b, b, b), summed gradientscheckpoints’ head grouping
test_gqa_equals_mha_with_repeated_kv_weightsdifferentialGQA = MHA with grouped kv weightswhat GQA is
test_cache_chunks_equal_full_forwarddifferentialprefill then decode through the cache, with and without window and sinksL8.2’s cached generation
test_cache_holds_kv_heads_onlypropertythe hook gets HkvH_{kv} heads, bytes = kv_bytes_per_tokenL10.2 admission by KV bytes
test_concat_cacheunitappends, copies, rejects bad chunksthe reference hook
test_causality_bitwisepropertyfuture tokens change nothing, bit for bita decoder
test_window_reachboundarykey t−Wt - W invisible, t−W+1t - W + 1 visibleMistral-style windows, L7.7
test_sinks_take_weight_from_every_keypropertylarge sink: output 0; very negative: plain attentiongpt-oss heads
test_fully_masked_row_is_zeroboundaryzeros, no NaN, other rows intactpadded batches
test_positions_shift_and_per_rowpropertyshift invariance; per-row positions equal separate callsleft padding, continued caches
test_attention_scaling_squares_into_the_scorespropertyYaRN scaling = s2s^2 on the scoresrope_scaling in L7.9
test_parameter_names_shapes_and_draw_orderunitHF keys, no o_proj bias, d_head default, draw orderthe safetensors keys
test_validationboundaryheads, d_head, window, rotary width, input, positions, maskwiring bugs fail loudly

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 ⌊h/nrep⌋\lfloor h / n_\text{rep} \rfloor, 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.

PitfallSymptomCaught by
1. tiling kv heads instead of repeating themquery head 1 reads kv head 1, not 0: loads fine, wrong logitstest_repeat_kv_repeats_each_head_in_a_row, test_golden_hf (mutant s01)
2. scaling by 1/d1/\sqrt{d} instead of 1/dh1/\sqrt{d_h}attention too flat whenever d≠dhd \ne d_htest_golden_hf (mutant s02)
3. a causal mask that ignores the cached prefixdecoded tokens see only the first keystest_cache_chunks_equal_full_forward (mutant s03)
4. adding the sink to every scoreno effect at all (softmax is shift invariant)test_sinks_take_weight_from_every_key (mutant s04)
5. a window one key too wide (≤W\le W)one extra key per query; outputs drift from the referencetest_window_reach (mutant s05)
6. rotating keys after the cacheevery cached key re-rotated to the newest position: generation degrades after the prompttest_cache_chunks_equal_full_forward (mutant s06)
mask polarity invertedonly padding is readtest_fully_masked_row_is_zero (mutant s07)
per-row positions ignoredleft-padded rows rotated wronglytest_positions_shift_and_per_row (mutant s08)
a bias on o_proj with qkv_biasQwen2 checkpoints fail to loadtest_parameter_names_shapes_and_draw_order (mutant s09)
caching the repeated headsthe cache is nrepn_\text{rep} times too largetest_cache_holds_kv_heads_only (mutant s10)
softmax over the wrong axisweights sum to 1 over queriestest_hand_example (mutant s11)
sinks as constantsthe sinks never traintest_gradcheck_every_parameter (mutant s12)
ignoring the attention scalingYaRN models lose their temperaturetest_attention_scaling_squares_into_the_scores (mutant s13)
DirectionModuleHow it uses this
BackL7.3rope_cos_sin, apply_rope, RopeSpec rotate qq and kk
BackL0.4the four Linear projections
BackL0.2F.matmul, F.masked_fill, F.softmax, F.concat give the backward
BackL0.1Tensor and its indexing (repeat_kv)
BackM06.3PCG32 initializes the layers when no rng is given
BackM05.1kv_bytes_per_token is what the cache hook must receive
ForwardL7.9self_attn of every Llama layer
ForwardL7.7sliding-window caches and StreamingLLM sinks build on window and sinks
ForwardL8.2KVCache.update is the hook
ForwardL9.3the C attention kernel takes Hkv and reads kv head h / n_rep
Your pieceProduction equivalentWhat it addsWhere to look
GQAttentionHF LlamaAttention, GptOssAttentioninterchangeable attention backends (eager, SDPA, FlashAttention)transformers/models/llama/modeling_llama.py
repeat_kvFlashAttention, vLLM paged attentionno repetition in memory: the kernel indexes kv head h / n_rep directlyFlashAttention-2 num_heads_k; vLLM paged_attention
the cache hookvLLM PagedAttention, SGLang RadixAttentionblock tables, prefix sharing, evictionL8.3, L8.4 in this course
learned sinksStreamingLLMkeeps the first tokens’ KV forever so a sliding window stays stableL7.7 in this course; Xiao et al. 2023