Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

37. Speculative decoding in the server

In this chapter

  • Moving speculation from one request (Chapter 26) into the batched engine: drafts become scheduled tokens, verification is one more kind of logits row, and rollback is a change of num_computed_tokens.
  • Three drafters with different costs: n-gram prompt lookup, a draft model whose KV pool shares the target's block tables, and Medusa heads that read the target's hidden state.
  • Rejection sampling with deterministic drafts, and a statistical test that the output distribution is exactly the target's.
  • When speculation pays in a server and when it doesn't: batch size, acceptance and the memory-bound regime.

You will build

NgramDrafter.propose_one, DraftModelDrafter.propose, verify and SpeculativeStep.step in engine/serve/spec.py.

Time: 6-8 hours. GPU: recommended for the speedups (every test runs on a CPU).

From one request to a batch

Chapter 26 showed why speculation works: a decode step reads every weight to produce one token, so checking $k+1$ tokens in one forward pass costs about the same as checking one, and every accepted draft token is a step saved. Its implementation served one request with a contiguous cache, its own loop and its own truncate. A serving engine can’t run a separate loop per request. The engine core of Chapter 31 already has everything speculation needs:

speculation needsthe engine core has
feed the newest token and $k$ draftsrequest.spec_token_ids; the scheduler asks for num_tokens_with_spec - num_computed_tokens tokens
$k+1$ logits rows per requestbuild_batch emits rows from the last real token onward (Chapter 31’s logits_indices)
blocks for the drafts’ K/Vallocate_slots for those tokens, like any other
roll back rejected tokensset num_computed_tokens to the accepted length, BlockManager.trim the rest

So speculation in the server is a different second half of EngineCore.step: verify instead of sample, roll back, then draft for the next step.

class SpeculativeStep:
    """Replaces the sample-and-update half of EngineCore.step when speculation is on."""

    def __init__(self, drafter):
        self.drafter = drafter
        self.stats = {"proposed": 0, "accepted": 0, "verify_steps": 0}

    def eligible(self, request):
        p = request.params
        return not (p.needs_penalties or p.logit_bias or p.allowed_token_ids is not None or p.logprobs is not None
                    or p.prompt_logprobs is not None or "guide" in request.extra
                    or p.min_tokens > request.num_output_tokens)

    def step(self, engine, plan):
        """Verify every request's scheduled drafts, roll back, then draft for the next step.  (Your engine: Chapter 37)"""
        batch = engine.runner.prepare(plan.scheduled, engine.blocks.req_blocks)
        logits = engine.runner.execute(batch)
        hidden = engine.runner.last_hidden
        if batch.prompt_spans:
            engine.record_prompt_logprobs(batch, logits[batch.num_sample_rows:])
        outputs, now, row, last_rows = [], time.monotonic(), 0, {}
        plain = [(r, k) for (r, _), k in zip(plan.scheduled, batch.sample_counts)
                 if k and not (r.spec_token_ids and r.num_computed_tokens == r.num_tokens - 1)]
        plain_rows = []
        for (request, n), k in zip(plan.scheduled, batch.sample_counts):
            before = request.num_computed_tokens
            request.num_computed_tokens += n
            if not k:
                engine.blocks.cache_full_blocks(request)
                continue
            if request.spec_token_ids and before == request.num_tokens - 1:     # a decode step with drafts
                drafts = request.spec_token_ids[:k - 1]
                new, accepted = verify(logits[row:row + k], drafts, request.params, engine.generators.get(request.request_id))
                self.stats["proposed"] += len(drafts)
                self.stats["accepted"] += accepted
                self.stats["verify_steps"] += 1
                request.spec_token_ids = []
                request.num_computed_tokens = before + 1 + accepted         # newest + accepted drafts
                engine.blocks.trim(request, request.num_computed_tokens + 1)      # free blocks of rejected drafts
                last_rows[request.request_id] = row + accepted
                self.finish_tokens(engine, request, new, now, outputs)
            else:
                plain_rows.append(row)
                last_rows[request.request_id] = row
            row += k
        if plain:                                                   # rows with nothing to verify: normal sampling
            tokens, logprobs = engine.sampler(logits[plain_rows], [r for r, _ in plain], engine.generators)
            for (request, _), new, lp in zip(plain, tokens, logprobs or [None] * len(plain)):
                self.finish_tokens(engine, request, new, now, outputs, lp)
        running = [r for r in engine.scheduler.running if r.request_id in last_rows and not r.status.finished]
        engine.spec_hidden = hidden[[last_rows[r.request_id] for r in running]] if hidden is not None and running else None
        candidates = [r for r in running if self.eligible(r)]
        proposals = self.drafter.propose(engine, candidates) if candidates else {}
        for request in candidates:
            room = min(request.params.max_tokens - request.num_output_tokens, engine.max_model_len - request.num_tokens) - 1
            request.spec_token_ids = proposals.get(request.request_id, [])[:max(room, 0)]
        engine.steps += 1
        return outputs

    def finish_tokens(self, engine, request, new, now, outputs, logprobs=None):
        """Append tokens one at a time so that a stop condition in the middle cuts the rest."""
        emitted = []
        for token in new:
            request.append(token)
            emitted.append(token)
            status = engine.check_stop(request)
            if status is not None:
                engine.scheduler.finish(request, status)
                engine.generators.pop(request.request_id, None)
                break
        request.first_token_time = request.first_token_time or now
        if not request.status.finished:
            engine.blocks.cache_full_blocks(request)
        outputs.append(engine.make_output(request, emitted, logprobs))

