Lesson 13.4 · 25 min
The KV Cache: Avoiding Redundant Attention Computation
If an LLM re-read its entire conversation before writing every single word, how slow would it be, and how does it avoid that?
In short: When an LLM generates text one token at a time, the keys and values of earlier tokens never change, so recomputing them at every step wastes enormous work. The KV cache stores them once and reuses them, turning quadratic recomputation into linear work per step. The price is memory: the cache grows with every token, layer and user, which is why so much of inference engineering is about managing it.
How LLMs generate text
An LLM is a next-token predictor. Given a sequence of tokens (pieces of text), it outputs a probability for every token in its vocabulary. We pick one (the most likely, or a random sample), append it to the sequence and ask again. This one-token-at-a-time loop is called autoregressive generation.
Our running example is a support chatbot. The prompt is “Where is my order” and the bot replies “Your order ships today”. To produce those 4 reply tokens, the model runs 4 times, and each run sees everything before it: the prompt plus the reply so far.
Think of it like a meeting note-taker Before saying anything, a careful assistant reviews what everyone said. A forgetful one re-listens to the whole recording before each sentence. A smart one keeps a summary card per remark and just adds one new card each time someone speaks. The KV cache is the stack of cards.
What happens inside the model
Each token first becomes a vector (an embedding). The vectors then pass through a stack of identical Transformer layers; a 7B-class model has about 32. Each layer has two parts: attention, where tokens gather information from earlier tokens, and a feed-forward network, which processes each token on its own.
In attention, each token's vector is multiplied by three learned matrices to make three new vectors:
- Query (q): what this token is looking for in earlier tokens.
- Key (k): what this token can be matched on.
- Value (v): the information this token hands over when it is attended to.
In a causal (decoder-only) model, a token can only attend to itself and earlier tokens, never later ones. That rule, enforced by the causal mask, is what makes caching possible, as we will see.
The problem: repeated computation
Without any cache, each generation step feeds the whole sequence back in. Step 1 processes the 4 prompt tokens. Step 2 processes 5 tokens, step 3 processes 6, and so on. In every step, each layer recomputes keys and values for tokens whose keys and values it already computed in the previous step.
Why are they the same? A token's key and value in layer 1 depend only on that token's embedding. In layer 2 they depend on that token's layer-1 output, which (because of the causal mask) depends only on that token and earlier ones. Adding a new token at the end changes nothing about earlier positions. So the recomputation produces byte-for-byte identical numbers.
For a sequence that ends up with n tokens, the no-cache approach computes roughly 1 + 2 + … + n ≈ n²/2 key/value rows per layer. For 1,000 tokens that is about half a million rows instead of one thousand.
The solution: the KV cache
The KV cache stores the key and value vectors of every processed token, separately for each layer and each attention head. Generation now has two phases:
Generation with a KV cache
- Prefill the prompt: Run all prompt tokens through the model in one pass. Save every layer's keys and values in the cache. Predict the first reply token.
- Feed only the new token: The next step's input is just the one new token, not the whole sequence.
- Compute its q, k, v: In each layer, project the new token to its query, key and value.
- Append k and v: Add the new key and value as one extra row in that layer's cache.
- Attend over the cache: The new query is scored against all cached keys and mixes all cached values. Then the feed-forward runs for this one token only.
- Repeat: Pick the next token and go back to step 2 until the reply ends.
The code below proves the claim on a tiny single-head attention layer: computing outputs with and without a cache gives the same numbers, while the cached version does far less key/value work.
kv_cache_demo.py
import numpy as np
rng = np.random.default_rng(0)
d, steps = 8, 6 # tiny head size, 6 tokens
Wq, Wk, Wv = (rng.normal(size=(d, d)) for _ in range(3))
X = rng.normal(size=(steps, d)) # token vectors entering the layer
def attend(q, K, V):
s = K @ q / np.sqrt(d) # one score per past token
w = np.exp(s - s.max()); w /= w.sum() # softmax
return w @ V # weighted mix of values
# No cache: at every step recompute K and V for the whole prefix.
out_nc, proj_nc = [], 0
for t in range(steps):
K, V = X[:t+1] @ Wk, X[:t+1] @ Wv
proj_nc += 2 * (t + 1) # K and V rows computed
out_nc.append(attend(X[t] @ Wq, K, V))
# With cache: compute K and V only for the newest token, append.
K_cache, V_cache, out_c, proj_c = [], [], [], 0
for t in range(steps):
K_cache.append(X[t] @ Wk); V_cache.append(X[t] @ Wv); proj_c += 2
out_c.append(attend(X[t] @ Wq, np.array(K_cache), np.array(V_cache)))
print("same outputs:", np.allclose(out_nc, out_c))
print("K/V projections without cache:", proj_nc)
print("K/V projections with cache: ", proj_c)
for n in [100, 1000]:
print(f"n={n}: no cache {n*(n+1)} vs cache {2*n} projections")Output:
same outputs: True K/V projections without cache: 42 K/V projections with cache: 12 n=100: no cache 10100 vs cache 200 projections n=1000: no cache 1001000 vs cache 2000 projections
Pause and think: In the no-cache loop, what else is being wasted besides the key/value projections?
In a full model, the whole forward pass for all earlier tokens: their queries, attention and feed-forward computations in every layer. Only the last position's output is used to pick the next token, so all of that is thrown away each step.
Why only keys and values are cached, not queries
A natural question: if we cache K and V, why not Q too? Look at who uses what in one decoding step.
- The new token's query is compared with every key. It is needed only in this step, by this token.
- The keys and values of all earlier tokens are needed in this step and every future step, because every future token may attend to them.
- Earlier tokens' queries are never used again: their outputs were already computed, and causal masking means they never need to look at the new token.
One-line rule Queries are consumed once by the token that asks; keys and values are reused by every token that comes after. Only reusable things are worth caching.
Pause and think: Pause and think: would caching still be exact in a model where every token can attend to every other token, including future ones (bidirectional attention, like BERT)?
No. Adding a new token would change what earlier tokens attend to, so their representations and therefore later layers' keys and values would change. The KV cache relies on the causal mask.
How much faster does it get?
Count the work for generating a reply when the sequence ends at n tokens. Without a cache, step t processes all t tokens through every layer, so total work grows roughly with n². With a cache, step t processes one token through the projections and feed-forward layers, plus attention over t cached entries.
Attention itself still grows: step t's query must read t cached keys and values, so total attention work is still about n²/2 dot products. But the expensive projections and feed-forward computations, which dominate FLOPs in big models, drop from quadratic to linear. In practice, generation without a KV cache becomes unusably slow beyond short sequences, and every production inference engine uses one.
Speed is now limited by memory reads With the cache, each decode step does little math but must read all the model weights plus the whole cache from GPU memory. That is why decode is called memory-bound, and why the cache's size matters for speed, not just capacity.
The trade-off: speed vs memory
For a model with 32 layers, 32 KV heads of size 128 in FP16: 2 × 32 × 32 × 128 × 2 = 524,288 bytes, 512 KiB per token. A 4,000-token conversation needs 2 GiB; 32 such conversations need 64 GiB. For long contexts and many users, the cache can exceed the size of the model weights.
- Shrink it: grouped-query attention (fewer KV heads), KV quantization, eviction (lesson on KV cache compression).
- Store it without waste: paged memory blocks (PagedAttention).
- Share it: reuse the cache of a common system prompt across requests (prefix caching).
Common mistakes Thinking the KV cache changes the model's answers (it does not; it is an exact optimisation). Forgetting the cache when estimating GPU memory: teams size a GPU for the weights only and then hit out-of-memory errors under real traffic. And assuming the cache only holds the prompt: it grows with every generated token too.
Going one level deeper
So far the cache lived for one reply. A chat has many turns, and each new turn sends the whole history again. Do we have to prefill all of it each time? No. The cache rows for a token stay valid as long as every token before it is unchanged. So we can keep the cache between turns and compute only what is new.
Two turns of our support chat
- Turn 1: Prompt “Where is my order” (4 tokens) plus reply “Your order ships today” (4 tokens). The cache ends with 8 rows per layer.
- Turn 2 arrives: The customer adds “Can I change the address” (5 tokens). The full sequence is now 13 tokens, and its first 8 tokens are exactly the ones we cached.
- Reuse the shared prefix: We keep the 8 cached rows and prefill only the 5 new tokens. Without reuse we would compute 13 rows; with reuse, 5.
- What if the history is edited?: Suppose the app rewrites the third token of the history. Rows 1 and 2 are still valid. Every row from the third onward is stale, even though tokens 4 to 13 look the same, so we recompute 11 rows.
- Why later rows go stale: Above the first layer, a token's key and value are built from its hidden state, and that state has already mixed in information from all earlier tokens. Change one early token and every later hidden state changes.
| What we see | Likely cause | How to check |
|---|---|---|
| Answers differ from a run with the cache turned off | Stale rows kept after the earlier text changed | Compare the cached token ids with the new sequence, position by position |
| Replies repeat or drift after the first turn | The whole history was fed in again on top of an existing cache, so rows are duplicated | Cache rows should equal tokens processed, never more |
| Every turn is as slow as a fresh prompt | Something at the very start changes each turn, such as a timestamp in the system prompt | Find the first position where two turns differ |
| Memory grows although chats have ended | Caches of finished chats are never freed | Count cached rows against rows owned by live chats |
The rule behind all four rows is the same: a cache is tied to one exact sequence of tokens. Because a cache is exact, the first row gives us a simple test for any cache code we write: run the same input with and without the cache and compare the outputs, as kv_cache_demo.py did.
Practice: try it yourself
We will build a toy cache that lives across chat turns. It does not run a model. It only tracks which token ids have rows in the cache, reuses the longest shared prefix, and counts rows computed and bytes held. The model shape is tiny and made up: one cached token costs 1,024 bytes.
practice_prefix_reuse.py
# A toy KV cache that survives across chat turns.
# We reuse rows for the longest shared prefix and compute only the rest.
LAYERS, KV_HEADS, HEAD_DIM, BYTES = 4, 4, 16, 2 # tiny made-up model, FP16
ROW_BYTES = 2 * LAYERS * KV_HEADS * HEAD_DIM * BYTES # K + V for one token
def common_prefix(a, b):
n = 0
while n < min(len(a), len(b)) and a[n] == b[n]:
n += 1
return n
cached = [] # token ids whose K/V rows are currently in the cache
total_computed = 0
def run_turn(name, tokens):
global cached, total_computed
keep = common_prefix(cached, tokens) # rows that are still valid
new = len(tokens) - keep # rows we must compute now
cached = list(tokens) # drop stale rows, append new ones
total_computed += new
print(f"{name:22s} tokens={len(tokens):2d} reused={keep:2d} "
f"computed={new:2d} cache={len(cached) * ROW_BYTES} bytes")
turn1 = [1, 7, 7, 3, 9, 4, 2, 8] # prompt + reply, 8 tokens
turn2 = turn1 + [5, 6, 3, 1, 2] # same history + new message
edited = [1, 7, 0, 3, 9, 4, 2, 8, 5, 6, 3, 1, 2] # the third token was changed
run_turn("turn 1", turn1)
run_turn("turn 2 (appended)", turn2)
run_turn("turn 2 (edited early)", edited)
print("rows computed in total:", total_computed, "| without any reuse:", 8 + 13 + 13)Output:
turn 1 tokens= 8 reused= 0 computed= 8 cache=8192 bytes turn 2 (appended) tokens=13 reused= 8 computed= 5 cache=13312 bytes turn 2 (edited early) tokens=13 reused= 2 computed=11 cache=13312 bytes rows computed in total: 24 | without any reuse: 34
Appending reused all 8 rows and computed 5. Changing one early token left only 2 reusable rows and forced 11 new ones. Now change it:
- Change the edit so that the last token of
editeddiffers instead of the third. Predict the reused and computed counts before you run it. - Model a chat app that drops the oldest message when the context is full: call
run_turnwithturn2[4:]. Predict how many rows can be reused. - Set
KV_HEADS = 1(as in multi-query attention). Predict the new bytes per token and the cache size after turn 2. Do the reused and computed counts change?
Pause and think: In the edited turn, tokens 4 to 13 have the same ids as before. Why can we not keep their cached rows?
Their rows above the first layer were computed from hidden states that had already attended to the old third token. With a different third token those hidden states, and so their keys and values, would be different numbers. Same token id does not mean same key and value; the whole prefix must match.
Pause and think: A chat app trims old messages from the front of the history to stay under the context limit. What does that do to cache reuse, and what is a cheaper place to trim?
Trimming from the front changes the very first tokens, so the shared prefix drops to zero and the whole remaining history must be prefilled again. Reuse survives only up to the first changed position. Keeping a fixed start (for example the system prompt) and trimming rarely, in large steps, means most turns still only append.
Key takeaways
- LLMs generate one token at a time, and each step needs the keys and values of all earlier tokens.
- Because of causal masking, earlier tokens' keys and values never change, so recomputing them is pure waste.
- The KV cache stores them once per layer and head; each step computes only the new token's q, k and v.
- Queries are not cached because each is used only once, by its own token.
- The cache gives identical outputs and huge savings, but its memory grows with tokens, layers, heads and users.
Key terms
- Autoregressive generation: Generating text one token at a time, each conditioned on everything before it.
- Query, key, value: Three vectors per token in attention: what it seeks, what it can be matched on, and what it passes on.
- Causal mask: The rule that a token may attend only to itself and earlier tokens.
- KV cache: Stored keys and values of past tokens for every layer and head, reused at each generation step.
- Prefill: Processing all prompt tokens in one pass to fill the KV cache before generation starts.
- Decode step: One forward pass that processes a single new token using the cache.
← 13.3 Prefill-Decode Disaggregation: Splitting the Two Phases · 13.5 KV Cache Compression: Trading Some Accuracy for Speed →