Lesson 8 of 8

Generating fast: the KV cache

Each new token needs the keys and values of every token before it. Storing them instead of recomputing them makes generation fast, and the store grows with every token.

Advanced16 min

In this lesson you will

  • Explain why the causal mask lets earlier keys and values be computed once and reused
  • Describe what the KV cache saves, and what work remains
  • Compute the cache size for a model, a context length, and a batch
  • Explain how multi-query and grouped-query attention shrink the cache

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 tt, a head compares the query qtq_t with the keys of every position up to tt and takes a weighted sum of their values:

outt=∑i≤tsoftmaxi ⁣(qt⋅kidhead)vi\mathrm{out}_t = \sum_{i \le t} \mathrm{softmax}_i\!\left(\frac{q_t \cdot k_i}{\sqrt{d_{\text{head}}}}\right) v_i

So the new position needs kik_i and viv_i for every earlier position ii, 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 ii only ever attends to positions at or before it, so everything computed at position ii, in every layer, depends only on tokens 11 through ii. 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

Without a cache: which positions have their keys and values computed or read from the cache at each step
StepThe·cat·sat·on·the·warm·mat.·ItPredicts
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

With a KV cache: which positions have their keys and values computed or read from the cache at each step
StepThe·cat·sat·on·the·warm·mat.·ItPredicts
1·the
2
3
4
5
6
Token passes so far
4
Attention scores so far
10
This step
4 passes, 10 scores
Step1 of 6
Token passes, 100-token prompt then 1,000 new tokens599,500 without, 1,099 with

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 nn tokens after a pp-token prompt processes p+(p+1)+⋯+(p+n−1)p + (p+1) + \dots + (p+n-1) token positions, which grows with the square of the length. With a cache it processes the prompt once and then one position per step: p+n−1p + n - 1 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 tt 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:

bytes=2×L×hkv×dhead×b×T×B\text{bytes} = 2 \times L \times h_{\text{kv}} \times d_{\text{head}} \times b \times T \times B

where LL is the number of layers, hkvh_{\text{kv}} the number of key/value heads, dheadd_{\text{head}} the head size, bb the bytes per stored number, TT the tokens in the context, and BB 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.

Model
4,096
1
Precision of weights and cache
Cache per token512 KiB
Whole cache2.00 GiB
Cache compared with weights0.16x

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 2×L×hkv×dhead2 \times L \times h_{\text{kv}} \times d_{\text{head}} 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.

Check yourself

Pick an answer to see why it is right or wrong. Nothing is graded. Your first answer is saved in this browser so the question can come back for review.

1With a KV cache, what does a generation step after the first compute for its one new token?
2Llama 3 8B has 32 layers, 8 key/value heads, and a head size of 128. At 16-bit precision, how much cache does one token take?
3Why is generating one token at a time with small batches usually limited by memory bandwidth rather than arithmetic?

Progress is saved in this browser only.

Up next in FrontiersDiffusion: from noise to data
Next
Transformers and LLMs
  1. 1Predicting the next token
  2. 2Attention, step by step
  3. 3Many heads and the causal mask
  4. 4Where words are
  5. 5The transformer block
  6. 6Inside a real language model
  7. 7How LLMs are trained
  8. 8Generating fast: the KV cache

Try "embedding", "softmax", "overfitting", or "backpropagation".