= KV Cache & Grouped-Query Attention Autoregressive transformers do two different jobs at inference time. *Prefill* reads the whole prompt in parallel; *decode* appends one token at a time, where redoing all previous key and value projections would waste most of the work. A KV cache stores those projected keys and values, and grouped-query attention shrinks that cache without changing the number of query heads that model different subspaces. [#sec-prefill-decode] == Prefill, decode, and the cache For one self-attention layer, let token states stem:[\vx_t] be projected to queries, keys, and values. A query at position stem:[t] may read only keys stem:[s \le t], so causal attention is [latexmath#eq-causal-attention] ++++ \va_t = \sum_{s \le t} \softmax_s\!\left(\frac{\vq_t^\T \vk_s}{\sqrt{d_h}}\right)\vv_s . ++++ Prefill computes all stem:[\vq_t, \vk_t, \vv_t] for the prompt, applies the triangular mask, and fills the cache with stem:[(\vk_s, \vv_s)]. Decode computes only the new token's stem:[\vq_t, \vk_t, \vv_t], appends stem:[(\vk_t, \vv_t)] to the cache, and attends the new query to every cached key. The value is identical to full recomputation because the causal formula for stem:[\va_t] depends only on stem:[s \le t], and those cached vectors are exactly the old projections. The important distinction is parallelism, not mathematics. During prefill, all prompt tokens are known, so the implementation can form a block of scores and mask its upper triangle. During decode, only one new row of scores is needed. Without a cache, the layer would rebuild the same old keys and values at every generated token; with a cache, it performs one new projection and one attention row. The output stream is unchanged as long as the cache stores the same dtype and the same mask is used. .Full prefill and cached decode use the same attention equation [source,python] ---- include::../../scratch/kv_cache.py[tag=cache] ---- The chapter tests build a small random causal attention layer and assert that decoding token by token matches the full prefill output to floating-point tolerance. [#sec-cache-memory] == Cache memory The cache owns two arrays per layer: keys and values. If a sequence has stem:[T] cached tokens, stem:[L] layers, stem:[G] key-value heads, head width stem:[d_h], and stem:[b] bytes per stored number, then the bytes per sequence are [latexmath#eq-kv-memory] ++++ M_{KV} = 2\,L\,G\,d_h\,T\,b . ++++ The factor stem:[2] is not a constant hidden in implementation folklore; it is one stored key and one stored value. For stem:[L=32], stem:[G=8], stem:[d_h=128], stem:[T=4096], and stem:[b=2], the formula gives stem:[536{,}870{,}912] bytes, or stem:[512] MiB, per sequence. The test suite computes and asserts those numbers, because cache accounting is the difference between a batch that fits and a batch that crashes. This formula is per sequence. A server batch multiplies it by the number of active requests, then adds model weights, temporary activations, and allocator overhead. Reducing stem:[G] is therefore attractive because it lowers both stored bytes and decode-time memory traffic. Reducing stem:[T] with a sliding window is a different trade: it bounds memory by forgetting part of the past. Reducing stem:[b] through lower precision changes storage, but serving code must still preserve enough numerical accuracy for attention scores and values. .Bytes for one sequence's cache [source,python] ---- include::../../scratch/kv_cache.py[tag=memory] ---- [#sec-gqa] == Multi-query and grouped-query attention Multi-head attention has stem:[H] query heads and stem:[H] independent key-value heads. *Multi-query attention* keeps stem:[H] query heads but shares one key-value head across all of them, reducing the cache by roughly stem:[H] times <>. *Grouped-query attention* chooses a middle value stem:[G] with stem:[1 < G < H]: each key-value head serves a group of stem:[H/G] query heads <>. Let stem:[g(h) = \lfloor hG/H \rfloor] map query head stem:[h] to its key-value group. Then [latexmath#eq-gqa] ++++ \va_{t,h} = \sum_{s \le t} \softmax_s\!\left(\frac{\vq_{t,h}^\T \vk_{s,g(h)}}{\sqrt{d_h}}\right)\vv_{s,g(h)} . ++++ The computation is the same dot product after expanding each KV head across its query group. When stem:[G=H], every query head gets its own KV head, so the implementation reduces exactly to ordinary multi-head attention; the test checks that equality numerically. GQA leaves the query projection wide. That matters because query heads choose different ways to look at the same context, while the shared KV heads decide how many different memories are stored. With stem:[G] between the MQA and MHA extremes, the model can keep several kinds of memory while paying less cache bandwidth than full multi-head attention. The grouping must be fixed by the architecture; changing it at serving time would change the attention computation. .Grouped-query attention by repeating KV heads [source,python] ---- include::../../scratch/kv_cache.py[tag=gqa] ---- [#sec-windows-sinks] == Windows, sinks, and modern variants A full cache grows linearly with the generated length. Sliding-window attention caps the visible past: token stem:[t] attends only to positions stem:[s] with stem:[t-W+1 \le s \le t]. The mask is still causal, but old positions outside the window are hidden. Windowing is exact only for a model trained or adapted to that mask. If a full-context model is served with old keys silently removed, its later layers may ask for evidence that is no longer visible. Sink tokens soften the boundary by leaving a small global landing pad that every later position can still read. They do not keep arbitrary facts from the dropped middle of the prompt; they mainly stabilize the attention pattern during streaming. .Causal sliding-window masks with optional sink tokens [source,python] ---- include::../../scratch/kv_cache.py[tag=masks] ---- Attention sinks are a small set of early tokens kept visible even when the rest of the distant past slides away; StreamingLLM observes that this preserves stable streaming behavior better than dropping every old key <>. QK-norm normalizes queries and keys before their dot product, or equivalently controls the dot-product scale, and Qwen reports using it in current models <>. Gated attention adds learned gates around attention outputs or weights to introduce extra nonlinearity and sparsity; recent work studies it as an attention-sink-free alternative <>. [NOTE,caption=In practice] ==== KV-cache size is often the memory bottleneck of long-context serving, so production decoders pair caching with MQA, GQA, paging, or windowing. GQA is widely used because it recovers much of multi-head quality while moving fewer KV bytes per token <>. Attention sinks and sliding windows are streaming tools, not magic memory erasers: they trade access to distant middle tokens for a bounded cache <>. Exact cache behavior is part of a model's architecture, so serving code must match training-time masks and head grouping. ==== [.key-equations#key-equations] .Key equations **** [latexmath] ++++ \va_t = \sum_{s \le t}\softmax_s\!\left(\vq_t^\T\vk_s/\sqrt{d_h}\right)\vv_s ++++ [latexmath] ++++ M_{KV} = 2\,L\,G\,d_h\,T\,b ++++ [latexmath] ++++ g(h) = \lfloor hG/H \rfloor, \qquad 1 \le G \le H ++++ [latexmath] ++++ \text{visible}(t, s) = (s \le t) \land (s \ge t-W+1 \;\lor\; s < S) ++++ **** [.teach] [#sec-teach] == Teach it *The one-sentence version.* A KV cache remembers the old keys and values, while GQA stores fewer kinds of keys and values than queries. *An analogy.* Prefill is reading a whole book and making index cards. Decode is answering the next question by adding one new card and searching the cards already written. GQA lets several searchers share one drawer of cards. *At the board.* . Write the causal attention sum and circle that stem:[\va_t] only uses stem:[s \le t]. . Replace recomputing old stem:[\vk_s, \vv_s] with reading them from a cache. . Count cache bytes: two tensors, layers, KV heads, head width, tokens, bytes. . Draw stem:[H] query heads pointing to stem:[G] KV heads; set stem:[G=1] and then stem:[G=H]. *Misconceptions to address.* * "The cache approximates attention." It is exact for the same mask and weights. * "GQA reduces query heads." It reduces KV heads; query heads remain. * "A sliding window keeps all long-range information." It keeps only the window and any sinks. *Check for understanding.* If a model has fewer KV heads but the same query heads, which term in stem:[2LGd_hTb] changed? [#sec-exercises] == Exercises [#ex-kv-cache-prefill-decode.exercise] .★ Prefill versus decode ==== Explain why cached decoding gives the same output as full causal recomputation for the newest token. Name the condition on the mask that makes the argument true. ==== [#ex-kv-cache-memory.exercise] .★★ Cache accounting ==== Derive stem:[M_{KV} = 2LGd_hTb]. Then compute the MiB per sequence for the dimensions used in <>. ==== [#ex-kv-cache-gqa.exercise] .★★ Grouped-query extremes ==== Using stem:[G] for the number of KV heads, describe the cases stem:[G=1] and stem:[G=H]. Why does stem:[G=H] reduce to ordinary multi-head attention? ==== [#ex-kv-cache-window.exercise] .★★★ Implement a streaming mask ==== Write a NumPy function that returns a causal sliding-window mask with an optional number of sink tokens. For stem:[T=6], stem:[W=3], and one sink, what may the last token read? ==== [bibliography] [#sec-references] == References include::../../book/sources.adoc[tags=vaswani2017attention;shazeer2019fast;ainslie2023gqa;xiao2023efficient;yang2025qwen3;qiu2025gated]