What this deep dive assumes

Start from the shape the base article establishes. A decoder produces a final hidden state h_t at each position, and the original language-model head maps h_t to a distribution over the next token, t+1. Medusa adds K extra heads: head i reads the same h_t and predicts the token at t+i+1, so from one forward pass you now have distributions for t+1 through t+K+1.

That offset detail is worth stating precisely: the base model still owns t+1; Medusa head i owns t+i+1. Everything that follows is about turning those parallel-but-independent guesses into correct tokens cheaply. We take the head shapes and tree-attention mask as given and spend our words on the training the base article does not cover.

Advertisement

The heads are residual blocks, not linear layers

A Medusa head is deliberately tiny but not trivial. Rather than a bare projection, each head is a single residual block: it takes h_t, passes it through one learned linear layer with a SiLU non-linearity, adds the result back to h_t, and only then applies a language-model projection to logits. In shorthand, head_i(h_t) = LMproj( h_t + SiLU(W_i · h_t) ).

The residual connection is the whole trick. It initialises the head close to the identity, so at the start of training a head behaves almost like the model’s own next-token predictor and then learns a small correction for its offset. That keeps each head cheap while giving it enough capacity to specialise. It also explains why the heads add little inference cost: they run in parallel off one hidden state, so the extra work is a few small matmuls, not another pass through the network.

Advertisement

The conditional-independence gap

Here is the flaw the deep dive has to confront. Every Medusa head conditions on the same h_t and nothing else. Head 2 predicting t+3 never sees what head 1 guessed for t+2. So the heads produce marginal predictions, not a joint one — they cannot model ‘if the next word is New, the one after is probably York.’

The consequence is a steady decay: head 1 is fairly accurate, head 2 less so, and by head 4 or 5 the top guess is often wrong, because predicting several tokens ahead from a single frozen context is harder and the head cannot see the intervening choices. This is precisely the gap EAGLE closes with an autoregressive feature recurrence, and it is why Medusa leans on a tree of several candidates per position rather than betting on one chain: breadth compensates for the heads’ inability to condition on each other.

From top-k heads to a sparse candidate tree

Because no single head is reliable, Medusa keeps the top few guesses from each. Take the top k_1 tokens from head 1, the top k_2 from head 2, and so on. The naive combination is a Cartesian product — every head-1 guess paired with every head-2 guess — which is k_1 × k_2 × … candidate continuations, exploding fast.

Medusa does not verify them all. It prunes to a fixed sparse tree chosen ahead of time: paths through more-confident, earlier heads get more branches; unlikely deep paths get dropped. The tree is laid out as a flat sequence of nodes with a matching attention mask and position ids, so a node attends only to its own ancestors, and the entire tree is pushed through the backbone in one forward pass. The base article details that mask; the point here is that tree shape is a tuned budget, trading more candidates (higher acceptance) against a heavier pass.

Typical acceptance, not exact rejection sampling

Classic draft-model speculative decoding uses rejection sampling: each drafted token is accepted or rejected with a probability crafted so the final output is distributed exactly as if you had sampled from the target model token by token. That guarantee is its selling point — speculation changes speed, never the distribution.

Medusa deliberately abandons that guarantee. Its heads do not define a clean proposal distribution you can do exact rejection against, so instead it uses typical acceptance: accept a proposed token if the original model’s probability for it clears a threshold of the form min(ε, δ · exp(−H)), where H is the entropy of the model’s distribution at that step. An absolute floor ε keeps confidently-predicted tokens flowing; the entropy term loosens the bar when the model itself is uncertain and tightens it when the model is peaked. The trade is explicit: you give up exact distribution matching for a higher acceptance rate and more speed. In practice it tracks greedy or low-temperature decoding closely, the regime most deployments run.

Medusa-1: training on a frozen backbone

The first and simplest recipe freezes the entire pretrained model and trains only the new heads. Each head i is supervised with cross-entropy against the true token at offset t+i+1, summed over positions, with a per-head weight that decays as the offset grows — roughly Σ_i λ^i · CE_i — because the far heads are intrinsically harder and should not dominate the gradient.

Freezing the backbone has two virtues. It is cheap: only a few small matrices train, so you can fit heads onto a large model on modest hardware, often with the backbone quantised. And it is safe — because the original weights never move, the model’s quality cannot regress; Medusa-1 either speeds things up or, at worst, does nothing. The ceiling is that frozen features were never shaped to be predicted several steps ahead, so the heads’ acceptance rate, and thus the speedup, is capped below what a co-adapted model could reach.