Quantised LLM Inference

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.

On this page 9
  1. The short answer
  2. The analogy you have already lived
  3. Why it exists
  4. How it works
  5. The part that is not obvious
  6. What is honestly hard here
  7. Where you have already seen this
  8. Remember this
  9. What to learn next

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 fit

Store 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

Developer — Code and libraries.

Setup

bash
pip install torch

Runs on a CPU in a couple of seconds.

How big the cache actually is, and what compressing it costs

kv_cache.py
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)
Output
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

python
# 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:

LeverEffectCost
Grouped-query attention4–8× smallerarchitectural, decided at pretraining
Multi-head latent attention~10× smallerarchitectural
fp8 / int8 cache2× smallernegligible quality loss
int4 cache4× smallermeasurable, workload-dependent
Sliding-window attentioncaps the cache at the windowloses long-range recall
Prefix / prompt cachingshares one cache across requestsneeds a shared prefix
PagedAttentionremoves fragmentation wastenone; 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

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:

  1. Evicting the first tokens from a sliding-window cache collapses quality. Keeping four of them restores it.
  2. Sink tokens have extreme activation magnitudes and must be excluded from, or handled separately by, any quantisation scheme.

Beyond quantisation

ApproachIdeaCost
PagedAttention (Kwon et al., 2023)block-based allocation, no contiguity requirementnone — pure win, eliminates 60–80% fragmentation waste
Token eviction (H2O, SnapKV)drop low-attention tokensirreversible; the dropped token may matter later
Prefix cachingshare the cache for a common prefixneeds an exact prefix match
Cross-layer sharing (YOCO, CLA)share KV across layersarchitectural
Offloadingmove cold cache to CPU or NVMePCIe 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

What to learn next

What to learn next

These follow on from what you just read.

  • Quantised LLM Inference

    fp32, bf16, fp8 and int4

    Every weight and activation is stored in some number format, and the choice between them is a trade between range, precision and memory.

  • Quantised LLM Inference

    GPTQ

    GPTQ quantises a layer one column at a time and pushes each rounding error into the weights it has not reached yet, so the layer's output stays close to the original.

  • Looking Inside a Trained Model

    The logit lens

    The logit lens reads out a model's current best guess at every intermediate layer, showing a prediction sharpen gradually rather than appear all at once.