Medusa and speculative decoding: getting more tokens per forward pass
Autoregressive decoding pays for a full pass over the weights to produce one token. Speculative decoding and now Medusa let a model verify several draft tokens in one pass. A walkthrough of how the trick works and where the speedup actually appears.
Why one token per pass is the wrong unit
Generating a token from a large model means reading every parameter out of high bandwidth memory into the accelerator's compute units and then, for a single sequence, using each parameter roughly once. The arithmetic is trivial next to the memory traffic. At batch size one the accelerator spends most of its time waiting on bandwidth, and decoding K tokens costs K of those waits, one after another, because token t+1 cannot be computed until token t exists.
The observation behind both papers we want to walk through is that the wait is the expensive part and the arithmetic is nearly free. If you already have several candidate tokens in hand, checking all of them in a single forward pass costs about the same as generating one, because the weights are read once either way. The question becomes how to get plausible candidates cheaply, and how to accept them without changing what the model would have said.
Speculative decoding, the original version
Yaniv Leviathan, Matan Kalman and Yossi Matias at Google published the approach in late 2022 and presented it at ICML 2023. A small draft model proposes gamma tokens autoregressively. The large target model then runs one pass over the prompt plus all gamma drafts, producing its own next-token distribution at every position. The drafts are walked left to right. A draft token x sampled from the draft distribution q is accepted outright if q(x) is at most p(x) under the target, and otherwise rejected with probability 1 minus p(x)/q(x), after which a replacement is sampled from an adjusted distribution. The first rejection ends the walk.
The elegant part is that this procedure is exactly equivalent to sampling from the target model. The output distribution is unchanged, so the speedup is free in quality terms. The expected number of tokens accepted per pass is a closed form in the per-token acceptance rate alpha, which the paper defines as one minus an expected divergence between draft and target. When alpha is high and gamma is a few tokens, each target pass yields two or three tokens instead of one.
On T5-XXL at 11 billion parameters, using T5-small at 77 million as the draft, the paper reports 2.6x on WMT English to German at temperature one and 3.4x greedy, and 2.3x and 3.1x on CNN/DailyMail summarisation. Those numbers come with the memory bandwidth assumption stated up front. If the hardware is compute bound, running the extra positions is no longer free and the gains shrink.
What Medusa changes
The Medusa paper from Tianle Cai, Tri Dao and colleagues, out this month, keeps the verify-in-one-pass idea and throws away the draft model. Their complaint is practical. Finding a small model whose distribution tracks the large one well enough is hard, it has to share a tokenizer, and shipping a second model inside a distributed serving stack is its own engineering project. Instead they bolt extra decoding heads onto the backbone. Each head is a single feed forward layer with a residual connection sitting on the final hidden state, and head k is trained to predict the token k positions ahead. The paper suggests five heads is about the most that helps.
Because each head produces a distribution over its position rather than a single guess, Medusa takes the top few tokens from each head and builds a tree of candidate continuations. A custom attention mask lets the backbone verify the whole tree in one pass, with each candidate token attending only to its own ancestors. That gives many candidate sequences for the price of one forward pass, which raises the odds that at least one long prefix is accepted.
Acceptance is the other departure. Rather than the exact rejection sampling rule, Medusa uses what the authors call typical acceptance. A candidate token is accepted if its probability under the backbone clears a threshold derived from the entropy of the backbone's distribution at that position. This is no longer distribution preserving in the strict sense. The authors argue that in practice it produces text of the same quality as the backbone while accepting more tokens, and they present it as a trade rather than hide it.
The two training recipes and the numbers
Medusa-1 freezes the backbone and trains only the heads. The paper reports that on Vicuna 7B this takes about five hours on one A100 with 60,000 ShareGPT samples, which puts it within reach of anyone serving an open model. Since the backbone is untouched, the outputs are the backbone's outputs, and the speedup is 2.18x on Vicuna 7B and 2.33x on 13B.
Medusa-2 trains the heads and the backbone together with a recipe designed to keep the backbone's capabilities intact. It gets to 2.83x on both Vicuna sizes, 2.35x on the 33B model, and 2.66x on Zephyr 7B. The paper also includes a self-distillation route for cases where the original fine-tuning data is unavailable, which is the usual situation with open chat models.
One caveat sits in the experimental setup and should be read before anyone plans a deployment around these figures. All the reported speedups are at batch size one, which the authors describe as the locally hosted, personal use setting. They say the ideas generalise to larger batches but do not show it. At high batch sizes the accelerator is already busy with arithmetic, the spare compute that speculative verification exploits disappears, and a 2.8x single stream speedup can turn into very little throughput gain.
Where the gain actually shows up
The way to think about this class of technique is as a conversion of idle compute into lower per-request latency. That matters most when a single user is waiting on a single stream, which describes local inference, interactive coding assistants, and anything where time to last token is the metric. It matters least for offline batch jobs where the serving stack already packs enough sequences to saturate the chip, because there is no idle compute to convert.
What we would like to see next is a careful comparison of the two acceptance rules on the same backbone under the same tree, so the gain from typical acceptance can be separated from the gain from multiple heads. We would also like the batch size curve. Until someone publishes throughput at batch sizes of 8, 32 and 128, the practical advice is that Medusa-1 is cheap enough to try on any open model you serve interactively, and that you should expect the headline number to shrink as your batch grows.
Sources
From the foundation