numerics-explained
← /learn · 09

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 cachebits per numberof FP16's memorylogit drifttop prediction unchanged
FP1616100.0%0.000106128 of 128
FP8 E4M3, one scale per tensor850.0%0.01259127
INT8 per token1062.5%0.002026127
FP4 E2M1, E4M3 scale per 164.528.1%0.04504127
INT4 per token637.5%0.04383126
INT4, K per channel, V per token637.5%0.0417126
INT2 per token425.0%0.266798
INT2, K per channel, V per token425.0%0.2744105

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 2×nlayers×nkv heads×dh2 \times n_{\text{layers}} \times n_{\text{kv heads}} \times d_h numbers per token. At bb bits per number plus one ss-bit scale per group of gg, the cost per number is b+s/gb + s/g bits: 4+16/8=64 + 16/8 = 6 for this model's INT4 with an FP16 scale per head of 8, 4+8/16=4.54 + 8/16 = 4.5 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 ϵ\epsilon multiplies a token's attention weight by about eϵe^{\epsilon}, 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);