llm-inference-explained
← /learn · 08

Quantisation for inference

Fewer bytes per weight and per cached value, faster decode.

Concept

If a decode step's time is mostly "read every weight, then read the cache", the most direct way to speed it up is to make those bytes fewer. That is what quantisation for inference is mostly about. Storing numbers in fewer bits also lets bigger models, longer contexts and larger batches fit in the same memory.

Three things can be quantised, and they behave differently:

  • Weights only (W8A16, W4A16). Weights are stored in 8 or 4 bits and converted back to 16-bit inside the matrix-multiply kernel; activations and arithmetic stay 16-bit. Decode reads half or a quarter of the weight bytes, so memory-bound steps get almost proportionally faster. Methods such as GPTQ and AWQ choose the rounding carefully so 4-bit weights lose little accuracy.
  • Weights and activations (W8A8, e.g. INT8 or FP8). Now the multiply itself runs in 8 bits on tensor cores that support it, which also raises the compute roof: that helps prefill, which is compute-bound. Activations have outliers that make this harder; LLM.int8() handles them in 16-bit separately, and SmoothQuant moves the difficulty from activations into weights.
  • The KV cache (FP8, INT8, or lower: KIVI goes to 2 bits). At long context and large batch the cache can outweigh the weights in bytes per step, so this can matter more than weight quantisation. It also doubles or quadruples how many tokens fit.

The widget shows where a decode step's bytes go. At batch 1 they are almost all weights; at long context and large batch, the cache takes over.

Interactive

Fewer bytes per step, faster decode

Llama-3-8B on 1× H100 (the simulator's cost model). The FLOP rate stays at BF16: weight-only formats dequantise before multiplying, and faster 8-bit tensor-core maths is not modelled. Accuracy is not modelled either.

Model
Weights
KV cache
Weight bytes / step
15.01 GB
was 15.01 GB
KV bytes / step
8.59 GB
was 8.59 GB
Step time
9.31 ms
memory-bound; was 9.31 ms
Tokens / s
1719
1.00× BF16

Memory: 16.06 GB of weights + 8.59 GB of cache = 24.65 GB of 80 GB. Fits.

W bf16 · KV bf169.3 ms
W bf16 · KV fp87.7 ms
W int8 · KV bf166.5 ms
W int8 · KV fp84.9 ms
W int4 · KV bf165.1 ms
W int4 · KV fp83.5 ms

Blue: weight bytes. Amber: KV-cache bytes. Right: step time.

Concept

Some figures from the same cost model, Llama-3-8B on one H100:

StepBF168-bit weights4-bit weights
batch 1, context 2,0486.201 ms, 161 tok/s3.400 ms, 294 tok/s2.000 ms, 500 tok/s

and at batch 32 with 8,192 tokens of context each, where the cache (34.36 GB per step in BF16) is more than twice the weights (15.01 GB): an FP8 cache takes the step from 18.92 ms to 12.51 ms, and 8-bit weights on top to 9.71 ms. The step stays memory-bound throughout, so time tracks bytes.

What the model leaves out: the cost of dequantising (usually hidden inside the kernel, but not always), faster 8-bit arithmetic for prefill, and, most importantly, accuracy. Fewer bits always risk quality; every method above is about keeping the loss small, and how small depends on the model and the task. Measure it on your workload.

Maths

With ww bytes per weight and κ=2L nkv dh bkv\kappa = 2 L\,n_{kv}\,d_h\,b_{kv} bytes of cache per token, a decode step over batch bb and total context cc reads

B=Pmw+b d w⏟weights+(c+b) κ⏟KV cache,B = \underbrace{P_m w + b\,d\,w}_{\text{weights}} + \underbrace{(c + b)\,\kappa}_{\text{KV cache}},

and while it is memory-bound, t≈B/β+t0t \approx B / \beta + t_0. Halving ww halves the first term; halving bkvb_{kv} halves the second. Which matters more depends on cκc\kappa against PmwP_m w: for Llama-3-8B in BF16 the cache passes the weights once the batch holds about 115,000 tokens of context in total (15.01 GB/128 KiB15.01\,\text{GB} / 128\,\text{KiB}).

A symmetric kk-bit quantiser with scale ss stores x^=s⋅round⁡(x/s)\hat{x} = s \cdot \operatorname{round}(x/s) with round⁡(x/s)∈[−2k−1,2k−1−1]\operatorname{round}(x/s) \in [-2^{k-1}, 2^{k-1}-1]; methods differ in how they choose ss (per tensor, per channel, per group) and which errors they minimise.

Code

// src/lib/inference/quantisation.ts (excerpt)
export function decodeBreakdown(m, device, nDevices, ctxPerSeq, batch) {
  const s = costModel(m, device, { nDevices }).decode(ctxPerSeq * batch, batch);
  const kv = (ctxPerSeq * batch + batch) * kvBytesPerToken(m);
  return { weightBytes: s.bytes - kv, kvBytes: kv, time: s.time /* … */ };
}

Formats are just weight_bytes and kv_bytes on the model shape, the same fields the simulator's ModelSpec has; the unit tests check the BF16 case reproduces results.md §1 (6.201 ms).