Serving · Decoding

Speculative Decoding

draft with a small model, verify with the large one

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

  1. 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 pass

  2. One 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 target

  3. A 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 text

  4. When 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))

  5. 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

Go deeper