Compressing the KV cache
The memory a model uses while answering is dominated by its cache of past tokens, and storing that cache in fewer bits is how you fit more users on one GPU.
- 15 min read
- 3 reading levels
- Published
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
While a model answers you, it keeps notes on everything said so far. Those notes take more memory than the model itself, so they get stored in a shorter form.
The analogy you have already lived
Sit through a three-hour meeting taking notes. If you write every sentence in full, you fill the notebook and run out of paper.
So you switch to shorthand. Abbreviations, initials, arrows. It fits, and you can still follow the thread.
You have lost something. A word you abbreviated might now be ambiguous. That is the trade, and it is exactly the trade here.
Why it exists
A model generating a reply reads everything before it, for every single new word. Redoing that work each time would be absurdly slow.
So it keeps a summary of each past token, and reuses it. That store is the KV cache — the model's notes on the conversation so far.
Here is the surprise. On a long conversation, those notes get bigger than the model's own weights.
The measurement below shows a 7-billion-parameter model needing 4 gigabytes of notes for a single long conversation. Serve sixty-four users at once and the notes need 256 gigabytes. The model's weights are 14.
How it works
full-precision notes compressed notes
16 bits per number 4 bits per number
[ 0.4712 -1.2033 0.0891 ] [ 0.5 -1.2 0.1 ]
exact, and four times approximate, and
the memory four times as many fitStore a rounded version of each number instead of the exact one. One shared scaling factor covers a small group.
Four times the users on the same card. Or four times the conversation length. Same hardware, same model.
The part that is not obvious
The two halves of the notes behave differently.
One half is the keys, used to decide what to look at. A few of their channels hold unusually large values. Rounding those together with normal ones destroys them.
The fix is to group the numbers differently for that half. Share a scaling factor down a channel, not across a token. In the measurement below that halves the damage.
That is not a detail. It is most of the difference between a working system and a broken one.
What is honestly hard here
Cutting to 4 bits is visibly lossy. In the measurement below, the best 4-bit scheme still changes the model's attention output by about 17%.
Whether that matters depends entirely on the task. Chit-chat survives it. Long-document retrieval and code often do not.
Nobody can tell you from theory whether your workload survives. You measure it.
Where you have already seen this
- Taking notes in shorthand instead of full sentences.
- Saving a photo as a smaller JPEG so more fit on the phone.
- A voice recording at lower quality that is still understandable.
Remember this
- The KV cache is the model's notes on the conversation, and it grows with every token.
- On long conversations it outgrows the model's own weights.
- Storing it in fewer bits fits more users, and the loss must be measured, not assumed.
What to learn next
- fp32, bf16, fp8 and int4 — the formats this lesson has been assuming.
- Measuring what quantisation costs you — how to know whether 17% attention error matters.
- GPU memory anatomy — where all the memory goes.
Developer — Code and libraries.
Setup
pip install torchRuns on a CPU in a couple of seconds.
How big the cache actually is, and what compressing it costs
import torch
import torch.nn.functional as F
# ---------- part 1: how big is the cache, really ----------
# (layers, kv heads, head dim) - read from each model's published config
MODELS = {
"GPT-2 small": (12, 12, 64),
"TinyLlama 1.1B": (22, 4, 64),
"Qwen3 0.6B": (28, 8, 128),
"Mistral 7B": (32, 8, 128),
}
print(f"{'model':<16} {'KV bytes/token (fp16)':>22} {'32k ctx':>10} {'x64 users':>11}")
for name, (L, H, D) in MODELS.items():
per_token = 2 * L * H * D * 2 # 2 for K and V, 2 bytes for fp16
print(f"{name:<16} {per_token:>22,} {per_token * 32768 / 2**30:>9.2f}G"
f" {per_token * 32768 * 64 / 2**30:>10.1f}G")
# ---------- part 2: quantising the cache, and what it costs ----------
torch.manual_seed(0)
T, H, D = 512, 8, 128 # 512 cached tokens, 8 kv heads
K = torch.randn(H, T, D)
V = torch.randn(H, T, D)
K[:, :, 3] *= 12 # a few channels carry outliers,
K[:, :, 57] *= 9 # which is what real K caches look like
# 64 queries: half generic, half recency-biased, so one lucky query cannot
# decide the ranking
q = torch.cat([torch.randn(H, 32, D),
K[:, -6:].mean(1, keepdim=True).expand(H, 32, D)
+ 0.5 * torch.randn(H, 32, D)], dim=1)
def quant(x, bits, axis, group=None):
"""Symmetric integer quantisation along `axis`, optionally in groups."""
qmax = 2 ** (bits - 1) - 1
if group:
s = x.reshape(*x.shape[:-1], -1, group).abs().amax(-1, keepdim=True) / qmax
r = (x.reshape(*x.shape[:-1], -1, group) / s).round().clamp(-qmax - 1, qmax) * s
return r.reshape(x.shape)
s = x.abs().amax(axis, keepdim=True) / qmax
return (x / s).round().clamp(-qmax - 1, qmax) * s
ref = F.scaled_dot_product_attention(q, K, V)
def report(name, Kq, Vq, bits_k, bits_v, extra_bits=0.0):
out = F.scaled_dot_product_attention(q, Kq, Vq)
rel = (out - ref).norm() / ref.norm()
kerr = (Kq - K).norm() / K.norm()
bits = (bits_k + bits_v) / 2 + extra_bits
print(f"{name:<34} {bits:>7.2f} {kerr:>10.4f} {rel.item():>13.5f}")
print(f"\n{'cache format':<34} {'bits':>7} {'K error':>10} {'attn error':>13}")
report("fp16 (the baseline)", K.half().float(), V.half().float(), 16, 16)
report("int8, one scale per token", quant(K, 8, -1), quant(V, 8, -1), 8, 8, 16 / 128)
report("int8, one scale per channel", quant(K, 8, -2), quant(V, 8, -2), 8, 8, 16 / 512)
report("int4, one scale per token", quant(K, 4, -1), quant(V, 4, -1), 4, 4, 16 / 128)
report("int4, groups of 32 per token", quant(K, 4, -1, 32), quant(V, 4, -1, 32),
4, 4, 16 / 32)
def quant_k_per_channel(x, bits, chunk=32):
"""KIVI's K path: a scale per channel, within a chunk of `chunk` tokens."""
H, T, D = x.shape
qmax = 2 ** (bits - 1) - 1
g = x.reshape(H, T // chunk, chunk, D)
s = g.abs().amax(2, keepdim=True) / qmax # per (chunk, channel)
return ((g / s).round().clamp(-qmax - 1, qmax) * s).reshape(H, T, D)
report("K per-channel in 32-token chunks",
quant_k_per_channel(K, 4), quant(V, 4, -1, 32), 4, 4, 16 / 32)
# KIVI's other trick: keep the most recent tokens in full precision
def with_window(x, qfn, w=32):
return torch.cat([qfn(x[:, :-w]), x[:, -w:]], dim=1)
report("...plus the last 32 tokens in fp16",
with_window(K, lambda t: quant_k_per_channel(t, 4)),
with_window(V, lambda t: quant(t, 4, -1, 32)), 4, 4, 16 / 32 + 12 * 32 / 512)model KV bytes/token (fp16) 32k ctx x64 users GPT-2 small 36,864 1.12G 72.0G TinyLlama 1.1B 22,528 0.69G 44.0G Qwen3 0.6B 114,688 3.50G 224.0G Mistral 7B 131,072 4.00G 256.0G cache format bits K error attn error fp16 (the baseline) 16.00 0.0002 0.00153 int8, one scale per token 8.12 0.0187 0.02676 int8, one scale per channel 8.03 0.0073 0.02240 int4, one scale per token 4.12 0.3169 0.37879 int4, groups of 32 per token 4.50 0.1791 0.20190 K per-channel in 32-token chunks 4.50 0.0967 0.16943 ...plus the last 32 tokens in fp16 5.25 0.0941 0.16866
Written against PyTorch 2.5.1, CPU, seeded — reproducible on this build. The layer and head counts are read from each model's published config; the tensors in part 2 are synthetic, with outlier channels added to mimic real key caches.
The size table first
Qwen3-0.6B needs three times GPT-2's cache per token, while having fewer than five times its parameters. Cache size is 2 × layers × kv_heads × head_dim × bytes, and it has nothing to do with the parameter count. A deep model with a large head dimension has an enormous cache regardless of how small it looks.
TinyLlama is the counter-example, and grouped-query attention is why. 32 query heads, but only 4 key–value heads. The cache is sized by the key–value heads, so GQA cuts it eightfold. This is the single largest architectural lever on cache size, and it is why every recent model uses it.
64 users × 32k context on Mistral 7B is 256 GB of cache against 14 GB of weights. Read that ratio again. At serving scale the weights are a rounding error and the cache is the whole problem.
Now the quantisation table
int8 is close to free. 0.022 relative attention error at half the memory. Almost every serving stack offers this and most workloads will not notice it.
int4 with one scale per token is a disaster: 0.379. A single scale must cover a whole 128-dimensional vector, and the outlier channels force that scale huge, crushing every ordinary value to zero or one.
Smaller groups fix most of it. 32-value groups cut the error from 0.379 to 0.202, at a cost of 0.38 extra bits per value for the scales.
Quantising keys per channel was the biggest single win: 0.169. The outliers in a key cache live in particular channels, consistently across tokens. Giving each channel its own scale isolates them. Values do not have that structure, so they stay per-token. This asymmetry is KIVI's central finding and it is visible directly in the K error column: 0.317 → 0.179 → 0.097.
Keeping the last 32 tokens in fp16 barely helped here: 0.16943 → 0.16866. In a real model, attention concentrates far more sharply on recent tokens than the synthetic queries here do, so this trick matters more in practice than this table shows. Reported honestly rather than dressed up.
In real serving stacks
# vLLM: the cache dtype is a server flag
# vllm serve <model> --kv-cache-dtype fp8
# transformers: a quantized cache class
# model.generate(**inputs, cache_implementation="quantized",
# cache_config={"backend": "quanto", "nbits": 4})No output block — both need a GPU and a real model, and the exact flag names move between releases. Check your version's documentation rather than trusting a snippet.
The other levers on cache size, in rough order of impact:
| Lever | Effect | Cost |
|---|---|---|
| Grouped-query attention | 4–8× smaller | architectural, decided at pretraining |
| Multi-head latent attention | ~10× smaller | architectural |
| fp8 / int8 cache | 2× smaller | negligible quality loss |
| int4 cache | 4× smaller | measurable, workload-dependent |
| Sliding-window attention | caps the cache at the window | loses long-range recall |
| Prefix / prompt caching | shares one cache across requests | needs a shared prefix |
| PagedAttention | removes fragmentation waste | none; use it |
Common mistakes
Sizing the cache from the number of query heads. It is the key–value heads that matter. On a GQA model this is an eightfold error.
Quantising keys and values the same way. The measurement above shows why. They have different outlier structure and want different scale axes.
Forgetting the scales in your memory budget. Groups of 32 with an fp16 scale is 4.5 bits per value, not 4. At group size 16 it is 5.
Benchmarking with short prompts. Cache quantisation costs nothing on 100-token prompts. Test at the context lengths you will actually serve.
Assuming a quantised cache is faster. It saves memory. Decoding is memory-bandwidth-bound, so a smaller cache usually is faster — but dequantisation costs something, and on short contexts it can be a net loss.
Try it yourself
Delete the two lines that create outlier channels and re-run. Per-token int4 will improve dramatically and the per-channel advantage will shrink. That single edit shows that the whole design of KV quantisation exists because of outliers.
What to learn next
- fp32, bf16, fp8 and int4 — the formats this lesson has been assuming.
- Measuring what quantisation costs you — how to know whether 17% attention error matters.
- GPU memory anatomy — where all the memory goes.
Researcher — Mathematics and papers.
The size equation
For a transformer with $L$ layers, $H_{kv}$ key–value heads, head dimension $d_h$, batch $B$, sequence length $S$ and $b$ bytes per element:
$$ M_{\text{KV}} = 2 \cdot L \cdot H_{kv} \cdot d_h \cdot B \cdot S \cdot b $$
Linear in sequence length and batch, and independent of the parameter count. Since weights are fixed while $B \cdot S$ grows with load, there is a crossover past which the cache dominates — and for any serious serving deployment it is already past.
Architectural reductions to $H_{kv} \cdot d_h$:
- Multi-query attention (Shazeer, 2019): $H_{kv} = 1$.
- Grouped-query attention (Ainslie et al., 2023): $H_{kv} = H_q / g$, typically $g \in {4, 8}$. The standard choice, because MQA loses quality and GQA does not.
- Multi-head latent attention (DeepSeek-V2, 2024): cache a low-rank latent $c_t^{KV}$ and reconstruct K and V by projection. Reported at roughly 1/10 the KV size of MHA with quality above GQA.
Where the outliers are
The empirical structure that determines quantisation design (Liu et al., 2024, KIVI):
- Key cache: a small number of channels have magnitudes one to two orders larger than the rest, consistently across tokens. This is the same fixed-channel outlier phenomenon Dettmers et al., 2022 documented in activations for LLM.int8().
- Value cache: no such channel structure; magnitudes are comparatively uniform.
Under symmetric quantisation, error scales with the group's dynamic range. Grouping along the axis that contains the outliers puts them in their own group; grouping across it forces every ordinary value to share a scale set by an outlier. Hence KIVI's asymmetric design:
- K: per-channel, in chunks along the token axis (needed because tokens arrive incrementally, so a whole-sequence per-channel scale would have to be recomputed).
- V: per-token, grouped along channels.
They report 2-bit KIVI reaching near-fp16 quality with 2.6× less peak memory and up to 3.5× throughput. The demonstration above reproduces the qualitative ordering at 4 bits.
Hooper et al., 2024 (KVQuant) push to 3-bit and below with four components: per-channel key quantisation applied before RoPE (rotary mixes channels and destroys the fixed-channel structure), non-uniform (sensitivity-weighted) quantisation levels, per-vector dense-and-sparse decomposition isolating outliers, and normalisation of the attention-sink token. The pre-RoPE detail is the one most likely to be missed in a reimplementation.
fp8, and why hardware prefers it
Integer quantisation needs scales, dequantisation, and awkward kernels. fp8 needs none of that on Hopper and later, where the tensor cores consume it natively. e4m3 has range ±448 and is the standard choice for the cache; e5m2 trades mantissa for range and is used for gradients.
vLLM's --kv-cache-dtype fp8 is close to free on supported hardware, which is why it, rather than int4, is the default recommendation in production.
Attention sinks
Xiao et al., 2024 (StreamingLLM) observed that the first few tokens of a sequence receive disproportionate attention regardless of content — attention "sinks" absorbing probability mass the softmax must place somewhere. Two consequences:
- Evicting the first tokens from a sliding-window cache collapses quality. Keeping four of them restores it.
- Sink tokens have extreme activation magnitudes and must be excluded from, or handled separately by, any quantisation scheme.
Beyond quantisation
| Approach | Idea | Cost |
|---|---|---|
| PagedAttention (Kwon et al., 2023) | block-based allocation, no contiguity requirement | none — pure win, eliminates 60–80% fragmentation waste |
| Token eviction (H2O, SnapKV) | drop low-attention tokens | irreversible; the dropped token may matter later |
| Prefix caching | share the cache for a common prefix | needs an exact prefix match |
| Cross-layer sharing (YOCO, CLA) | share KV across layers | architectural |
| Offloading | move cold cache to CPU or NVMe | PCIe bandwidth becomes the limit |
PagedAttention is the one with no downside and is the reason vLLM became the default serving engine. Kwon et al., 2023 report 2–4× throughput over the then-standard systems purely from eliminating memory waste, with the gain largest at long sequences and large batches.
Papers
- Shazeer, Fast Transformer Decoding: One Write-Head is All You Need, 2019 — arxiv.org/abs/1911.02150
- Dettmers et al., LLM.int8(), NeurIPS 2022 — arxiv.org/abs/2208.07339
- Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models, EMNLP 2023 — arxiv.org/abs/2305.13245
- Kwon et al., Efficient Memory Management for LLM Serving with PagedAttention, SOSP 2023 — arxiv.org/abs/2309.06180
- Zhang et al., H2O: Heavy-Hitter Oracle for Efficient Generative Inference, NeurIPS 2023 — arxiv.org/abs/2306.14048
- Xiao et al., Efficient Streaming Language Models with Attention Sinks, ICLR 2024 — arxiv.org/abs/2309.17453
- Liu et al., KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV Cache, ICML 2024 — arxiv.org/abs/2402.02750
- Hooper et al., KVQuant, NeurIPS 2024 — arxiv.org/abs/2401.18079
- DeepSeek-AI, DeepSeek-V2, 2024 — arxiv.org/abs/2405.04434 (multi-head latent attention)
What to learn next
- fp32, bf16, fp8 and int4 — the formats this lesson has been assuming.
- Measuring what quantisation costs you — how to know whether 17% attention error matters.
- GPU memory anatomy — where all the memory goes.