Requests that ask for features speculation can’t serve exactly are simply not drafted for: penalties, logit bias, allowed tokens and grammars change the target distribution at each position depending on the tokens before it, and logprobs would have to be reported for accepted drafts as well. They decode one token per step in the same batch, sampled by Chapter 34’s sampler, while their neighbors speculate.

Verification with deterministic drafts

Every drafter in this chapter proposes its single best guess. The proposal distribution $q$ is then a point mass on the draft $d$, and Chapter 26’s rule simplifies:

  • accept $d$ with probability $\min(1, p(d)/q(d)) = p(d)$;
  • on rejection, sample from the residual $\max(p - q, 0)$, which is $p$ with $d$’s entry set to zero, renormalized;
  • if all $k$ drafts are accepted, sample one bonus token from the last row.

The total probability of emitting token $t$ at the first position is $p(d)$ for $t = d$, and $(1 - p(d)) \cdot p(t) / (1 - p(d)) = p(t)$ for every other $t$: exactly $p$. Greedy decoding is the special case “accept while the target’s argmax equals the draft”.

def verify(logits, drafts, params, generator=None):
    """logits [len(drafts) + 1, V] from the target; returns (accepted drafts + one new token, accepted).  (Your engine: Chapter 37)

    Greedy: accept while the target's argmax equals the draft. Sampled: with p the target's
    filtered distribution at each row, accept draft d with probability p(d); on the first
    rejection, draw from p with d removed. If all are accepted, the last row gives a bonus token.
    """
    if params.greedy:
        best = logits.argmax(-1).tolist()
        out = []
        for i, d in enumerate(drafts):
            if best[i] != d:
                return out + [best[i]], len(out)
            out.append(d)
        return out + [best[len(drafts)]], len(drafts)
    x = logits.float() / params.temperature
    if params.top_k or params.top_p < 1 or params.min_p:
        n = x.shape[0]
        x = top_k_top_p_min_p(x, torch.full((n,), params.top_k or x.shape[-1], device=x.device),
                              torch.full((n,), params.top_p, device=x.device), torch.full((n,), params.min_p, device=x.device))
    p = x.softmax(-1)
    out = []
    for i, d in enumerate(drafts):
        if torch.rand((), generator=generator, device=p.device) < p[i, d]:
            out.append(d)
            continue
        residual = p[i].clone()
        residual[d] = 0
        if residual.sum() <= 0:
            residual = p[i]
        return out + [int(torch.multinomial(residual / residual.sum(), 1, generator=generator))], len(out)
    return out + [int(torch.multinomial(p[len(drafts)], 1, generator=generator))], len(drafts)

