The bottleneck Medusa attacks
Autoregressive decoding is memory-bandwidth-bound: to emit one token the model streams every weight (and the KV cache) from memory, does a tiny amount of arithmetic, and repeats. The GPU’s compute units sit almost idle — a single-token step’s arithmetic intensity is far below the hardware’s FLOPs-per-byte ratio. So generating n tokens costs roughly n full weight reloads, and latency scales linearly with output length no matter how much spare compute exists.
Every speculative method exploits the same slack: a forward pass that verifies k tokens at once costs barely more than a pass that emits one, because both are dominated by loading the same weights — the extra tokens ride along nearly free on unused compute. If you can cheaply guess the next few tokens and verify them in parallel, you amortize one weight reload over several accepted tokens. Medusa’s contribution is a guessing mechanism that needs no separate model: extra heads on the model you already have.
The Medusa head: math and shapes
Let h_t ∈ R^d be the backbone’s final hidden state at position t (the vector its own LM head would turn into the next-token distribution). Medusa attaches K heads. Head k is a single residual block feeding an unembedding:
p_t^(k) = softmax( W2_k · ( SiLU(W1_k · h_t) + h_t ) )
h_t : [d] backbone final hidden state
W1_k : [d, d] residual projection (init ≈ 0)
W2_k : [V, d] unembedding (init = base LM head)
p_t^(k): [V] distribution over the (k+1)-th future tokenTwo initializations make this work. W1_k starts near zero, so initially SiLU(0)+h_t ≈ h_t and the head behaves like the original LM head; W2_k copies the base model’s unembedding, reusing its token geometry. Only the heads train — the backbone stays frozen, so training is cheap (hours, one GPU) and cannot regress the base model. Head k predicts the token at t+k+1; with the backbone’s own t+1 prediction, one pass drafts up to K+1 positions. Crucially the heads are conditionally independent given h_t — head 2 never sees head 1’s guess — Medusa’s main quality limitation, addressed structurally by the tree.
From heads to a candidate tree
A point prediction per head would be brittle: if head 1 is wrong, everything downstream is wasted. Instead Medusa keeps the top-s_k tokens from each head and expands their combinations into a tree. With per-head widths (s_1, …, s_K) the Cartesian product has ∏_k s_k leaf continuations — e.g. (4,3,2,2) gives 48 candidate paths from a single hidden state.
The tree shares prefixes: continuations that begin with the same head-1 token share that node, those agreeing on heads 1–2 share those two, and so on, so the tree has far fewer nodes than leaves×depth — each token appears once. In practice you skip the full product and keep only the highest-probability nodes (ranked by the product of head confidences) up to a fixed budget, since low-probability branches are almost never accepted and only cost compute. The shape is fixed ahead of time so the attention mask can be precomputed.
Tree attention: the mask math
All tree nodes are flattened into one input sequence and run through the backbone in a single forward pass. The trick is the attention mask M: a node attends to the prompt and to its own ancestors along its path from the root — and to nothing on sibling branches. Each node’s hidden state is then computed as if only its own candidate prefix existed, so many mutually-exclusive continuations are scored at once without contaminating each other.
M[i, j] = 0 if node j is an ancestor of node i (or i itself)
M[i, j] = -∞ otherwise
attn_i = softmax( (q_i K^T + M[i]) / sqrt(d_k) ) VOrdinary causal decoding uses a lower-triangular mask — a linear chain where every position attends to all earlier ones. Tree attention generalizes that to a partial order: still ‘attend to your ancestors,’ but ancestry follows tree edges, not sequence position, so two cousins on different branches sit adjacent yet cannot see each other. Because the shape is fixed, M and the positional indices (each node’s position id = prompt length + its depth) are precomputed once. A tree of m nodes thus verifies ∏_k s_k candidate paths at the cost of one m-token forward pass — the core efficiency of Medusa.
Typical acceptance, not exact match
Classic speculative decoding accepts draft tokens by a rejection rule that provably reproduces the target’s exact sampling distribution — elegant but conservative, and it needs the draft’s probabilities. Medusa instead uses typical acceptance, a threshold on the backbone’s own probability for the proposed token given the accepted prefix:
accept x iff p_base(x | prefix) > min( ε, δ · exp(-H(p_base)) )
H(p) = -Σ_x p(x) log p(x) entropy of the base next-token dist
ε, δ : tunable constants (e.g. ε=0.09, δ=0.3)The threshold adapts to uncertainty. When the model is confident (low entropy) the bar is high, so only near-certain tokens pass; when it is genuinely uncertain (high entropy, many acceptable words) the bar drops and more candidates qualify. You walk each tree path from the root, accept the longest prefix whose every token clears its threshold, then take the model’s own next token as a free bonus at the first rejection. This does not match the base sampling distribution token-for-token — it is a quality-preserving heuristic that at temperature 0 collapses to exact greedy matching. Dropping the exactness guarantee is what lets Medusa accept longer prefixes and go faster.
Expected speedup, with a worked example
Let τ be the mean accepted length — tokens confirmed per verification pass (the accepted prefix plus the bonus token). Without Medusa, τ tokens need τ passes. With Medusa they need one pass, slightly costlier because it processes m tree tokens and runs K heads. With per-step overhead factor c ≥ 1:
speedup ≈ τ / c
τ : accepted tokens per pass (typically 2.3 – 3.6)
c : overhead of wider pass + heads (typically 1.05 – 1.2)Overhead stays small because decoding is memory-bound: a few dozen tree tokens and a few heads barely change a step dominated by streaming weights. Worked case: a one-token step takes 10 ms; add K=4 heads and a 40-node tree and the wider pass measures 11.5 ms, so c = 1.15. If the heads clear thresholds for τ = 2.8 tokens on average:
baseline for 2.8 tokens : 2.8 × 10 ms = 28.0 ms
Medusa per pass : 11.5 ms → 2.8 tokens
speedup = 28.0 / 11.5 = 2.43× (check: τ/c = 2.8/1.15 = 2.43×)Growing the tree to 100 nodes might lift τ to 3.1 but c to 1.25, giving only 2.48× — the tree is crossing into the compute-bound regime where extra nodes stop being free. The sweet spot is the largest tree still riding the memory-bound slack, and it is workload-dependent: code and templated text accept far longer runs than open-ended prose.