KV cache and speculative decoding
Generating text is memory-bound, not compute-bound - the GPU spends most of its time waiting on weights, which is why you can get several tokens for the price of one.
The surprising thing about running a language model is that generating text barely uses the GPU’s compute at all. Producing one token requires reading every weight in the model out of memory, and then doing a trivial amount of arithmetic with each. The hardware spends its time waiting on memory bandwidth.
Everything interesting about inference optimisation follows from that one fact.
The KV cache
Attention at position n needs the keys and values of every position before it. Computed
naively, generating a sequence of length n recomputes those projections O(n²) times in
total, and all but the newest are identical to what was computed the step before.
So you keep them. The KV cache stores keys and values per layer per position, turning each step into “project the new token, append, attend over the cache”. The recomputation disappears.
What you pay is memory, and it is not a small amount: the cache grows linearly with sequence length and batch size, and for long contexts it can exceed the size of the model weights. Serving systems live or die on how they manage it, which is why vLLM’s PagedAttention - treating the cache as pages rather than one contiguous block - was such a step change.[3]paperEfficient Memory Management for Large Language Model Serving with PagedAttention[4]source codevLLM - the paged KV cache implementation
Press Evict the KV cache in the widget. Every position turns red: all of that work has to be redone before generation can continue.
Why speculation is free
Because decoding is memory-bound, a forward pass that verifies five tokens costs almost exactly what a forward pass that produces one token costs. The weights get read either way; the extra arithmetic is nearly free.
Speculative decoding exploits this.[1]paperFast Inference from Transformers via Speculative Decoding A small, cheap draft model
generates k tokens. The large model then runs a single forward pass over all of them and
checks, at every position, whether it would have chosen the same token.
Everything up to the first disagreement is kept. The rest is thrown away, because those tokens were conditioned on a prefix the target model rejected.
The essential property is that the output distribution is unchanged. This is not an approximation or a quality trade - with the right acceptance rule, the result is distributed identically to ordinary sampling from the large model.[2]paperAccelerating Large Language Model Decoding with Speculative Sampling
The formula
With per-token acceptance rate α and draft length k:
E[tokens per forward pass] = (1 - α^(k+1)) / (1 - α)
Two things worth reading out of it.
It is never worse than 1. Even when every draft token is rejected, the verification pass itself still produces one correct token - the bonus token. Speculation cannot lose. Drag α down to 0.1 in the widget and the rate approaches 1 without ever going below it.
Returns diminish quickly in k. Each additional draft token is only reached if all
previous ones were accepted, so its contribution is discounted by α^k. At α = 0.7, going
from k=1 to k=2 buys a lot; going from k=7 to k=8 buys almost nothing while costing a full
extra draft-model run every time.
What this changes about serving
Once you accept that decoding is memory-bound, a whole family of techniques follows: batch more requests so the weight reads are amortised across them, page the KV cache so memory does not fragment, and speculate so each expensive pass yields more than one token.
They compose. They also all compete for the same scarce resource, which is memory - a bigger batch means more KV cache, and more cache means fewer concurrent sequences. That is the real tuning problem in production inference, and it is a memory budget problem wearing a throughput problem’s clothes.
The dial
The line - when someone asks
Transformer decoding recomputes attention over every previous token at each step, so the KV cache stores past keys and values and turns that quadratic repetition into a linear append. Because generating one token then reads the entire model's weights for a single token of output, decoding is memory-bandwidth-bound and the GPU is mostly idle - so speculative decoding has a small draft model guess k tokens ahead and the big model verify them all in one pass. Accepted guesses are free, and the maths gives (1-α^(k+1))/(1-α) tokens per pass, which is never worse than 1.
Recall
Loading…
Where are you with this?
Saved on this device. Sign in to keep it across devices.
Sources
Primary
- [1]
- [2]
- [3]
Secondary
- [4]