The target’s distribution $p$ includes the request’s temperature and its top-k, top-p and min-p filters, computed with Chapter 34’s batched filter, so a request gets the same distribution with or without speculation. The test checks this statistically: with a perfect drafter at temperature 1, the empirical distribution of the second generated token over 3,000 samples must be within a total variation distance of 0.07 of its exact law, which the test computes by summing over every possible first token. The test also checks that drafts were both accepted and rejected, so both branches were exercised.

Rollback

After verification the request’s token list holds the accepted drafts and one new token. Its K/V cache is valid through the last accepted draft: the newest token (the correction or bonus) hasn’t been fed yet, which is exactly Chapter 31’s invariant “the newest token is uncomputed”. So:

$$ \texttt{num_computed_tokens} = (\text{before}) + 1 + \text{accepted}, $$

and BlockManager.trim releases blocks that only held rejected drafts. Nothing is copied, and nothing about the rejected positions needs to be erased: those slots will be overwritten before any query can see them, by the same causal rule that has protected stale slots since Chapter 19.

Two bookkeeping rules were found by this chapter’s tests. A request that is preempted while holding drafts must drop them (the scheduler clears spec_token_ids in _preempt); otherwise, when it resumes, recomputing its whole history would be mistaken for a verification step. And a request that stops in the middle of its accepted tokens, because one of them is a stop token, must discard the rest: finish_tokens appends one token at a time and checks check_stop after each.

Drafter 1: n-gram prompt lookup

The cheapest drafter needs no model at all. If the last few tokens of the context appeared earlier, guess that what followed them then follows them now:

class NgramDrafter:
    """Prompt-lookup decoding: find the most recent earlier occurrence of the context's last n
    tokens (longest n first) and propose the k tokens that followed it. Free to run, and very
    effective when the output copies the input: code edits, extraction, RAG answers, chat history.
    """

    def __init__(self, k=4, max_n=4, min_n=1):
        self.k, self.max_n, self.min_n = k, max_n, min_n

    def propose_one(self, tokens, k):
        """(Your engine: Chapter 37)"""
        for n in range(min(self.max_n, len(tokens) - 1), self.min_n - 1, -1):
            suffix = tokens[-n:]
            for start in range(len(tokens) - n - 1, -1, -1):          # most recent match first
                if tokens[start:start + n] == suffix:
                    follow = tokens[start + n:start + n + k]
                    if follow:
                        return follow
        return []

    def propose(self, engine, requests):
        return {r.request_id: self.propose_one(r.token_ids, self.k) for r in requests}

This “prompt lookup decoding” (Saxena, 2023) costs microseconds and is remarkably effective whenever the output copies the input: editing code, extracting fields, answering questions about a document in the prompt (RAG), continuing a conversation that quotes itself. For free-form generation it rarely finds a match, proposes nothing, and costs nothing. Most engines enable it by default for that reason (vLLM’s ngram method).

Drafter 2: a draft model that shares block tables

A small model from the same family (Qwen3-0.6B drafting for Qwen3-32B) guesses well on all kinds of text. It needs its own K/V cache, and that’s where the paged design pays off again: give the draft model a pool with the same number of blocks as the target’s, and every request’s block table addresses both pools. The block manager allocates once, for two models.

