Inference optimization (4): speculative decoding and multi-token decoding
In this post we look at techniques to generate several tokens per forward pass, to remove the overheads around the math, and looks at architectures that are built for fast inference. As before, all the code is in the repository of this post.
Speculative decoding: more than one token per forward pass
The idea
Decode is memory-bound: a forward pass over one token costs about the same as a forward pass over a handful of tokens, because the time goes into reading the weights. Speculative decoding exploits this. A cheap method drafts \(K\) tokens, and the large target model checks all of them in a single forward pass. Every draft token the target agrees with is a token generated almost for free. When the target disagrees, we keep the tokens before the disagreement, take the target’s own token at the disagreement, and draft again. With the right acceptance rule, the output is exactly the same as with the target model alone: speculative decoding is lossless.
If each draft token is accepted with probability \(\alpha\), the expected number of tokens produced per forward pass of the target model is:
\[\frac{1 - \alpha^{K+1}}{1 - \alpha}\]The speedup therefore depends on how often the drafts are right and on how cheap they are. Speculative decoding helps most at small batch sizes, where decode is most memory-bound; at large batch sizes, the GPU is already busy, and the extra verification work competes with useful work. The methods below differ in where the drafts come from.
Draft-model speculative decoding
The original method (Leviathan et al., Chen et al.) drafts with a small language model that uses the same tokenizer as the target, ideally a distilled version of the target, such as the student we train in the distillation section. Each iteration:
- the draft model generates \(K\) tokens autoregressively;
- the target model runs one forward pass over the context and the \(K\) draft tokens, which gives its distribution at each of the \(K\) draft positions, plus one more after the last draft;
- the draft tokens are checked from left to right.
The acceptance rule depends on how we decode:
- With greedy decoding, draft tokens are accepted while they are equal to the target’s most likely token. At the first mismatch, the draft token is replaced by the target’s most likely token, and the remaining drafts are discarded. If all \(K\) drafts are accepted, the target’s prediction after the last draft is appended as a bonus token. The output is identical to greedy decoding with the target alone.
- With sampling (speculative sampling), the draft model samples each token \(x\) from its distribution \(p\), and the target accepts it with probability \(\min\left(1, q(x) / p(x)\right)\), where \(q\) is the target’s distribution. On rejection, a replacement token is sampled from the normalized residual distribution \(\max(0, q - p)\), and the remaining drafts are discarded. This modified rejection sampling produces tokens distributed exactly as \(q\). It requires the drafts to be sampled from \(p\): mixing greedy drafts with this probabilistic rule recovers neither greedy decoding nor sampling.
In a batch, each sequence accepts a different number of tokens, so the sequences progress at different speeds, and engines track their lengths separately, as in continuous batching.
To keep it simple, our implementation generates one sequence at a time, and the whole sequence must fit in the context window. Both models use a KV cache: after each verification, the cache of the target is truncated to the accepted tokens, and the draft model reuses the part of its own cache that is still valid.
The following code implements greedy speculative decoding in GPTlite, with the distilled model as the draft. The full code is in main_speculative_decoding.py, which checks that the output is the same as greedy decoding with the target alone:
Show code
def generate_speculative(model, prompt, n_tokens, propose, n_drafts):
""" Greedy speculative decoding of a single sequence (prompt of shape [1, S]). At every iteration,
propose(tokens, k) returns up to k draft tokens, and the target model checks all of them in a single
forward pass. The cache of the target holds all the accepted tokens except the last one, which is
processed together with the drafts. The whole sequence must fit in the context window. """
assert prompt.size(1) + n_tokens <= model.seqlen, "the sequence must fit in the context window"
tokens, kv_cache = prompt, None
if prompt.size(1) > 1:
_, kv_cache = model(prompt[:, :-1])
stats = {"target forward passes": 0, "proposed drafts": 0, "accepted drafts": 0}
while tokens.size(1) < prompt.size(1) + n_tokens:
drafts = propose(tokens, min(n_drafts, model.seqlen - tokens.size(1))) # [1, k]
# target predictions after the last accepted token and after each draft: [1, k+1]
logits, kv_cache = model(torch.cat([tokens[:, -1:], drafts], dim=1), kv_cache=kv_cache)
predictions = torch.argmax(logits, dim=-1)
n_accepted = 0 # accept the drafts while they match the predictions of the target
while n_accepted < drafts.size(1) and drafts[0, n_accepted] == predictions[0, n_accepted]:
n_accepted += 1
# keep the accepted drafts, plus the target's own token at the first mismatch (or the bonus token, after
# the last draft if all were accepted), and discard the cache entries of the rejected drafts
tokens = torch.cat([tokens, drafts[:, :n_accepted], predictions[:, n_accepted:n_accepted+1]], dim=1)
kv_cache = truncate_cache(kv_cache, tokens.size(1) - 1)
stats["target forward passes"] += 1
stats["proposed drafts"] += drafts.size(1)
stats["accepted drafts"] += n_accepted
return tokens[:, :prompt.size(1) + n_tokens], stats
class DraftModel:
""" Proposes draft tokens with a small model and its own KV cache. On every call, it reuses the cache
of the longest prefix of the accepted tokens it has already processed """
def __init__(self, model):
self.model = model
self.kv_cache, self.processed = None, None # cache and the tokens it holds
def __call__(self, tokens, n_drafts):
if n_drafts == 0:
return tokens[:, :0]
n_cached = 0
if self.processed is not None: # number of leading tokens processed before and still valid
n = min(self.processed.size(1), tokens.size(1) - 1)
mismatches = (self.processed[0, :n] != tokens[0, :n]).nonzero()
n_cached = mismatches[0].item() if len(mismatches) > 0 else n
kv_cache = truncate_cache(self.kv_cache, n_cached) if n_cached > 0 else None
new_tokens, drafts = tokens[:, n_cached:], []
for _ in range(n_drafts): # greedy generation with the draft model
logits, kv_cache = self.model(new_tokens, kv_cache=kv_cache)
new_tokens = torch.argmax(logits[:, -1:], dim=-1) # [1, 1]
drafts.append(new_tokens)
drafts = torch.cat(drafts, dim=1)
self.kv_cache, self.processed = kv_cache, torch.cat([tokens, drafts[:, :-1]], dim=1) # the last draft was not processed
return drafts
Prompt lookup decoding
Prompt lookup decoding needs no draft model at all. It takes the last few generated tokens (an n-gram), searches for the same n-gram earlier in the prompt or in the generated text, and proposes the tokens that followed it as drafts; the verification is the same as above. It costs almost nothing, and gives large speedups when the output copies from the input, as in summarization, code editing, question answering over documents, or rewriting a previous answer. Related methods draft from other sources: lookahead decoding collects n-grams from Jacobi iterations of the model itself, and REST retrieves continuations from a datastore of text.
The following code implements prompt lookup decoding, reusing the verification loop of speculative decoding with a different source of drafts. The full code is in main_prompt_lookup.py:
Show code
def prompt_lookup(tokens, n_drafts, max_ngram=3):
""" Drafts the tokens that followed the most recent earlier occurrence of the last n-gram of the sequence,
trying the longest n-grams first. Returns no drafts if the n-gram never occurred before """
sequence = tokens[0]
for n in range(max_ngram, 0, -1):
if sequence.size(0) <= n:
continue
ngram = sequence[-n:]
windows = sequence[:-1].unfold(0, n, 1) # all the earlier n-grams of the sequence
matches = (windows == ngram).all(dim=1).nonzero()
if len(matches) > 0:
start = matches[-1].item() + n # first token after the most recent match
return sequence[start:start + n_drafts].unsqueeze(0)
return tokens[:, :0]
# usage: the same verification loop as speculative decoding, with a different source of drafts
tokens, stats = generate_speculative(model, prompt, n_tokens, prompt_lookup, n_drafts)
Medusa
Medusa adds a few extra decoding heads on top of the last hidden state of the target model: the original output head predicts the next token, and head \(i\) predicts the token \(i+1\) positions ahead. The top candidates of each head are combined into a tree of possible continuations, and all of them are verified in a single forward pass with tree attention, a mask that lets each candidate attend only to its own ancestors in the tree. The longest accepted path is kept. Training only the heads, with the model frozen (Medusa-1), gives speedups above 2.2 times; training the heads together with the model (Medusa-2) reaches 2.3 to 3.6 times. Medusa needs no separate draft model, but each head guesses its token independently, without seeing the tokens guessed by the other heads, so the accuracy drops quickly for later positions. For sampling, Medusa also proposes a faster typical acceptance rule, which is not exactly lossless.
EAGLE
EAGLE drafts at the level of features rather than tokens. Its draft model is a single lightweight transformer layer that takes the top-layer hidden states (features) of the target, together with the embeddings of the tokens sampled so far, predicts the next feature, and turns it into a token with the target’s own output layer. Feature sequences are more regular than token sequences, so these drafts are much more accurate than Medusa’s independent heads. EAGLE-2 shapes the draft tree dynamically, expanding the branches where the draft model is confident. EAGLE-3 (NeurIPS 2025) noticed that EAGLE barely improved with more training data, and traced this to its feature prediction objective. It drops feature prediction and predicts tokens directly, fuses low-, mid- and high-level features of the target instead of the top layer only, and uses training-time test: during training, the draft model is fed its own predictions, to simulate multi-step drafting. EAGLE-3 reaches speedups of up to 6.5 times, about 1.4 times more than EAGLE-2, and 1.38 times higher throughput at a batch size of 64 in SGLang. EAGLE-style drafting is supported by vLLM, SGLang and TensorRT-LLM.
Self-speculative decoding: LayerSkip
LayerSkip uses the early layers of the target model as the draft. The model is trained with layer dropout, which skips the later layers more often, and with an early exit loss, which trains every layer to make good predictions through the shared output head. At inference, the draft exits after the first \(E\) layers, and the verification runs only the remaining layers, reusing the computation and the cache of the first \(E\). There is no second model to store, and the reported speedups reach up to about 2 times on summarization and coding tasks.
DFlash
DFlash (Chen, Liang and Liu, UC San Diego, ICML 2026) replaces the autoregressive drafter with a small block diffusion model. Autoregressive drafters such as EAGLE-3 need one forward pass per drafted token, so they are kept very shallow (a single transformer layer for EAGLE-3), which, according to the authors, limits their acceptance length and keeps practical speedups around 2-3 times. DFlash instead takes the last verified token, appends 15 mask tokens, and fills the whole block in a single forward pass, so its drafting cost barely grows with the block size and it can afford a deeper drafter: a 5-layer DFlash fills a 16-token block faster than a 1-layer EAGLE-3 drafts 8 tokens, and gets more of them accepted. A drafter this small predicts poorly on its own (without target features, a 5-layer diffusion drafter only reaches about 3 times), so DFlash conditions it on the target: the hidden states of 5 target layers, spread from shallow to deep, are concatenated, projected, and injected into the keys and values of every draft layer, instead of only into the drafter’s input as in EAGLE-3, so the acceptance length keeps improving as draft layers are added. The drafter reuses the target’s frozen embedding and output head, and only its transformer layers are trained, on responses generated by the target: each training block starts at a randomly sampled anchor token, as at inference, all blocks of a sequence are trained in one pass with a block-sparse attention mask, and the loss favors early positions, since one early mistake invalidates the rest of the block. Verification is unchanged, so the output stays lossless. On Qwen3-8B with greedy decoding, the paper reports about 6.5 tokens committed per verification step on average and up to 6.1 times faster decoding than the baseline, about 2.5 times the speedup of EAGLE-3, with the largest gains on math and code and the smallest on open-ended chat (about 2.8 times). In SGLang on a B200, the speedup drops from about 4-5 times for a single request to about 2.5-3 times at 32 concurrent requests, as verifying 16-token blocks becomes compute-bound. Drafters trained on 4K-token contexts also lose acceptance on longer inputs, which a short long-context fine-tune recovers. Unlike MTP, which only helps models pre-trained with MTP layers, a DFlash drafter is trained after the fact against a frozen target, so it can be added to almost any model: z-lab releases drafters for Qwen3.5/3.6, gpt-oss, Gemma 4 and Kimi K2.5, among others, and DFlash runs in vLLM and SGLang, with an MLX version for Apple Silicon.
Multi-token prediction
Multi-token prediction (MTP) trains the model itself, during pre-training, to predict several future tokens. It includes training objective where the model predicts several future tokens, creating a denser learning signal — and doubling as a speculative-decoding drafter at inference.
- Gloeckle et al. (Meta, 2024) add \(n\) independent output heads on a shared trunk, each predicting one of the next \(n\) tokens. Besides improving sample efficiency, notably on code, the extra heads serve as drafts for self-speculative decoding, making the inference of a model trained to predict 4 tokens up to 3 times faster.
- DeepSeek-V3 uses sequential MTP modules: each module combines the representation of the previous depth with the embedding of the next token, and runs one transformer block, so every prediction stays conditioned on the previous ones. MTP is an extra training objective, and at inference the modules can either be discarded or used as drafts: the second predicted token is accepted 85-90% of the time, giving 1.8 times more tokens per second.
- Sequential / cascaded modules — NOT parallel towers. This is the defining detail. The main model predicts $t_{i+1}$ from the shared trunk. Then MTP module 1 (a single lightweight Transformer layer) takes the trunk’s hidden states plus the embedding of the just-predicted token to predict $t_{i+2}$; MTP module 2 chains off module 1 to predict $t_{i+3}$; and so on. Each module depends on the previous one, preserving the full causal chain. (Gloeckle et al.’s original MTP uses parallel independent heads off one trunk — V3 deliberately chose the sequential version to model inter-token dependencies.)
- Shared embedding & output head. Every MTP module reuses the main model’s embedding and unembedding/output head — only the small per-module Transformer layer is extra. Open-source V3 ships 1 MTP module (depth = 1).
- Auxiliary training loss. Each module’s future-token prediction carries its own cross-entropy loss, added (weighted) to the main loss. Penalizing bad “look-ahead” guesses forces the trunk to build representations useful for longer-range planning — helping coherence in reasoning, code, and math.
- Inference: default single-token, optional speculative decoding. By default only the main head runs (max accuracy). But the MTP modules can act as a built-in draft model: they propose the next few tokens, and the main model verifies them in parallel in one forward pass. V3’s MTP-1 has an ~80–90% acceptance rate, giving ~1.8× higher generation throughput.
Medusa and EAGLE add drafting to an existing model, while MTP builds it in during pre-training.
Summary of speculative decoding methods
| Method | Source of the drafts | Extra training | Extra memory | Reported speedup |
|---|---|---|---|---|
| Draft model | a separate small model | none, or distilling a draft | a second model | about 2-3x |
| Prompt lookup | n-grams of the context | none | none | large on copy-heavy tasks |
| Medusa | extra heads | the heads (or joint fine-tuning) | small | 2.2-3.6x |
| EAGLE-3 | a feature-level draft layer | the draft layer | small | up to 6.5x |
| LayerSkip | early layers of the model | special fine-tuning | none | up to about 2x |
| DFlash | a block-diffusion drafter conditioned on the target’s hidden states | the drafter (target frozen) | small (a few layers) | over 6x, about 2.5x EAGLE-3’s speedup |
| Multi-token prediction | MTP heads or modules | pre-training | small | 1.8x (DeepSeek-V3) to 3x |
The reported speedups come from the respective papers, with different models, tasks and hardware, and they shrink as the batch size grows.