Quantising the KV cache
Keys and values in 8, 4 and 2 bits, per token and per channel, and FP4 with a scale per 16: memory saved against logit drift, measured on a live transformer.
Loading the animation…
Concept
During generation a transformer keeps every past token's keys and values, the KV cache, and reads all of it for every new token. For long contexts and large batches it outgrows the weights, and decode time is mostly spent reading it (the Inference site's KV-cache chapter measures how much). Storing it in fewer bits shrinks both the memory and the time.
The animation quantises the tiny model's keys and values as they are computed, in every layer, and compares its predictions with the unquantised model's. Over the 128 predictions of four prompts:
| KV cache | bits per number | of FP16's memory | logit drift | top prediction unchanged |
|---|---|---|---|---|
| FP16 | 16 | 100.0% | 0.000106 | 128 of 128 |
| FP8 E4M3, one scale per tensor | 8 | 50.0% | 0.01259 | 127 |
| INT8 per token | 10 | 62.5% | 0.002026 | 127 |
| FP4 E2M1, E4M3 scale per 16 | 4.5 | 28.1% | 0.04504 | 127 |
| INT4 per token | 6 | 37.5% | 0.04383 | 126 |
| INT4, K per channel, V per token | 6 | 37.5% | 0.0417 | 126 |
| INT2 per token | 4 | 25.0% | 0.2667 | 98 |
| INT2, K per channel, V per token | 4 | 25.0% | 0.2744 | 105 |
Eight bits are close to free; four bits cost a few per cent of drift and almost no flipped predictions; two bits break it. The scatter at the bottom of the animation is the trade-off at a glance: memory along the bottom, damage up the side.
Per token or per channel? A per-token scale covers one token's channels (here, one head's 8); a per-channel scale covers one channel across the tokens. KIVI found that in real LLMs the keys have a few fixed channels with very large magnitudes, so keys should be quantised per channel, and the values per token. KVQuant reached the same conclusion for keys and quantises them before the rotary embedding mixes their channels. This random model has no such outlier channels, and the two layouts give about the same drift: the rule comes from trained models' statistics, not from the arithmetic.
FP4 with a fine-grained scale. The FP4 row uses E2M1 elements with one FP8 E4M3 scale per 16 channels, the layout DeepSeek-V4.1-Flash reports for its main KV cache (its section 2.4.4), which it says nearly halves the cache against FP8. With an FP8 scale per 16 values it costs 4.5 bits per number, and here it lands near INT4 per token. The Architectures site's encoder-decoder chapter covers the model; this site's reconstruction of the rounding is its own.
Maths
A KV cache holds numbers per token. At bits per number plus one -bit scale per group of , the cost per number is bits: for this model's INT4 with an FP16 scale per head of 8, for FP4 with an E4M3 scale per 16. Real heads are 64 to 128 wide, so the scale overhead there is a fraction of a bit.
Errors in the keys enter through the softmax: a score error multiplies a token's attention weight by about , so a key error matters most for the tokens a query attends to strongly. Errors in the values are averaged by the attention weights, so per-token value scales that keep each token's vector accurate are what matter.
Code
FP4 E2M1 with one E4M3 scale per 16 channels, cut from src/lib/num/tiny.ts (roundTo is the library's tested encoder):
for (let c0 = 0; c0 < r.length; c0 += block) {
const blk = r.slice(c0, c0 + block);
let s = roundTo(amaxOf(blk) / 6, "e4m3", "rne", true);
if (s === 0) s = INFO.e4m3.min_sub;
for (const v of blk) row.push(roundTo(v / s, "e2m1", "rne", true) * s);
}
The quantised keys and values replace the real ones between the projection and the attention, inside an otherwise unchanged copy of the explainer's attention:
const Q = matmul(x, w.W_q);
let K = matmul(x, w.W_k);
let V = matmul(x, w.W_v);
if (kv) [K, V] = kv(layer, K, V);