class DraftModelDrafter:
    """A small model (same tokenizer) drafts greedily. Its KV pool has the target's shape in blocks,
    so a request's block table addresses both pools: the block manager allocates once for two
    models. The draft's own progress is request.extra["draft_computed"]."""

    def __init__(self, draft_model, engine, k=4):
        self.k = k
        self.flat = draft_model if isinstance(draft_model, FlatModel) else FlatModel(draft_model)
        runner = engine.runner
        self.backend = get_backend(runner.backend.name) if hasattr(runner.backend, "name") else runner.backend
        layers, kv_heads, head_dim = self.flat.kv_spec()
        shape = (runner.num_blocks + 1, runner.block_size, kv_heads, head_dim)
        p = next(self.flat.parameters())
        self.caches = [self.backend.allocate(shape, p.dtype, runner.device) for _ in range(layers)]

    @torch.inference_mode()
    def propose(self, engine, requests):
        """Catch the draft up to each request's newest token, then draft k tokens greedily.  (Your engine: Chapter 37)"""
        live = []
        for r in requests:
            done = r.extra.get("draft_computed", 0)
            if r.extra.get("draft_epoch") != r.num_preemptions:       # its blocks were freed or swapped
                done = 0
            done = min(done, r.num_computed_tokens)
            # Room for the drafts' K/V: positions up to num_tokens + k - 2.
            if engine.blocks.allocate_slots(r, self.k):
                live.append((r, done))
        drafts = {r.request_id: [] for r in requests}
        if not live:
            return drafts
        tables = engine.blocks.req_blocks
        views = [_View(r, done) for r, done in live]
        scheduled = [(v, v.num_tokens - v.num_computed_tokens) for v in views]
        for step in range(self.k):
            batch = build_batch(scheduled, tables, engine.runner.block_size, engine.runner.device)
            hidden = self.flat(batch.input_ids, batch.positions, self.caches, batch.meta, self.backend)
            tokens = self.flat.compute_logits(hidden[batch.logits_indices]).argmax(-1).tolist()
            for (view, n), token in zip(scheduled, tokens):
                view.num_computed_tokens += n
                view.spec_token_ids.append(token)
            scheduled = [(v, 1) for v in views]
        for (r, _), view in zip(live, views):
            drafts[r.request_id] = view.spec_token_ids
            r.extra["draft_computed"] = r.num_tokens + self.k - 1    # draft K/V is valid this far, if accepted
            r.extra["draft_epoch"] = r.num_preemptions
        return drafts

Each step, the drafter first catches up: it runs the draft model over every token the draft hasn’t seen yet (the whole prompt the first time; afterwards, the newest token and any correction), as one flattened batch over all requests. Then it drafts $k$ tokens greedily, one batched forward per token. It tracks its own progress in request.extra["draft_computed"], which is valid only up to the tokens that were actually accepted, and resets when the request is preempted (its blocks were freed or swapped, and only the target’s pool is swapped).

Before drafting, it asks the block manager for room for the drafts’ K/V. If the pool is too full, that request simply gets no drafts this step. Speculation never causes a preemption on its own.

Drafter 3: Medusa heads

The draft model above repeats work the target already did: it re-reads the context and builds its own representation of it. Medusa (Cai et al., 2024) instead adds $k$ small heads to the target. Each reads the target’s final hidden state at the newest position and guesses the token $1, 2, \ldots, k$ places further on, through the target’s own LM head. The heads have no attention and no cache, so drafting is a few matrix-vector products:

class MedusaHeads(nn.Module):
    """k residual heads on the target's last hidden state; head i guesses the token i + 1 places
    after the one the target just sampled. They share the target's LM head."""

    def __init__(self, hidden, k):
        super().__init__()
        self.blocks = nn.ModuleList(nn.Linear(hidden, hidden) for _ in range(k))
        for block in self.blocks:
            nn.init.zeros_(block.weight)
            nn.init.zeros_(block.bias)

    def forward(self, h):                                   # [R, D] -> [k, R, D]
        return torch.stack([h + nn.functional.silu(block(h)) for block in self.blocks])


