Inference optimisation
Be able to measure latency and throughput, and use the KV cache, batching and speculative decoding.
Prerequisites
- DTransformers — the architecturerequired
- FQuantisationrequired
Intuition
LLM inference has two phases with entirely different bottlenecks:
| Phase | What | The bottleneck | The remedy |
|---|---|---|---|
| Prefill | reads the whole prompt in | computation (parallel over the tokens) | faster kernels, flash attention |
| Decoding | generates one token at a time | memory bandwidth (all the weights are read per token) | quantisation, MQA/GQA, speculative decoding |
That is why quantisation helps most in decoding: half as many bytes to read ≈ twice as fast token generation.
Two measures that are not the same thing:
- Latency — the time to the first token (TTFT) and the time per token. What one user experiences.
- Throughput — tokens per second in total. What the system copes with.
A larger batch raises the throughput and raises the latency. You cannot optimise both.
Formal
Continuous batching is the single largest gain in a service. Static batching waits for a whole batch and everybody has to wait for the longest sequence. Continuous batching (vLLM, TGI) lets new requests jump in as soon as a slot becomes free — 2–10× higher throughput on the same hardware.
PagedAttention handles the KV cache in pages instead of contiguous blocks, which in practice eliminates fragmentation and makes it possible to run far more simultaneous sequences.
Speculative decoding: a small «draft model» generates k tokens, the large one verifies them all in one forward pass and accepts the longest correct prefix. At 70 % acceptance it becomes ~2× faster — without changing the output distribution (it is mathematically exact, not an approximation).
A measurement order that produces results: measure first (where does the time go — prefill or decoding?), then turn continuous batching on, after that quantisation, and last speculative decoding. Starting at the wrong end is the most common waste of time.
Code
import time, numpy as np
def measure_inference(generate, prompt, n_tokens=128, repeats=5):
ttft, per_token = [], []
for _ in range(repeats):
t0 = time.perf_counter(); first = None
for i, _tok in enumerate(generate(prompt, max_new_tokens=n_tokens, stream=True)):
if i == 0:
first = time.perf_counter() - t0
total = time.perf_counter() - t0
ttft.append(first); per_token.append((total - first) / (n_tokens - 1))
return {"ttft_ms": round(np.median(ttft) * 1000, 1),
"ms_per_token": round(np.median(per_token) * 1000, 2),
"tokens_per_s": round(1 / np.median(per_token), 1)}
print(measure_inference(generate, "Explain attention briefly."))
# {'ttft_ms': 180.4, 'ms_per_token': 22.1, 'tokens_per_s': 45.2}
Always separate TTFT from the time per token in the report — they have different causes and different remedies.
Mastery means
- Measures latency and throughput separately
- Uses batching, the KV cache and quantisation correctly
- Knows what the bottleneck is in each phase
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Efficient Memory Management for LLM Serving with PagedAttention (vLLM) — arXiv (open access; licence per article)
- arXiv — Fast Inference from Transformers via Speculative Decoding — arXiv (open access; licence per article)