Lesson 13.5 · 29 min
KV Cache Compression: Trading Some Accuracy for Speed
Can we throw away three quarters of an LLM's memory of a conversation and still get nearly the same answer? Sometimes yes, and it depends on which quarter we keep.
In short: The KV cache stores a key and a value vector for every past token in every layer, and for long contexts it outgrows the model itself. KV cache compression shrinks it in four main ways: fewer bits per number (quantization), fewer tokens (eviction), fewer key/value heads (sharing across heads), or fewer dimensions per token (low-rank). This lesson measures the quality cost of each on real attention math and gives a decision guide.
How an LLM writes text, and what attention needs
An LLM generates text autoregressively: it predicts one token, appends it, and predicts again. Inside each Transformer layer, attention lets the newest token look back at earlier ones. Each token is projected into a query (what it is looking for), a key (what it can be matched on) and a value (what it contributes).
This formula tells us what compression can safely change. If we perturb keys slightly, the weights wᵢ shift slightly. If we perturb values slightly, the average shifts slightly. If we delete token i entirely, its weight is redistributed to the remaining tokens, which is harmless only if wᵢ was tiny.
Think of it like packing for a long trip Your suitcase is the GPU's memory. You can fold clothes tighter (quantization), leave some items at home (eviction), share one toiletry bag among family members (sharing heads), or vacuum-pack bulky things and expand them on arrival (low-rank). Each saves space; only one of them means you might arrive without something you need.
The KV cache and why it becomes huge
The KV cache keeps every past token's key and value so that each new step does not recompute them. Its size is a product of factors:
For a 32-layer model with 32 KV heads of size 128 in FP16, that is 512 KiB per token. Our running example, a legal-document assistant that keeps whole 100-page contracts (about 60,000 tokens) in context, needs about 29 GiB of cache per conversation, nearly twice the 16 GB of weights of an 8B model. Every decode step must also read that cache, so size costs speed too.
What KV cache compression is
KV cache compression means storing the information in the cache in fewer bytes while keeping attention outputs close to the uncompressed ones. Each method attacks one factor of the size formula: bytes_per_number (quantization), tokens (eviction), kv_heads (sharing) or the per-token dimension (low-rank).
To compare them fairly we need a quality measure. A simple one is the relative error of one attention output: ‖out_compressed − out_exact‖ / ‖out_exact‖. Zero means identical; 1.0 means an error as large as the answer itself. Real evaluations use task accuracy (for example long-document question answering), but this measure shows the mechanics clearly.
Approach 1: Quantization
Quantization stores each number as a small integer plus a shared scale. With b bits there are 2ᵇ levels; symmetric INT8 uses −127…127. To quantize a group of numbers, we find the largest magnitude, set scale = max|x| / 127, store round(x / scale), and multiply back when reading.
Quantizing one key vector to INT8
- Find the range: Say the key is [0.42, −1.27, 0.05, 0.88]. The largest magnitude is 1.27.
- Compute the scale: scale = 1.27 / 127 = 0.01.
- Round to integers: 0.42/0.01 = 42, −1.27/0.01 = −127, 0.05/0.01 = 5, 0.88/0.01 = 88. Store [42, −127, 5, 88] in one byte each plus the scale.
- Dequantize when used: Multiply back: [0.42, −1.27, 0.05, 0.88]. Here it is exact; in general each number is off by at most half a step (0.005).
Grouping matters. Per-token scales use one scale per token vector; per-channel scales use one scale per dimension across tokens. Keys often contain a few channels with consistently large values (outliers), which ruin a per-token scale for the other channels; the KIVI paper therefore quantizes keys per-channel and values per-token. In practice FP8 KV caches are widely supported in serving engines, and INT4 or 2-bit schemes need such careful grouping to hold quality.
Approach 2: Token eviction
Token eviction keeps at most a fixed number of tokens (the budget) in the cache and discards the rest. The policy decides which ones survive:
- Sliding window: keep the most recent W tokens. Simple, but forgets the start of the conversation.
- Attention sinks + window: StreamingLLM observed that many heads put large attention on the first few tokens regardless of content. Keeping those few “sink” tokens plus a recent window keeps generation stable on very long streams.
- Score-based: methods like H2O keep “heavy hitters”, tokens that accumulated the most attention so far, plus recent tokens. SnapKV-style methods pick important tokens using the attention of the last part of the prompt.
Eviction is irreversible Once a token's keys and values are gone, no later query can see them. A token that looked unimportant early on may become crucial when the user asks a follow-up. Scores from the past are only a guess about the future.
Approach 3: Sharing keys and values across heads
Standard multi-head attention (MHA) gives each query head its own key and value head. Multi-query attention (MQA) shares one K/V head among all query heads. Grouped-query attention (GQA) shares each K/V head among a group: 32 query heads with 8 K/V heads means groups of 4 and a 4× smaller cache.
Unlike quantization and eviction, sharing is an architecture decision: the model must be trained (or converted and briefly “up-trained”, as the GQA paper showed) with shared heads. Llama 3, Mistral 7B and many other recent open models use GQA, so you often get this saving simply by choosing a modern model.
Approach 4: Low-rank compression
A matrix has low rank if its rows mostly live in a smaller subspace: a few basis directions combine to approximate every row. Key and value matrices in trained models are often like this. Then we can store each token as a short vector of r coefficients plus a shared r × d basis, instead of d numbers per token.
Post-hoc low-rank methods (applied to an already trained model) exist too, but quality depends heavily on how much true redundancy the model has, so they are less of a free lunch than learned designs like MLA.
Code: measuring the quality cost of each approach
Let us build 512 cached tokens with keys that have hidden low-rank structure, a query that looks back at token 200, and measure the relative error of the attention output under each method.
compress_kv.py
import numpy as np
rng = np.random.default_rng(42)
T, d = 512, 64 # 512 cached tokens, head dim 64
K = rng.normal(size=(T, 16)) @ rng.normal(size=(16, d)) / 4 # keys: low-rank structure
K += 0.1 * rng.normal(size=(T, d))
V = rng.normal(size=(T, d))
q = K[200] / 2 + rng.normal(size=d) / 2 # query that "looks back" at token 200
def attn(q, K, V):
s = K @ q / np.sqrt(d); w = np.exp(s - s.max()); w /= w.sum()
return w @ V, w
ref, w = attn(q, K, V)
err = lambda out: np.linalg.norm(out - ref) / np.linalg.norm(ref)
def quantize(x, bits): # per-token symmetric quantization
qmax = 2 ** (bits - 1) - 1
scale = np.abs(x).max(axis=1, keepdims=True) / qmax
return np.round(x / scale).clip(-qmax, qmax) * scale
for bits in [8, 4, 2]:
print(f"quantize K,V to INT{bits}: error {err(attn(q, quantize(K, bits), quantize(V, bits))[0]):.3f}")
keep = np.r_[0:4, T-124:T] # 4 "sink" tokens + 124 most recent
print(f"evict: sinks + recent 128: error {err(attn(q, K[keep], V[keep])[0]):.3f}")
top = np.argsort(w)[-128:] # keep the 128 highest-attention tokens
print(f"evict: top-128 by score: error {err(attn(q, K[top], V[top])[0]):.3f}")
U, S, Vt = np.linalg.svd(K, full_matrices=False) # low-rank: store K as (T x r)(r x d)
for r in [32, 16, 8]:
K_r = (U[:, :r] * S[:r]) @ Vt[:r]
print(f"low-rank keys, rank {r:2d}: error {err(attn(q, K_r, V)[0]):.3f}")Output:
quantize K,V to INT8: error 0.007 quantize K,V to INT4: error 0.102 quantize K,V to INT2: error 0.878 evict: sinks + recent 128: error 1.568 evict: top-128 by score: error 0.645 low-rank keys, rank 32: error 0.041 low-rank keys, rank 16: error 0.052 low-rank keys, rank 8: error 0.345
- INT8 is almost lossless (0.7% error); INT4 is usable (about 10%) and benefits from smarter grouping; INT2 with naive scales breaks.
- Recency eviction scores worst because the query needed token 200, which the window dropped. Score-based eviction is better but still lossy, because this query's attention is spread across many tokens.
- Low rank is nearly free down to the hidden rank (16) and degrades sharply below it (rank 8).
Pause and think: Why is the rank-32 error (0.041) only slightly lower than the rank-16 error (0.052)?
The keys were built from 16 hidden directions plus small noise. Rank 16 already captures the real structure; dimensions 17–32 only add back some of the noise, which matters little.
Comparison, and when to use which
Our legal-document assistant Users ask about any clause, so eviction is risky. We choose a GQA model with 8 KV heads (4×) and an FP8 cache (2×): 60,000 tokens drop from about 29 GiB to about 3.7 GiB per conversation, with every clause still in memory.
Common mistakes Evaluating compression on short prompts or general benchmarks only; the damage appears on long-context retrieval. Quantizing keys with per-token scales and being surprised by outlier channels. And forgetting that compressed caches need kernel support in your serving engine, otherwise the dequantization overhead can eat the speed gain.
When not to compress: if contexts are short and batches small, the cache is a minor share of memory. Spend effort on weight quantization or a smaller model first.
Common mistakes and how to spot them
Most compression surprises come from two habits: quoting the saving on one part of the cache as if it were the saving on all of it, and trusting one average error number. Let us work through the first with small numbers. Take one head with d = 64, so each token stores 64 key numbers and 64 value numbers.
What is the real saving?
- Baseline: FP16 keys and values: (64 + 64) × 2 bytes = 256 bytes per token per head.
- Rank-32 keys only: Keys shrink from 64 to 32 numbers, values stay at 64: (32 + 64) × 2 = 192 bytes. The keys halved, but the whole cache is only 256 ÷ 192 ≈ 1.33× smaller. (We ignore the small shared basis.)
- INT8 values only: 64 × 2 + 64 × 1 = 192 bytes. Again 1.33×, not 2×.
- INT8 for both: (64 + 64) × 1 = 128 bytes: a true 2×. A saving counts in full only when it covers keys and values.
- Add the small extras: Scales and bases also take bytes. They are tiny for long caches, but they are why a measured ratio is always a little below the headline one.
| What we see | Likely cause | How to check |
|---|---|---|
| Short tests pass, but answers about early pages of a long contract are wrong | Eviction dropped the tokens that held the fact | Ask questions whose answers sit at the start and middle of a long context |
| Overall error is low at 4 bits, yet outputs degrade | A few large channels set the scale, and the ordinary channels round to zero | Measure the error on ordinary channels separately, as in the practice below |
| Memory fell less than the quoted ratio | Only keys or only values were compressed | Add up bytes for keys, values and scales together |
| Memory fell but tokens per second did not rise | The engine pays extra work to expand the cache at each step | Time a decode step with and without compression |
The pattern in every row: measure the thing we care about directly. Bytes for the whole cache, error where the query actually looks, and speed on a real decode step.
Practice: try it yourself
The lesson said keys often have a few channels with very large values, and that this is why keys are quantized per channel. We will see it happen. We build 200 keys with two outlier channels, quantize them to 4 bits with one scale per token and then with one scale per channel, and measure the error in two ways.
practice_outlier_channels.py
import numpy as np
rng = np.random.default_rng(3)
T, d, BITS = 200, 16, 4 # 200 cached keys, 16 channels, INT4
K = rng.normal(size=(T, d))
K[:, 5] *= 20 # channel 5 is an outlier: always large
K[:, 11] *= 8 # channel 11 is large too
def quantize(x, axis):
# axis=1: one scale per token (row). axis=0: one scale per channel (column).
qmax = 2 ** (BITS - 1) - 1 # INT4 keeps the integers -7..7
scale = np.abs(x).max(axis=axis, keepdims=True) / qmax
return np.round(x / scale).clip(-qmax, qmax) * scale, scale.size
small = [c for c in range(d) if c not in (5, 11)] # the ordinary channels
for name, axis in [("per-token", 1), ("per-channel", 0)]:
Kq, n_scales = quantize(K, axis)
err_all = np.linalg.norm(Kq - K) / np.linalg.norm(K)
err_small = np.linalg.norm((Kq - K)[:, small]) / np.linalg.norm(K[:, small])
zeroed = np.mean(Kq[:, small] == 0) # ordinary numbers rounded to 0
size = T * d * BITS / 8 + n_scales * 2 # packed ints + FP16 scales
print(f"{name:11s} error(all)={err_all:.3f} error(ordinary)={err_small:.3f} "
f"zeroed={zeroed:.0%} bytes={size:.0f}")
print("FP16 bytes:", T * d * 2)Output:
per-token error(all)=0.119 error(ordinary)=0.674 zeroed=69% bytes=2000 per-channel error(all)=0.108 error(ordinary)=0.123 zeroed=17% bytes=1632 FP16 bytes: 6400
The overall error is almost the same for both (0.119 and 0.108). But with per-token scales, 69% of the ordinary numbers were rounded to zero and their error is 0.674. The overall number hid it because the two huge channels dominate the norm. Now change it:
- Set
BITS = 8. Predict whether per-token scales still wipe out the ordinary channels, then check thezeroedcolumn. - Delete the two lines that scale channels 5 and 11. Predict which grouping wins when there are no outlier channels.
- Set
T = 20. Predict which grouping now stores more scales, and compare the two byte counts. Why does the per-token overhead grow with T while the per-channel one does not?
Pause and think: In the practice output the two groupings have nearly the same overall error. Why would a query still get much worse attention scores with per-token scales?
The score is q · k, a sum over all channels. If the query's useful signal is in the ordinary channels, per-token quantization has turned most of those numbers into zero, so the scores lose that signal. The overall error looks fine only because the two large channels dominate the norm and are stored well.
Pause and think: A report says: “We store keys at rank 16 instead of 64, so the cache is 4× smaller.” What is missing from that claim?
The values. If values are unchanged, each token stores 16 + 64 = 80 numbers instead of 128, which is only 1.6× smaller. A ratio on keys alone is not the ratio for the whole cache.
Key takeaways
- KV cache size = 2 × layers × kv_heads × head_dim × bytes × tokens × batch; each compression method shrinks one factor.
- Quantization (FP8/INT8) is nearly lossless and works on existing models; very low bits need careful grouping.
- Eviction can save the most but deletes information; recency-only policies fail when old tokens matter.
- GQA/MQA and latent attention (MLA) are architectural choices that shrink the cache from the start.
- Measure quality on your real long-context tasks before turning on aggressive compression.
Key terms
- KV cache compression: Storing cached keys and values in fewer bytes while keeping attention outputs close to exact.
- Scale (quantization): The number that maps stored integers back to real values, often max|x| divided by the largest integer level.
- Per-channel quantization: Using a separate scale for each vector dimension across tokens, which handles outlier channels in keys.
- Attention sink: Early tokens that receive large attention regardless of content; keeping them stabilises eviction-based caches.
- Heavy hitter: A token that has received a large share of attention, kept by score-based eviction policies.
- Low-rank approximation: Representing a matrix with fewer basis directions, storing short coefficient vectors per row.
- Multi-head Latent Attention: DeepSeek's design that caches a compact learned latent per token and reconstructs keys and values from it.
← 13.4 The KV Cache: Avoiding Redundant Attention Computation · 13.6 Paged Attention: OS-Inspired Memory Management for KV Caches →