A language model writes one token at a time. Each new token is appended to the input and the model runs again to predict the next one. Done naively, the 500th token of a reply would reprocess all 499 tokens before it, and the 501st would reprocess 500. Almost all of that work repeats work already done. The KV cacheStored keys and values for every token already processed, kept between generation steps so each new token only needs its own keys and values computed. It grows with every token in the context.Open in glossary is the reason it does not have to.
What each step actually needs
Go back to attention. To produce the output for position , a head compares the query with the keys of every position up to and takes a weighted sum of their values:
So the new position needs and for every earlier position , in every layer and every head. Here is the important fact: those keys and values do not change when later tokens arrive. The Causal maskA mask that stops each token from attending to tokens after it, by setting those attention scores to negative infinity before the softmax. Used by models that generate text left to right.Open in glossary means position only ever attends to positions at or before it, so everything computed at position , in every layer, depends only on tokens through . Appending a new token cannot change it.
That makes the cache possible. Compute each position’s keys and values once, store them, and on every later step compute only the new token’s query, key, and value.
Generating with and without a cache
Each row is one generation step. Squares mark the positions whose keys and values the step needs.
- Keys and values computed in this step
- Read from the cache
Without a cache
| Step | The | ·cat | ·sat | ·on | ·the | ·warm | ·mat | . | ·It | Predicts |
|---|---|---|---|---|---|---|---|---|---|---|
| 1 | ·the | |||||||||
| 2 | ||||||||||
| 3 | ||||||||||
| 4 | ||||||||||
| 5 | ||||||||||
| 6 |
- Token passes so far
- 4
- Attention scores so far
- 10
- This step
- 4 passes, 10 scores
With a KV cache
| Step | The | ·cat | ·sat | ·on | ·the | ·warm | ·mat | . | ·It | Predicts |
|---|---|---|---|---|---|---|---|---|---|---|
| 1 | ·the | |||||||||
| 2 | ||||||||||
| 3 | ||||||||||
| 4 | ||||||||||
| 5 | ||||||||||
| 6 |
- Token passes so far
- 4
- Attention scores so far
- 10
- This step
- 4 passes, 10 scores
Try this
- Step through to the end. In the left grid every row is a full recomputation; in the right grid, after the first row, each step adds a single orange square.
- Compare the token passes after six steps: 39 without a cache and 9 with one. The gap grows with every step.
- Look at the attention scores. The cache cuts them too, but each step still compares the new query with every earlier key, so the per-step count keeps growing.
- Read the second readout: for a 100-token prompt and a 1,000-token reply, the difference is nearly 600,000 token passes against 1,099.
What the cache saves, and what it does not
Without a cache, generating tokens after a -token prompt processes token positions, which grows with the square of the length. With a cache it processes the prompt once and then one position per step: in total.
The first step is special. The whole prompt is processed at once, in parallel, to fill the cache. This phase is called prefill, and it keeps the hardware busy with large matrix multiplications. Every step after it, called decode, processes a single token.
Attention itself does not go away. The new token’s query still has to be compared with all cached keys, so each decode step costs more than the last. And decode has a different bottleneck. To produce one token, the hardware must read every weight of the model and every cached key and value from memory, while doing only about one multiply and one add with each number. With small batches, moving the data takes longer than the arithmetic, so decoding speed is set by memory bandwidth (Shazeer, 2019). Anything that makes the cache smaller makes decoding faster.
How big the cache gets
Every layer stores a key and a value vector for every key/value head and every token:
where is the number of layers, the number of key/value heads, the head size, the bytes per stored number, the tokens in the context, and the number of sequences processed together. Llama 2 7B, with 32 layers, 32 key/value heads, head size 128, and 16-bit numbers, needs 512 KiB per token. A full 4,096-token context is 2 GiB, for a single sequence.
How big is the cache?
Keys and values for every layer, head, and token, compared with the model's own weights.
Llama 2 7B: 32 layers, 32 query heads, one key/value head per query head, head size 128. Trained context: 4,096 tokens.
Try this
- With Llama 2 7B, raise the batch to 8 at 4,096 tokens. The cache is now larger than the model’s weights.
- Switch to Llama 3 8B. It is a slightly larger model, yet its cache per token is four times smaller. Turn on the toggle to see what it would need without grouping.
- Drag the context to 131,072 tokens with Llama 3 70B. At long contexts the cache becomes a large share of memory, and the computation needed to read a long prompt grows too.
Sharing keys and values
The cache scales with the number of key/value heads, and nothing says that has to equal the number of query heads. Multi-query attention (Shazeer, 2019) lets all query heads share a single key head and a single value head, which shrinks the cache by the number of heads, at some cost in quality. Grouped-query attentionAn attention variant in which groups of query heads share one key head and one value head. It shrinks the KV cache with little loss in quality. Multi-query attention is the extreme case of a single shared key/value head.Open in glossary (Ainslie et al., 2023) sits in between: query heads are split into groups, and each group shares one key/value head. Ainslie et al. found that it recovers quality close to full multi-head attention with speed close to multi-query attention.
It is now the common choice. Llama 2 70B and both Llama 3 8B and 70B use 8 key/value heads; Llama 3 8B has 32 query heads, so its cache is a quarter of what full multi-head attention would need. SmolLM2-135M, which you ran in your browser two lessons ago, shares 3 key/value heads among 9 query heads.
Key ideas
- Because of the causal mask, a position’s keys and values never change once computed, so they can be cached and reused at every later step.
- With a cache, generation processes the prompt once (prefill) and then one token per step (decode), instead of reprocessing the whole context each time.
- Each step still attends over the whole context, and decoding is usually limited by how fast weights and cache can be read from memory.
- The cache takes numbers per token and can outgrow the weights at long contexts and large batches.
- Multi-query and grouped-query attention shrink the cache by sharing key/value heads across query heads.