2Day 2

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.

The KV cache

Keys and values do not change, so keep them

In causal attention, the key and value of position jj depend only on tokens up to jj. When token t+1t + 1 arrives, positions 1 to tt 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:

bytes=2×L×Hkv×dhead×T×bytes per number\text{bytes} = 2 \times L \times H_{kv} \times d_{\text{head}} \times T \times \text{bytes per number}

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 HkvH_{kv}.

Playground
Decoding
KVq01234567891011
K and V computed this step1per layer, per head
Computed so far5one per token
Cached positions5

Orange cells are keys and values computed at this step, blue cells are reused from the cache, green is the new query. The query must read every earlier key, but those keys never change once written, so the cache keeps them. Without it, generating t tokens projects about t²/2 keys and values. With it, t. The price is the memory the cells take, which is what GQA shrinks.

Generating tokens one at a time. Step forward and switch the cache off to see what gets recomputed.
MHA, GQA, MQA

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.

Playground
KV heads
Q0Q1Q2Q3Q4Q5Q6Q7K0 V0K1 V1
Shape
Element
Per token per layer256 B2 × 2 × 32 × 2
Cache for one sequence4 MiB× 4 layers × 4,096 positions
Against MHA4× smallerMHA needs 16 MiB

Each group of 4 query heads reads the same key and value head. Query heads still compute their own scores, so they can look for different things. They just search one shared set of keys.

Eight query heads and their key-value heads. Pick a shape and a length to see the cache for one sequence, computed by the same kv_cache_bytes octlm uses. A test checks the function against Python for six configurations.

What happened?

VariantKV headsBytes per token per layerCache at 4K
MHA8204896 MiB
GQA4102448 MiB
GQA251224 MiB
MQA125612 MiB
The tradeoff

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.

In octlm

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.

octlm/model.pyline 472
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)
EXP-013

The cache, measured at the Day 2 shape

At four layers, width 256, eight query heads and bfloat16, from runs/day2-exp013-cache.jsonl:

Measured
8 KV heads4 KV heads2 KV heads1 KV head
0.5 MiB1 MiB2 MiB4 MiB8 MiB16 MiB1,0244,096cached positionscache size8 KV heads4 KV heads2 KV heads1 KV head
Hover the chart to read values.
Cache bytes for one sequence at 1,024 and 4,096 positions for each KV-head count. Both axes are logarithmic, base 2.

Every halving of the KV heads halves the cache at every length.

EXP-012 and EXP-013

Quality, measured in the grid

Measured
Held-out split
baseline seed spread2.803.003.203.403.603.804.00code bits per byte, lower is betterbaselinermsnormswiglupost-normgqa-4gqa-2mqa-1modern
Hover a row to read its numbers.
Mean bits per byte over three seeds, spread drawn as a bar centered on the mean, baseline seed spread shaded. From notes/day2.md.

What happened?

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.

Skipped

What we did not build, and why