Serving · Decoding
- target
- GPT-2 small 12 layers
- draft
- distilgpt2 6 layers
- guesses per round
- 4
- this run
- 12 tokens 5 target passes
- output
- identical to GPT-2
All steps
Guess cheaply, check in one pass
A large model’s decode step is limited by reading its weights, so checking five tokens in one pass costs about as much as producing one. Speculative decoding uses that: a small draft model guesses the next few tokens cheaply, and the large model verifies all the guesses at once.
draft k tokens · verify in 1 passOne round
A real round: distilgpt2 (6 layers) guesses 4 tokens one after another; GPT-2 small (12 layers) runs once over all of them and writes down its own choice at each position. Guesses are kept up to the first disagreement, and GPT-2’s choice there is added, so every round gains at least one token.
keep the agreeing prefix + 1 token from the targetA whole run
The whole run: 12 new tokens from 5 passes of GPT-2 instead of 12, and the text is exactly what GPT-2 alone would have written greedily, because every token was checked by GPT-2.
12 tokens · 5 target passes · same textWhen sampling
When sampling instead of picking the top token, a drafted token x is accepted with probability min(1, p(x)/q(x)), where p is the target’s probability and q the draft’s. If rejected, the replacement is drawn from what p has in excess of q. The tokens that come out follow the target’s distribution exactly.
accept x with probability min(1, p(x) / q(x))How much faster
If each guess is accepted with probability α, one target pass yields (1 − α^(k+1)) / (1 − α) tokens on average. The draft must also be cheap: distilgpt2, half of GPT-2’s size, makes this run slower; real systems use drafts 10 to 100 times smaller.
tokens per pass = (1 − α^(k+1)) / (1 − α)
Code
# one round of greedy speculative decoding
draft = []
for _ in range(k): draft.append(small(ctx + draft).argmax()) # k cheap passes
picks = large(ctx + draft).argmax(-1)[-k-1:] # one pass checks all k
n = 0
while n < k and draft[n] == picks[n]: n += 1
ctx += draft[:n] + [picks[n]] # at least one new token
# sampling: accept x with probability min(1, p(x) / q(x))
if random() < min(1, p[x] / q[x]): keep(x)
else: x = sample(normalize(max(0, p - q))); stop