class MedusaDrafter:
    def __init__(self, heads, k=None):
        self.heads, self.k = heads, k or len(heads.blocks)

    @torch.inference_mode()
    def propose(self, engine, requests):
        hidden = engine.spec_hidden                        # [R, D]: each request's newest accepted position
        if hidden is None or not requests:
            return {r.request_id: [] for r in requests}
        logits = engine.runner.flat.compute_logits(self.heads(hidden)[:self.k])     # [k, R, V]
        guesses = logits.argmax(-1).T.tolist()
        return {r.request_id: g for r, g in zip(requests, guesses)}

The heads must be trained, which is cheap: freeze the target, run it over text (ideally its own outputs, so the heads learn to predict the target), and train each head with cross-entropy against the token $i + 1$ positions ahead. run.py drafters trains three heads for 300 steps on the CPU.

Medusa’s guesses for positions 2 and 3 ignore the tokens in between, which limits its acceptance. EAGLE (Li et al., 2024) and DeepSeek-V3’s multi-token prediction (MTP) modules fix that with one small transformer layer that takes the target’s hidden state and the embedding of the next token, and drafts autoregressively with a cache of its own. These are the most effective drafters in production today, and with this engine they’re a combination of the two drafters above: a draft model whose input embedding is fused with the target’s hidden state (stretch exercise 3).

When speculation pays in a server

Speculation trades compute for memory bandwidth. Verifying $k$ drafts multiplies the step’s tokens, and therefore its FLOPs, by up to $k + 1$, but barely changes its memory traffic. So:

  • At small batch sizes, decode is memory-bound, the extra FLOPs are nearly free, and every accepted token is a step saved. Speedups of 2-3× are common for chat with a good drafter.
  • At large batch sizes, decode is approaching compute-bound (Chapter 24). Extra verification tokens now cost real time, and rejected drafts are wasted compute that other requests could have used. Speculation can make throughput worse.

Production engines therefore turn speculation down as load rises: fewer drafts per request, or none, when the batch is large (stretch exercise 1). The expected tokens per step for acceptance rate $\alpha$ and $k$ drafts is Chapter 26’s $(1 - \alpha^{k+1})/(1 - \alpha)$; whether that’s worth the extra $k$ tokens of compute depends on where the step sits on the roofline.

Tree verification extends this: instead of one chain of $k$ guesses, verify a small tree of alternatives (Medusa’s and EAGLE-2’s top-2 at each position) in one forward pass, with an attention mask that lets each node see only its ancestors. It accepts more tokens per step at the price of more verification compute, so it shines at batch 1 and fades with load. With the reference backend’s allowed mask, it’s a stretch exercise.

Run it

python run.py drafters --requests 8 --slots 8 --new-tokens 64

Eight requests through the 6-layer test model (random weights) on a laptop CPU, with each drafter:

{"drafter": "none", "steps": 64, "tokens_per_step_per_request": 1.0, "acceptance": null, "tok_s": 448.7}
{"drafter": "n-gram, k=2", "steps": 61, "tokens_per_step_per_request": 1.05, "acceptance": 0.521, "tok_s": 437.3}
{"drafter": "n-gram, k=4", "steps": 61, "tokens_per_step_per_request": 1.05, "acceptance": 0.398, "tok_s": 435.0}
{"drafter": "draft model (first 2 of 6 layers), k=3", "steps": 32, "tokens_per_step_per_request": 2.0, "acceptance": 0.479, "tok_s": 443.6}
{"drafter": "draft model = target, k=3", "steps": 17, "tokens_per_step_per_request": 3.76, "acceptance": 1.0, "tok_s": 450.0}
{"drafter": "none (Medusa prompts)", "steps": 64, "acceptance": null, "mean_accepted_per_step": null}
{"drafter": "Medusa, 3 trained heads", "steps": 34, "acceptance": 0.441, "mean_accepted_per_step": 1.3}

Every configuration produced exactly the tokens of plain decoding (the command asserts it). The n-gram drafter rarely finds a repeat in a random model’s output, but when it does, half its guesses are right. A draft model made of the target’s first two layers, an “early exit” draft, halves the number of steps; the target drafting for itself reaches 3.76 tokens per step, the ceiling for $k = 3$ minus the requests’ final steps. Three Medusa heads, trained for 300 steps on the target’s own outputs, also halve the steps.

