All posts

KV cache in LLMs

/AI/3 min read

An LLM generates one token at a time, and each new token needs the keys and values of every token before it. The KV cache stores them instead of recomputing them.

A large language model does not produce a sentence in one go. It generates one token at a time, and each new token is predicted by looking at everything generated so far:

The                ->  The train
The train          ->  The train was
The train was      ->  The train was late

Every step re-reads the whole sequence. That is where the waste comes from, and the KV cache is what removes it.

Q, K and V

Inside an attention layer, each token is turned into three vectors:

what it is
Query (Q)what the current token is looking for
Key (K)what each token has to offer
Value (V)the actual information a token carries

To produce the next token, the model takes the current token's Query, compares it against the Keys of every token in the sequence to get attention scores, and uses those scores to collect the matching Values.

So a single decoding step needs one Query, plus the Keys and Values of every token so far.

The problem

Do that literally and you recompute the Keys and Values of the earlier tokens at every single step:

WITHOUT A KV CACHE step 1 t1 1 K/V computed step 2 t1 t2 2 K/V computed step 3 t1 t2 t3 3 K/V computed step 4 t1 t2 t3 t4 4 K/V computed total: 10 K/V computations WITH A KV CACHE step 1 t1 1 K/V computed step 2 t1 t2 1 K/V computed step 3 t1 t2 t3 1 K/V computed step 4 t1 t2 t3 t4 1 K/V computed total: 4 K/V computations
Solid = computed at this step. Dashed = read back from the cache.

Step 2 recomputes the K and V for token 1, which it already computed in step 1. Step 3 recomputes tokens 1 and 2. Step 4 recomputes 1, 2 and 3. The work piles up quadratically, and every repeat produces exactly the same numbers as the time before.

Put real numbers on it. Starting from a 4-token prompt and generating out to 200 tokens, the uncached version computes:

4+5+6+⋯+200=20,094 K/V computations4 + 5 + 6 + \dots + 200 = 20{,}094 \text{ K/V computations}

The solution

The Keys and Values for a token never change. Once token 1 has been processed, its K and V are fixed — nothing generated later alters them.

So store them. Compute each token's K and V once, keep them in memory, and read them back on every later step. Each step then computes K and V for exactly one new token.

For the same 200-token sequence:

4+196=200 K/V computations4 + 196 = 200 \text{ K/V computations}

20,094 down to 200 — about 100 times fewer.

Why only K and V

The Query belongs only to the token being generated right now. Once that step is done it has served its purpose and is never needed again, so there is nothing to gain by keeping it.

Keys and Values are the opposite. Every token's K and V are needed at every future step, because each new Query has to be compared against all of them. They get reused for the rest of the sequence, which is exactly what makes them worth caching.

The trade-off

The cache buys speed with memory. Nothing is free: the K and V of every token sit in GPU memory for as long as the sequence is alive, and that grows linearly with sequence length.

For short sequences this is a clear win. For very long ones the cache itself becomes the constraint, which is why serving systems manage it — commonly by keeping only a sliding window of recent tokens plus the first few tokens, the attention sinks, which cannot be dropped without quality falling apart.

Summary

  • LLMs generate one token at a time, and each step needs the K and V of every previous token.
  • Recomputing them every step is quadratic and entirely redundant — the values never change.
  • Caching K and V turns 20,094 computations into 200 for a 200-token sequence.
  • Q is not cached, because it is only ever used for the current token.
  • The cost is memory, which becomes the limiting factor on long sequences.