Chapter 24
KV Cache & Grouped-Query Attention
Prefill and decode, cache memory, MQA, GQA, sliding windows, sinks, and QK-norm.
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.
24.1 Prefill, decode, and the cache
For one self-attention layer, let token states be projected to queries, keys, and values. A query at position may read only keys , so causal attention is
Prefill computes all for the prompt, applies the triangular mask, and fills the cache with . Decode computes only the new token’s , appends to the cache, and attends the new query to every cached key. The value is identical to full recomputation because the causal formula for depends only on , 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.
def project_heads(x, weight):
"""Project token states to per-head vectors."""
return np.einsum("td,dhr->thr", x, weight)
def causal_self_attention(x, w_q, w_k, w_v, w_o):
"""Full prefill: project all tokens, then apply one causal attention."""
q = project_heads(x, w_q)
k = project_heads(x, w_k)
v = project_heads(x, w_v)
heads, _ = gqa_attention(q, k, v, causal_mask(x.shape[0]))
return np.einsum("thr,hro->to", heads, w_o)
def incremental_self_attention(x, w_q, w_k, w_v, w_o):
"""Decode one token at a time, reusing cached keys and values."""
q = project_heads(x, w_q)
k = project_heads(x, w_k)
v = project_heads(x, w_v)
outs = []
for t in range(x.shape[0]):
heads, _ = gqa_attention(q[t:t + 1], k[:t + 1], v[:t + 1])
outs.append(np.einsum("thr,hro->to", heads, w_o)[0])
return np.stack(outs)
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.
24.2 Cache memory
The cache owns two arrays per layer: keys and values. If a sequence has cached tokens, layers, key-value heads, head width , and bytes per stored number, then the bytes per sequence are
The factor is not a constant hidden in implementation folklore; it is one stored key and one stored value. For , , , , and , the formula gives bytes, or 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 is therefore attractive because it lowers both stored bytes and decode-time memory traffic. Reducing with a sliding window is a different trade: it bounds memory by forgetting part of the past. Reducing through lower precision changes storage, but serving code must still preserve enough numerical accuracy for attention scores and values.
def kv_cache_bytes(layers, num_kv_heads, head_dim, tokens, bytes_per_value):
"""Bytes for K and V caches for one sequence."""
return 2 * layers * num_kv_heads * head_dim * tokens * bytes_per_value
24.3 Multi-query and grouped-query attention
Multi-head attention has query heads and independent key-value heads. Multi-query attention keeps query heads but shares one key-value head across all of them, reducing the cache by roughly times [shazeer2019fast]. Grouped-query attention chooses a middle value with : each key-value head serves a group of query heads [ainslie2023gqa].
Let map query head to its key-value group. Then
The computation is the same dot product after expanding each KV head across its query group. When , 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 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.
def expand_kv_heads(kv, num_query_heads):
"""Repeat each KV head so it serves a group of query heads."""
num_kv_heads = kv.shape[1]
if num_query_heads % num_kv_heads != 0:
raise ValueError("query heads must be a multiple of KV heads")
repeats = num_query_heads // num_kv_heads
return np.repeat(kv, repeats, axis=1)
def gqa_attention(q, k, v, mask=None):
"""Scaled dot-product attention with H query heads and G KV heads."""
k_heads = expand_kv_heads(k, q.shape[1])
v_heads = expand_kv_heads(v, q.shape[1])
scale = np.sqrt(q.shape[-1])
scores = np.einsum("thd,shd->hts", q, k_heads) / scale
if mask is not None:
scores = np.where(mask[None, :, :], scores, -np.inf)
weights = softmax(scores, axis=-1)
return np.einsum("hts,shd->thd", weights, v_heads), weights
24.4 Windows, sinks, and modern variants
A full cache grows linearly with the generated length. Sliding-window attention caps the visible past: token attends only to positions with . 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.
def causal_mask(tokens):
"""True where a query position may read a key position."""
pos = np.arange(tokens)
return pos[None, :] <= pos[:, None]
def sliding_window_mask(tokens, window, sinks=0):
"""Causal local attention, optionally keeping early sink tokens visible."""
pos = np.arange(tokens)
causal = pos[None, :] <= pos[:, None]
recent = pos[None, :] >= pos[:, None] - window + 1
sink = pos[None, :] < sinks
return causal & (recent | sink)
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 [xiao2023efficient]. 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 [yang2025qwen3]. 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 [qiu2025gated].
|
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 [ainslie2023gqa]. Attention sinks and sliding windows are streaming tools, not magic memory erasers: they trade access to distant middle tokens for a bounded cache [xiao2023efficient]. Exact cache behavior is part of a model’s architecture, so serving code must match training-time masks and head grouping. |
24.5 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 only uses .
-
Replace recomputing old with reading them from a cache.
-
Count cache bytes: two tensors, layers, KV heads, head width, tokens, bytes.
-
Draw query heads pointing to KV heads; set and then .
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 changed?
24.6 Exercises
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.
Derive . Then compute the MiB per sequence for the dimensions used in Section 24.2.
Using for the number of KV heads, describe the cases and . Why does reduce to ordinary multi-head attention?
Write a NumPy function that returns a causal sliding-window mask with an optional number of sink tokens. For , , and one sink, what may the last token read?
References
-
[ainslie2023gqa] J. Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. 2023. arXiv:2305.13245
-
[qiu2025gated] Z. Qiu et al. Gated Attention for Large Language Models: Non-linearity, Sparsity, and Attention-Sink-Free. 2025. arXiv:2505.06708
-
[shazeer2019fast] N. Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019. arXiv:1911.02150
-
[vaswani2017attention] A. Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762
-
[xiao2023efficient] G. Xiao et al. Efficient Streaming Language Models with Attention Sinks. 2023. arXiv:2309.17453
-
[yang2025qwen3] A. Yang et al. Qwen3 Technical Report. 2025. arXiv:2505.09388