The last column is the honest one: on a CPU, tokens per second didn’t move. A CPU forward pass is compute-bound, so verifying four tokens costs about four times as much as decoding one, cancelling the saved steps. On a GPU, where the same decode step is memory-bound, the steps saved translate into time saved; measure it there with this command and compare it with the formula above.

Build it

Engine milestone 37: speculation in the engine. Implement NgramDrafter.propose_one, DraftModelDrafter.propose, verify and SpeculativeStep.step in engine/serve/spec.py (MedusaHeads, MedusaDrafter, finish_tokens and enable_speculation are provided). Turn it on with spec.enable_speculation(engine, drafter).

pytest tests/test_ch37_speculative.py
python run.py drafters --impl engine

The tests check n-gram lookup, greedy verification, the sampled verification law against the target distribution, n-gram speculation equal to plain decoding with and without preemption and prefix caching, a perfect draft model accepting every draft, an early-exit draft model under preemption, Medusa heads, a stop token inside accepted drafts, and the distribution of sampled speculative output against its exact law.

Stretch exercises

  1. ★★ Make $k$ adaptive: track each request’s recent acceptance rate and the batch size, and choose $k \in {0, 1, 2, 4}$ to maximize expected tokens per unit of step time, using the roofline model of Chapter 10 for the step’s cost. Where: draft-length selection in SpeculativeStep in engine/serve/spec.py.
  2. ★★ Capture CUDA graphs for verification batches: every decoding request schedules exactly $1 + k$ tokens, so the batch is uniform, and Chapter 33’s buckets apply with $(1 + k) \times$ the rows. Where: verification execution in engine/serve/spec.py, with graph buckets in engine/serve/graphs.py.
  3. ★★★ Implement an EAGLE-style drafter: one Qwen3 layer whose input is fc(concat(embed(token), target_hidden)), with its own KV pool sharing the target’s block tables. Train it on the target’s outputs and compare acceptance with Medusa’s. Where: add a drafter in engine/serve/spec.py; expose target hidden rows in engine/serve/model.py and train it in experiments/ch37.py (create it).
  4. ★★★ Tree verification: verify a tree of drafts (the top 2 at each of 3 positions) with the reference backend’s allowed mask, and accept the longest path the target agrees with. Where: tree proposals/verification in engine/serve/spec.py, with new tree-mask metadata in engine/serve/batch.py and mask handling in engine/serve/attention.py.

Check your understanding

  1. Which four mechanisms of the Chapter 31 engine make batched speculation a small change?
  2. Why does a deterministic drafter make the acceptance probability simply $p(d)$, and why does the output still follow $p$ exactly?
  3. After accepting 2 of 4 drafts, what is the request’s new num_computed_tokens, and which blocks can be freed?
  4. Why must a preempted request drop its drafts?
  5. How can a draft model use the target’s block tables without the block manager allocating twice?
  6. Why can speculation reduce throughput at large batch sizes even with a 70% acceptance rate?

Going deeper

  • Leviathan et al., Fast Inference from Transformers via Speculative Decoding (ICML 2023) and Chen et al., Accelerating Large Language Model Decoding with Speculative Sampling (2023), for the rejection rule.
  • Cai et al., Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads (2024); Li et al., EAGLE: Speculative Sampling Requires Rethinking Feature Uncertainty (2024) and EAGLE-2 (2024); DeepSeek-AI, DeepSeek-V3 Technical Report (2024), §2.2 on multi-token prediction.
  • Apoorv Saxena, Prompt Lookup Decoding (2023); Liu et al., Optimizing Speculative Decoding for Serving Large Language Models Using Goodput (2024), for adapting speculation to load.
  • vLLM’s vllm/v1/spec_decode/ (n-gram, EAGLE, Medusa, MTP proposers) and vllm/v1/sample/rejection_sampler.py.