Concept
In every attention layer, each position computes a query, a key and a value from its own hidden state. The causal mask means position i only ever looks at keys and values of positions 0…i.
Now generate one more token. The new position needs its own query, and the keys and values of every earlier position. But those earlier keys and values were computed from earlier hidden states, which depend only on earlier tokens, and appending a token can't change the past. So they are exactly what they were last step. Recomputing them is pure waste.
The KV cache keeps them. After prefill it holds the prompt's keys and values for every layer. Each decode step computes the new token's query, key and value, appends the key and value, and attends over the whole cache. Everything else in the step (the MLP, the output head) only ever needed the new position anyway.
This is not an approximation. Below, a real (untrained) model generates with the cache, and at every step the full sequence is also recomputed from scratch without one. The two logit vectors are compared number by number. They match exactly, because the cached path performs the same floating-point operations in the same order; the masked future positions that the full pass carries contribute exact zeros.
Interactive
Watch the cache fill, and check it against recomputation
Each click runs one step of the same untrained toy model with greedy decoding. The cache gains one column per layer; a full uncached forward pass runs alongside and the logits are compared exactly.
| pos | ||||||||||||||||||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| L0 K | ||||||||||||||||||||||||||||||||
| L0 V | ||||||||||||||||||||||||||||||||
| L1 K | ||||||||||||||||||||||||||||||||
| L1 V |
Green: written by prefill. Blue: written by decode steps. Amber: the row written this step. Grey: empty.
Concept
The cache turns each decode step from "process the whole sequence" into "process one token, read the cache". The price is memory. Every layer stores a key and a value vector per token per KV head, for as long as the sequence lives.
How big? For Llama-3-8B, which has 32 layers and 8 KV heads of 128 dimensions, in BF16:
bytes, 128 KiB per token. Llama-3-70B (80 layers, 8 KV heads) needs 320 KiB per token. Sixteen 8B sequences of 8,192 tokens hold 17.18 GB of cache, more than the 16.06 GB of weights. On a GPU with 80 GB, the cache, not the weights, decides how many users you can serve at once. That is why the next chapters spend so much effort on it.
Shrinking it. Multi-head attention (MHA) keeps K and V for every head. Multi-query attention (MQA) shares one K/V head across all query heads. Grouped-query attention (GQA), which Llama 3 uses, shares each K/V head across a group (32 query heads, 8 KV heads: 4 per group), cutting the cache 4× against MHA with little quality loss. Multi-head latent attention (MLA), from DeepSeek-V2, caches one compressed latent per token per layer (512 values, plus a 64-value key for rotary positions) and reconstructs the heads' keys and values from it. Use the calculator to compare them.
Interactive
How big is the KV cache?
Bytes = values per token × bytes per value × tokens. Llama shapes from Disaggregated_Inference_Sim's hardware.py; DeepSeek-V2 and the MLA sizes from its paper (arXiv:2405.04434).
| Variant | KV heads | Bytes / token | vs MHA |
|---|---|---|---|
| MHA | 32 | 512.0 KiB | 100.0% |
| GQA | 8 | 128.0 KiB | 25.0% |
| MQA | 1 | 16.0 KiB | 3.1% |
| MLA | latent 512+64 | 36.0 KiB | 7.0% |
No Llama model uses MLA; its row applies DeepSeek-V2’s latent sizes to this layer count, for comparison only.
Maths
Per token, summed over layers, the cache holds (DeepSeek-V2, Table 1):
| Variant | Values per token |
|---|---|
| MHA | |
| GQA | |
| MQA | |
| MLA |
for heads of dimension , KV groups, layers, KV compression dimension and decoupled rotary key dimension . Bytes are values × bytes per value × tokens.
Why it's exact. Layer by layer, the hidden state of position depends only on tokens (the causal mask), so and never change once computed. For the new position ,
uses (new) and for (cached, plus the new one).
Work saved. Recomputing a length- sequence costs FLOPs; a cached step costs . Over a whole generation the uncached total grows like (and in the attention term), the cached one like .
Code
// src/lib/transformer/kvcache.ts — one cached decode step (excerpt)
const q = mm(ln1, w.attn.W_q, counter)[0]!;
const k = mm(ln1, w.attn.W_k, counter)[0]!;
const v = mm(ln1, w.attn.W_v, counter)[0]!;
layer.K.push(k);
layer.V.push(v);
// … then per head: scores = q · Kᵀ / √d_head over every cached position.
The test that makes the claim (tests/unit/transformer/kvcache.test.ts)
generates 12 tokens greedily on two model sizes and asserts, with no
tolerance, that every cached logit vector equals full recomputation's.