Grouped-query attention and the KV cache
Generation stores every past key and value, and that cache grows with the number of key and value heads. Grouped-query attention shares them across query heads. Day 2 measured the saving exactly and the quality cost not at all, because the cost was below the noise.
Training processes a whole sequence at once. Generation does not. It produces one token, appends it, and runs again. Without care, every step recomputes the keys and values of every earlier token, although they never change. The fix is to keep them, and what they cost to keep is the subject of this post.
Keys and values do not change, so keep them
In causal attention, the key and value of position depend only on tokens up to . When token arrives, positions 1 to have exactly the same keys and values as before. A decoder can store them in a cache and compute only the new token's query, key and value at each step. The new query attends to the cached keys, and the new key and value are appended.
The cache costs memory for every layer, every key and value head, every head dimension and every cached position:
The leading 2 is one key and one value. For a real model at a long context this is large, and generation reads all of it for every new token. Decoding one token does little arithmetic and a lot of reading, so its speed is set by memory bandwidth. A smaller cache is both less memory and faster decoding.
The only term a design choice can change without losing information is .
Let query heads share key and value heads
Standard multi-head attention (MHA) has one key head and one value head per query head. Multi-query attention (MQA) keeps all the query heads and shares a single key and value head across them. Grouped-query attention (GQA) sits in between: the query heads split into groups, and each group shares one key and value head.
What happened?
- The cache shrinks by exactly the ratio of query heads to KV heads. Nothing else in the formula moves, so the saving is predictable before any training.
- The paper-exercise shape, 12 layers at width 512 and 4K positions in bfloat16, needs 96 MiB with MHA and 24 MiB with two KV heads. That is the Day 2 exercise, done by hand in the note.
- Each query head still has its own query projection. Heads in a group can look for different things, but they search the same keys.
| Variant | KV heads | Bytes per token per layer | Cache at 4K |
|---|---|---|---|
| MHA | 8 | 2048 | 96 MiB |
| GQA | 4 | 1024 | 48 MiB |
| GQA | 2 | 512 | 24 MiB |
| MQA | 1 | 256 | 12 MiB |
Why MQA is fastest and still loses quality
The GQA paper measured this on T5-XXL. Inference time fell from 1.51 s with MHA to 0.24 s with MQA, because decoding reads a cache one eighth or less of the size. Quality fell too, from 47.2 to 46.6 average ROUGE. With one key and value head, every query head retrieves from the same subspace, and heads can no longer specialize in what they look up. GQA with 8 groups took 0.28 s and scored 47.1, keeping most of the speed and almost none of the loss. The paper picked 8 groups because the slowdown from MQA stays small at first and grows as the group count approaches MHA.
The note flags one thing not to copy. Eight groups was a choice for a 64-head model. octlm has 8 query heads, where 8 KV heads is plain MHA. The ratio and the shape of the curve transfer, the number does not. So EXP-013 swept 1, 2, 4 and 8.
Fewer key and value heads, broadcast at attention time
The key and value projections output kv_heads × head_size columns instead of d_model. The fused kernel broadcasts each KV head across its group with enable_gqa=True. The handwritten Day 1 path does the same thing explicitly with repeat_interleave, which is how a test checks the two agree.
def kv_cache_bytes(config: DecoderConfig, length: int, element_bytes: int = 2) -> int:
"""Bytes held by the key and value cache for one sequence, all layers."""
per_token = config.cache_dims * element_bytes
return per_token * config.n_layers * config.cached_positions(length)The cache, measured at the Day 2 shape
At four layers, width 256, eight query heads and bfloat16, from runs/day2-exp013-cache.jsonl:
Every halving of the KV heads halves the cache at every length.
Quality, measured in the grid
What happened?
- All three grouped variants land inside the baseline's seed band on code. gqa-4 is 0.034 better, gqa-2 0.011 worse, mqa-1 0.005 better. None of that is a signal.
- They have fewer parameters, because the key and value projections shrink: 3,346,944 for gqa-2 against 3,740,160.
- All three are 20 to 24 percent slower per training step. Training never reads a cache, so it gets none of the benefit, and the broadcast inside the kernel costs time.
Decision: keep two KV heads. It cuts the 4K cache from 16 MiB to 4 MiB, and its quality difference is inside the noise. One KV head would halve the cache again at the same measured quality. The note does not take it, because the GQA paper's quality penalty for MQA grows with model size, and a 400-step run at width 256 cannot see it. Two heads is the hedge. The Day 4 model has two.
What we did not build, and why
- The cache itself. Day 2 measured cache sizes from the formula. Generation still recomputes the whole sequence each step. The KV cache is EXP-071 on Day 5, and its exit check is that cached and uncached generation produce the same tokens.
- Converting an MHA checkpoint to GQA. The paper mean-pools the key and value projections in each group, then trains for 5 percent of the original steps. octlm trains GQA from scratch, so it never needed the conversion.