KV Cache Explained: Why LLM Inference Keeps Keys and Values
8 min readBytePatterns
The KV cache explained: why generation recomputes nothing for earlier tokens, what it saves, how much memory it costs per token, and why it caps the batch size.
A language model writes one token at a time, and every new token has to look back at all the tokens before it. Done naively, that means reprocessing the whole conversation for every word of the reply. The KV cache is the trick that avoids it, and it is behind two facts every LLM engineer eventually has to explain: why the first token of a reply is slow and the rest are fast, and why a GPU with plenty of compute left still cannot take more concurrent requests.
The problem it solves
Generation is a loop. Given the tokens so far, the model computes a distribution for the next one, picks a token, appends it, and repeats. Inside each transformer layer, attention projects every token into a query, a key and a value. The newest token's query is compared against the keys of all earlier tokens, and the resulting weights mix their values.
Without a cache, step t recomputes the keys and values of all t tokens, then throws them away. Over a reply of n tokens that is 1 + 2 + ... + n, about n²/2 projections per layer, almost all of them repeats.
The intuition
In a decoder, attention is causal: a token only sees itself and the tokens before it. So the key and value of token 5 depend only on tokens 1 to 5, and appending token 6 cannot change them. They are computed once and are valid forever.
The cache stores them. Each step now projects only the new token, appends its key and value to the cache, and runs its query against everything stored. Projection work per step becomes constant instead of growing with the length.
Two things do not go away, and interviewers probe both:
- Attention itself still reads the whole prefix. The new query is compared against every cached key, so each step costs time proportional to the context length. The cache removes the recomputation, not the reading.
- The cost moves into memory. Each token leaves behind a key and a value per layer, per key-value head. The size per token is
2 × layers × kv_heads × head_dim × bytes per number. For 32 layers, 32 heads of dimension 128 in 16-bit precision, that is 524,288 bytes, 512 KiB per token, and about 4.3 GB for an 8,192-token context. Every concurrent request holds its own.
That last point is why serving is often memory-bound: the batch size is limited by how many caches fit next to the weights, not by arithmetic. It is also why the first token is slow: the prefill step processes the whole prompt at once to build the cache, and every later decode step adds one token to it. As of September 2026, the common ways to shrink the cache are grouped-query or multi-query attention, which share each key-value head among several query heads, quantizing the cache to fewer bits, and allocating it in fixed-size pages so memory is not reserved for tokens that never arrive.
Watch it run
The animation puts the sequence on top and the work each step must do underneath. Four tokens are in; to produce the fifth, attention needs every earlier token's key and value, four projections. Without a cache the next step recomputes all of them, plus the new one: 4 + 5 = 9. Do that for every token and the total is quadratic in the length, 5,050 projections for 100 tokens. But a key depends only on its own token and what came before it, so it never changes. Keep them: now the step projects exactly one token and reads the rest. The next token is the same amount of work again, and the next, 100 projections for 100 tokens. Fifty times less work at a hundred tokens, and the gap widens the longer the answer gets. Then the catch: the cost moved rather than vanished, and the cache grows with every token, 512 KiB each. And every concurrent request keeps its own, which is what bounds the batch size.
The KV Cache
Step 1 of 9
Four tokens are in. To produce the fifth, attention needs every earlier token's key and value.
The same interactive animation as the lesson — step through it with the controls.
The code
A toy model, not a real transformer: one attention head with random weights in pure Python. It generates 100 steps twice, once recomputing every key and value and once with a cache, and counts projections and query-key dot products. The outputs are identical; the work is not:
import math
import random
random.seed(26)
D = 8 # head dimension
def matrix():
return [[random.gauss(0, 1 / math.sqrt(D)) for _ in range(D)] for _ in range(D)]
def matvec(w, x):
return [sum(a * b for a, b in zip(row, x)) for row in w]
def attend(q, keys, values):
scores = [sum(a * b for a, b in zip(q, k)) / math.sqrt(D) for k in keys]
top = max(scores)
weights = [math.exp(s - top) for s in scores]
total = sum(weights)
return [sum(w * v[i] for w, v in zip(weights, values)) / total for i in range(D)]
WQ, WK, WV = matrix(), matrix(), matrix()
def generate(tokens, cached):
count = {"projections": 0, "dots": 0}
cache_k, cache_v, outputs = [], [], []
for t in range(1, len(tokens) + 1):
if cached:
new = tokens[t - 1] # only the newest token is projected
cache_k.append(matvec(WK, new))
cache_v.append(matvec(WV, new))
count["projections"] += 1
count["dots"] += t # its query reads every cached key
out = attend(matvec(WQ, new), cache_k, cache_v)
else:
keys = [matvec(WK, x) for x in tokens[:t]] # recompute the whole prefix
values = [matvec(WV, x) for x in tokens[:t]]
count["projections"] += t
for i in range(t): # the full forward pass redoes every position
count["dots"] += i + 1
out = attend(matvec(WQ, tokens[i]), keys[:i + 1], values[:i + 1])
outputs.append(out)
return outputs, count
tokens = [[random.gauss(0, 1) for _ in range(D)] for _ in range(100)]
slow, slow_count = generate(tokens, cached=False)
fast, fast_count = generate(tokens, cached=True)
same = max(abs(a - b) for x, y in zip(slow, fast) for a, b in zip(x, y)) < 1e-12
print(same, slow_count, fast_count)
# True {'projections': 5050, 'dots': 171700} {'projections': 100, 'dots': 5050}
The memory side, using the lesson's shape. Sharing each key-value head among four query heads, as grouped-query attention does, divides the cache by four:
def kv_bytes_per_token(layers, kv_heads, head_dim, bytes_per_number):
return 2 * layers * kv_heads * head_dim * bytes_per_number # 2: keys and values
full = kv_bytes_per_token(32, 32, 128, 2) # 32 layers, 32 heads, 16-bit numbers
grouped = kv_bytes_per_token(32, 8, 128, 2) # 8 key-value heads shared by 32 queries
print(full, full // 1024, round(full * 8192 / 1e9, 2)) # 524288 512 4.29
print(grouped // 1024, round(grouped * 8192 * 16 / 1e9, 1)) # 128 17.2
print(round(full * 8192 * 16 / 1e9, 1)) # 68.7 sixteen 8k requests
Checked on 60 seeded random sequences of 1 to 40 tokens: the cached generation must match the brute-force recomputation step by step, and the counters must match their closed forms, n and n(n+1)/2 projections, n(n+1)/2 and n(n+1)(n+2)/6 dot products:
ok = True
for _ in range(60):
n = random.randint(1, 40)
seq = [[random.uniform(-2, 2) for _ in range(D)] for _ in range(n)]
a, ca = generate(seq, cached=False)
b, cb = generate(seq, cached=True)
ok &= all(abs(p - q) < 1e-9 for x, y in zip(a, b) for p, q in zip(x, y))
ok &= cb == {"projections": n, "dots": n * (n + 1) // 2}
ok &= ca == {"projections": n * (n + 1) // 2, "dots": n * (n + 1) * (n + 2) // 6}
print(ok) # True
The complexity
- Projections per reply of
ntokens:O(n²)without a cache,O(n)with one. - Attention dot products:
O(n³)if every step reruns the whole prefix,O(n²)with a cache; each decode step still reads alltcached keys. - Memory:
O(n)per request per layer, times the number of concurrent requests.
Where it goes wrong
- Saying the cache makes generation linear. Projections become linear; attention over the cache is still linear per step, so quadratic over the reply.
- Forgetting it is per request. Two users never share a cache, except for identical prompt prefixes, which some servers reuse deliberately.
- Ignoring context length. Doubling the context doubles the cache; long-context serving is mostly a memory problem.
- Confusing it with a response cache. It stores intermediate tensors for one generation, not answers to repeated questions.
When it shows up in interviews
ML systems and LLM infrastructure interviews ask why decoding is memory-bound, why time to first token differs from time per output token, and how to fit more concurrent users on one GPU. It sits between attention, which defines the keys and values, and sampling, which picks each token the cache helps produce.
How to say it in an interview
"During decoding, each token's key and value depend only on the tokens before it, so they never change. The KV cache stores them per layer, and each step projects only the new token and attends over the cache. That turns quadratic projection work into linear. The price is memory: two times layers times key-value heads times head dimension times bytes, per token, per request. That is why decoding is memory-bound and why grouped-query attention, cache quantization and paged allocation matter for serving."