llm-inference-explained
← /learn · 02

The KV cache

What is cached, why it is exact, and how big it gets.

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.

KV cache: one row per layer and K or V, one column per cached position
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.

Cached positions
0
This step: cached
0
FLOPs
This step: uncached
0
FLOPs
Max |Δ logit|
–

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:

2×32×8×128×2=131,0722 \times 32 \times 8 \times 128 \times 2 = 131{,}072 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).

Model
Attention
KV precision
Values / token
65,536
Bytes / token
128.0 KiB
Whole batch
17.18 GB
80 GB GPUs of KV
0.21
cache alone, no weights
Bytes per token for each attention variant
VariantKV headsBytes / tokenvs MHA
MHA32512.0 KiB100.0%
GQA8128.0 KiB25.0%
MQA116.0 KiB3.1%
MLAlatent 512+6436.0 KiB7.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):

VariantValues per token
MHA2 nh dh l2\,n_h\,d_h\,l
GQA2 ng dh l2\,n_g\,d_h\,l
MQA2 dh l2\,d_h\,l
MLA(dc+dhR) l(d_c + d_h^R)\,l

for nhn_h heads of dimension dhd_h, ngn_g KV groups, ll layers, KV compression dimension dcd_c and decoupled rotary key dimension dhRd_h^R. Bytes are values × bytes per value × tokens.

Why it's exact. Layer by layer, the hidden state of position ii depends only on tokens 0..i0..i (the causal mask), so ki=WKhik_i = W_K h_i and vi=WVhiv_i = W_V h_i never change once computed. For the new position tt,

ot=∑j≤tsoftmax⁡j ⁣(qt⋅kjdh)vjo_t = \sum_{j \le t} \operatorname{softmax}_j\!\left(\frac{q_t \cdot k_j}{\sqrt{d_h}}\right) v_j

uses qtq_t (new) and kj,vjk_j, v_j for j≤tj \le t (cached, plus the new one).

Work saved. Recomputing a length-SS sequence costs L(8SD2+4S2D+4SDF)+2SDVL(8SD^2 + 4S^2D + 4SDF) + 2SDV FLOPs; a cached step costs L(8D2+4SD+4DF)+2DVL(8D^2 + 4SD + 4DF) + 2DV. Over a whole generation the uncached total grows like S2S^2 (and S3S^3 in the attention term), the cached one like SS.

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.