Medusa 深读:解码慢不是算力问题,是搬运问题

从"解码是内存带宽瓶颈"这个第一性约束出发,拆解 Medusa 的多头架构、Typical Acceptance 与 Tree Attention;再沿 Stern 2018 → Speculative Decoding → SpecInfer → Medusa → Hydra / EAGLE / DFlash 的 arXiv 系谱,看清这场"并行宽度换串行步数"的交易,以及评估任何投机解码方案的四个问题。

一块 A100 每秒能做 312 万亿次 FP16 乘加,但生成一个 token 时,它绝大部分时间在干一件小学生的活:把权重从显存搬到计算单元。以 70B FP16 模型粗略估算——140 GB 权重、约 2 TB/s 的 HBM 带宽,光搬一遍权重就要约 70 毫秒。这 70 毫秒里 GPU 的算力几乎在空转,而它换来的产出只有:一个 token

投机解码(Speculative Decoding)一族的所有设计,都是在跟这 70 毫秒讨价还价:既然搬一遍权重的钱已经花了,能不能一次多验证几个 token?

本文沿着这条主线拆解 Medusa(Cai et al., ICML 2024)——它把投机解码的三个环节(草稿、验证、接受)全部重新谈判了一遍:草稿模型被折叠进主模型、验证被摊平成一棵树、接受规则从”分布无损”放宽到”够像就行”。读完你会带走一个可复用的框架:评估任何投机解码方案的四个问题

本文的触发点是微信公众号”肖喵喵是喵大王”的《Speculative Decoding Evolution: Medusa》一文,在其骨架上补入了论文系谱与原始论文的核实细节。

一、第一性原理:为什么解码是”搬运问题”

自回归解码有一条铁律:第 t+1t+1 个 token 依赖第 tt 个 token。这意味着生成 100 个 token 就要 100 次串行的 forward pass,每一次都要把全部权重读一遍。

关键在于算术强度(arithmetic intensity,每读一字节做多少次运算):

  • 训练 / prefill:一次读权重,成千上万个 token 共享这次读取,算术强度高,GPU 是 compute-bound;
  • 解码(batch size 小时):一次读权重只服务一个 token 的计算,算术强度极低,GPU 是 memory-bound——瓶颈是 HBM 带宽,不是 FLOPs。

这就是投机解码的杠杆支点:验证 K 个候选 token 和生成 1 个 token,搬运的权重量几乎一样。只要能便宜地搞到”大概率正确”的候选序列,就能把 K 次串行搬运压缩成一次。

于是问题分解成三个子问题,也是这一族论文分化的三条战线:

  1. 草稿(Draft):便宜的候选 token 从哪来?
  2. 验证(Verify):怎么在一次 forward pass 里并行验证尽量多的候选?
  3. 接受(Accept):验证后按什么规则接受,接受得越长越快,但会不会改变输出分布?

二、系谱:Medusa 之前的三块基石

Medusa 不是凭空出现的,它的每个部件都有明确的出处。按时间线:

年份论文贡献对应子问题
2018Blockwise Parallel Decoding(Stern, Shazeer & Uszkoreit, NeurIPS 2018)多个输出头并行预测未来多个位置,再回退到验证通过的最长前缀;贪心不变时迭代数减半,实测最高 4× 提速草稿 + 验证的雏形
2022Fast Inference from Transformers via Speculative Decoding(Leviathan, Kalman & Matias, ICML 2023 Oral)用独立小模型起草 + 大模型并行验证 + 新颖采样规则,证明输出分布严格不变;T5-XXL 上 2–3×接受规则的无损标准
2023Accelerating LLM Decoding with Speculative Sampling(Chen et al., DeepMind)修正拒绝采样(modified rejection sampling),Chinchilla-70B 分布式场景 2–2.5×同上,工程验证
2023SpecInfer(Miao et al., ASPLOS 2024)把候选组织成 token 树,LLM 作为”树验证器”一次验证所有分支验证的树形化
2024Medusa(Cai et al., ICML 2024)多解码头 + Tree Attention + Typical Acceptance 的完整框架三条战线同时动刀

看出脉络了吗——Medusa ≈ Stern 2018 的多头草稿 + SpecInfer 的树验证 + 一个自己发明的有损接受规则。它的第一性创新不在任何单个部件,而在一个判断:草稿模型是多余的

传统投机解码需要养一个独立的 draft model:它要足够小(否则起草本身太慢),输出分布又要贴近 target model(否则接受率太低)——这两个要求天然打架。Medusa 的回答是:不养了,直接从 target model 的 hidden state 上长出几个头来

三、架构:把 draft model 折叠进 target model

3.1 Medusa Head 的定义

给定输入 X=(x1,,xt)X=(x_1,\dots,x_t),target model 在位置 tt 的最后一层 hidden state 为 htRdh_t \in \mathbb{R}^d。原 LM Head 用 hth_t 预测 yt+1y_{t+1};第 kk 个 Medusa Head 则跳过中间位置,直接预测 yt+k+1y_{t+k+1}

pt(k)=softmax ⁣(W2(k)(SiLU(W1(k)ht)+ht))p_t^{(k)} = \operatorname{softmax}\!\Big(W_2^{(k)}\big(\operatorname{SiLU}(W_1^{(k)} h_t) + h_t\big)\Big)

其中(按列向量约定)W1(k)Rd×dW_1^{(k)} \in \mathbb{R}^{d\times d}W2(k)RV×dW_2^{(k)} \in \mathbb{R}^{V\times d}dd 是 hidden 维度,VV 是词表大小。这就是一个带残差的单层 FFN,激活函数用 SiLU。

初始化很讲究:W1(k)W_1^{(k)} 置零W2(k)W_2^{(k)} 复制原 LM Head 的权重。这样残差支路在训练开始时是恒等的,每个 Medusa Head 的初始预测与原 LM Head 完全一致——训练从”预测下一个 token 的分布”平滑地漂向”预测第 k+1 个未来 token 的分布”。

KK 个头加原 LM Head,一次 forward pass 同时得到 K+1K+1 个位置的分布。代价也要算清:每个 W2(k)W_2^{(k)} 都是 V×dV \times d 的词表投影,KK 个头引入 O(KdV)O(KdV) 的额外参数和参数读取——头不是免费的,只是比一个完整的 draft transformer 便宜。

3.2 Medusa 的”原罪”:独立性假设

注意所有头都只吃同一个 hth_t,第 2 个头并不知道第 1 个头实际选了哪个 token。所以各头输出的不是联合分布,而是对不同未来位置的并行边缘预测(parallel marginal predictions)。把它们当作独立因子拼起来:

q(yt+1,,yt+K+1X)=kqk(yt+k+1X)q(y_{t+1},\dots,y_{t+K+1} \mid X) = \prod_{k} q_k(y_{t+k+1} \mid X)

这一般不等于 target model 的自回归联合分布 kp(yt+k+1X,yt+k)\prod_k p(y_{t+k+1} \mid X, y_{\le t+k})。位置越远,边缘分布与真实条件分布的差距越大,头的命中率越低。

两点值得强调的推论:

  1. 缺少序列依赖不妨碍理论上的无损。只要草稿真的从已知且归一化的 qkq_k 采样,并执行标准的 pk/qkp_k/q_k 接受 + (pkqk)+(p_k - q_k)_+ 残差修正,仍能严格恢复 target 分布——独立性只会降低 pkp_kqkq_k 的重叠度,从而压低接受率,不会破坏正确性。
  2. 但官方实现没有走这条路。Medusa 官方代码用 top-k 候选树 + Greedy/Typical/Nucleus 后验评估选最长前缀,并未实现完整的拒绝采样与残差修正(见其 medusa_generate / evaluate_posterior)。所以 Medusa 的验证本身不保证采样分布无损——这是读这篇论文最容易误会的一点。

这个独立性假设就是 Medusa 的阿喀琉斯之踵,后文会看到,Medusa 之后几乎每一篇改进论文都在瞄着它开枪。

四、训练:两档火力 + 自蒸馏

Medusa-1(冻结 backbone):只训练头。第 kk 个头拿位置 t+k+1t+k+1 的真实 token 当标签,做交叉熵:

L1=k=1KλkLCE(pt(k),yt+k+1),λk=0.8k\mathcal{L}_1 = \sum_{k=1}^{K} \lambda_k \cdot \mathcal{L}_{\mathrm{CE}}\big(p_t^{(k)},\, y_{t+k+1}\big), \qquad \lambda_k = 0.8^k

越远的位置越难预测、loss 越大,衰减系数 λk\lambda_k 防止远端头绑架总 loss。这档训练便宜、绝不伤害原模型(backbone 一个参数都不动)。

Medusa-2(联合微调):头和 backbone(或其 adapter)一起训,头的精度更高,但必须保住原有的 next-token 能力,所以要把原 LM loss 加回来:

L2=LLM+λ0L1\mathcal{L}_2 = \mathcal{L}_{\mathrm{LM}} + \lambda_0 \mathcal{L}_1

配合头先热身(warmup)、头与 backbone 用不同学习率等技巧,避免随机初始化的头在训练早期把 backbone 带偏。

自蒸馏(self-distillation):头最好用与 target model 输出分布相近的数据训练,但原始训练数据往往不公开,模型还可能经过 SFT/RLHF。Medusa 的做法是用 target model 自己对相近领域的 prompt 生成文本:生成的 token 给头当硬标签;对 Medusa-2,再用原模型的完整概率分布当 soft teacher 加一个 KL 约束,防止 backbone 在自己生成的数据上训练时能力退化。

五、接受规则:从”无损”到”够像就行”

Leviathan/Chen 式的 speculative sampling 有一个优雅但苛刻的性质:为了严格保持 target 分布,它必须按 min(1,p/q)\min(1, p/q) 拒绝采样并做残差修正——哪怕一个候选 token 在 target 看来完全合理,只要 draft 给它的概率偏高,也可能被拒。高温采样时 target 分布熵大、合理的续写很多,这种”洁癖”会明显压短接受长度。

Medusa 的 Typical Acceptance 换了一种哲学:不追求分布一致,只要求”这个 token 在 target 眼里是典型的(typical)“。候选 token xx 被接受当且仅当

poriginal(xcontext)>min ⁣(ϵ,  δeH(poriginal(context)))p_{\mathrm{original}}(x \mid \text{context}) > \min\!\big(\epsilon,\; \delta e^{-H(p_{\mathrm{original}}(\cdot \mid \text{context}))}\big)

其中 ϵ\epsilon 是硬阈值,δeH\delta e^{-H} 是熵相关阈值,H(p)=xp(x)logp(x)H(p) = -\sum_x p(x)\log p(x)。直觉:

  • 熵高(很多合理续写)→ δeH\delta e^{-H} 变小 → 门槛降低,更多候选过关;
  • 熵低(答案基本唯一)→ 门槛抬高,只放行高概率 token。

第一个 token 直接用原 LM Head 的贪心预测无条件接受(保证每步至少推进一个 token),后续按上式过滤,最后取最长的全接受前缀提交。

三档接受规则的适用边界,别搞成二元对立:

接受规则分布保证适用场景
Greedy 匹配保持 target 的确定性 argmax温度 0、要求可复现
完整 speculative sampling(p/qp/q + 残差)严格无损对分布一致性有硬要求(如评测、蒸馏数据生产)
Typical Acceptance有损,允许偏移高温、多样化生成,接受长度优先

六、Tree Attention:把”若干条候选”摊成”一棵树”

6.1 为什么是树

每个头不止出一个候选——第 kk 个头取 top-sks_k。假设原 LM Head 出 r0r_0,头 1 保留 (r11,r12)(r_{11}, r_{12}),头 2 保留 (r21,r22,r23)(r_{21}, r_{22}, r_{23}),那么完整候选序列就是笛卡尔积:2×3=62 \times 3 = 6 条。逐条验证要 6 次 forward pass,等于把省下来的搬运又还回去了。

树的观察是:这 6 条序列共享大量前缀。把它们挂成一棵树,需要送进验证的新增 token 节点只有 2+6=82 + 6 = 8 个(第一层 2 个,第二层每个父节点下挂 3 个),而不是 6×2=126 \times 2 = 12 个。一般地,深度 KK、每层取 top-sks_k 时:

候选序列数=k=1Ksk,新增节点数=k=1Kj=1ksj\text{候选序列数} = \prod_{k=1}^{K} s_k, \qquad \text{新增节点数} = \sum_{k=1}^{K} \prod_{j=1}^{k} s_j

一次 forward pass 验证整棵树的手法是改 attention mask:每个节点只能看到虚拟根和自己分支上的祖先,看不到兄弟分支;同一深度的节点共享同一个 position index。这样每条根到叶的路径在因果语义上都等价于一条独立的候选序列——树只是把它们的计算叠在了一个 batch 维度里。这个思路与 SpecInfer 的 token tree verification 同源。

6.2 稀疏树:预算内的贪心构造

笛卡尔积的节点数随深度指数增长,而不同 rank 的命中概率极不均衡(rank 1 命中率远高于 rank 2)。Medusa 的做法是拿一个 calibration 数据集统计每个头的命中率画像

ak(i)=P(Rk=i)a_k^{(i)} = P\big(R_k = i\big)

即真实 token 恰好排在第 kk 个头第 ii 名的概率。对一条 rank 路径 (i1,,ik)(i_1,\dots,i_k),用跨头独立性近似估计整条路径命中的概率(路径权重):

w(i1,,ik)j=1kaj(ij)w(i_1,\dots,i_k) \approx \prod_{j=1}^{k} a_j^{(i_j)}

然后在节点预算 BB 内贪心建树:每次从 frontier 里挑路径权重最大的节点加入树,把它的子节点补进 frontier,重复 BB 次。

为什么贪心选路径权重是在优化期望接受长度?LcoverL_{\mathrm{cover}} 是参考续写被候选树覆盖的最大深度,TkT_k 是树在第 kk 层保留的 rank 路径集合。由非负整数随机变量的尾和公式 E[L]=k1P(Lk)\mathbb{E}[L] = \sum_{k\ge 1} P(L \ge k)

E[Lcover]=k=1K(i1,,ik)Tkj=1kaj(ij)\mathbb{E}[L_{\mathrm{cover}}] = \sum_{k=1}^{K} \sum_{(i_1,\dots,i_k)\in T_k} \prod_{j=1}^{k} a_j^{(i_j)}

——期望覆盖深度恰好等于树中所有节点路径权重之和。所以”每次加入权重最大的节点”就是对 E[Lcover]\mathbb{E}[L_{\mathrm{cover}}] 的贪心最大化,每个新节点的边际贡献就是它自己的路径权重。

手算一个小例子。3 个头、每头只看前两名,calibration 得到:

Head kkak(1)a_k^{(1)}ak(2)a_k^{(2)}
10.750.20
20.600.25
30.500.20

预算 B=6B=6 个非根节点,贪心过程:

次序选中节点 uu(rank 路径)w(u)w(u)
1(1)(1)0.750.75
2(1,1)(1,1)0.75×0.60=0.450.75\times0.60=0.45
3(1,1,1)(1,1,1)0.45×0.50=0.2250.45\times0.50=0.225
4(2)(2)0.200.20
5(1,2)(1,2)0.75×0.25=0.18750.75\times0.25=0.1875
6(2,1)(2,1)0.20×0.60=0.120.20\times0.60=0.12
E[Lcover]0.75+0.45+0.225+0.20+0.1875+0.12=1.9325\mathbb{E}[L_{\mathrm{cover}}] \approx 0.75+0.45+0.225+0.20+0.1875+0.12 = 1.9325

注意树长成了左偏的形状:高 rank 节点及其后代的路径权重天然更大,预算自然向”每层都选第一名”的主干倾斜。也要记住这套优化的边界——它依赖跨头独立性和 calibration 分布,优化的是接受长度的近似,不直接等于 wall-clock 加速。

七、Trade-offs:账要算到硬件上

7.1 速度账

先分清两个量。Acceleration rate AA 是每个 decoding step 平均推进的 token 数(不是”接受的 token 占提案的比例”)。但每一轮 Medusa 验证比普通解码一步更贵(要处理几十个树节点),设单轮延迟比 r=TM/TARr = T_{\mathrm{M}} / T_{\mathrm{AR}},则

wall-clock speedupAr\text{wall-clock speedup} \approx \frac{A}{r}

论文 Table 1 的实测(Medusa-2,MT-Bench):

模型SpeedupAcc. Rate AAMT-Bench 质量变化
Vicuna-7B2.83×3.47+0.01
Vicuna-13B2.83×3.51−0.14
Vicuna-33B2.35×3.01+0.05

A3.03.5A \approx 3.0\text{–}3.5 而 speedup 只有 2.352.83×2.35\text{–}2.83\times,中间的差值就是 r>1r > 1 在吃收益。论文摘要口径:Medusa-1 无损质量下 >2.2×,Medusa-2 达 2.3–3.6×。

7.2 树不是越大越好

节点越多覆盖率越高,但验证计算和 KV Cache 操作也在涨。按原论文 Figure 4 的实测(转引自前述微信文章的解读),64 节点的优化稀疏树的 acceleration rate 优于部分 256 节点的稠密树;论文附录的 roofline 模拟进一步给出边界:Llama-7B、A100、batch size 1、序列长 1024 的设定下,候选 token 超过约 64 后模拟加速开始下降;batch size 超过 32 后收益递减甚至转负。

这印证了开头的第一性原理:投机解码的收益来自”memory-bound 时算力免费”。batch size 一大,解码本身就滑向 compute-bound,免费算力消失,多验证的每个候选都是真金白银的计算——候选树的尺寸必须在真实模型、真实 batch、真实硬件上标定,不能只看接受长度

7.3 质量账:证据的边界

Typical Acceptance 的分布偏移是算法预期内的结果,问题只在于偏移是否可感知。论文的证据:MT-Bench 质量变化在 −0.14 到 +0.05 之间;在 writing/roleplay、温度 0.7 的设定下,阈值调严则加速下降、质量分上升——旋钮是连续的。

但要诚实地标注证据边界:这些评估集中在 MT-Bench 的少数类目、单一温度、以 GPT-4 judge 均分为指标。均分相近既不能说明生成分布接近,也不能排除在 factuality、reasoning、code、safety 或长上下文一致性上出现退化。把它读作”特定设置下未观察到明显退化的初步证据”,而不是”Typical Acceptance 广泛无害”的证明。

八、Medusa 之后:每个弱点都长出了一篇论文

Medusa 最有生命力的部分是它的范式——共享 target hidden state、并行起草、树验证——而它的每个具体缺陷,都成了后续论文的靶子:

graph LR
    S[Stern 2018<br/>多头并行草稿] --> M[Medusa 2024]
    L[Leviathan/Chen<br/>无损投机采样] --> M
    T[SpecInfer 2023<br/>树验证] --> M
    M -->|补序列依赖| H[Hydra 2024]
    M -->|特征级自回归| E1[EAGLE 2024]
    E1 -->|静态树→动态树| E2[EAGLE-2 2024]
    E2 -->|training-time test| E3[EAGLE-3 2025]
    E3 -->|自回归草稿→扩散并行| D[DFlash 2026]
  • Hydra(Ankner et al., 2024)直击独立性假设:把草稿头改成顺序依赖的——后面的头条件化在前面头实际选出的 token 上。仅此一改,吞吐较 Medusa 提高至多 1.31×(较自回归 2.70×)。
  • EAGLE(Li et al., 2024)换了个更深的答案:与其让 K 个头独立猜边缘分布,不如在特征层(倒数第二层 hidden state)做自回归——特征序列比 token 序列更规整、更好预测;再把下一时刻的 token 采样结果喂给 draft 头以消解特征的不确定性。LLaMA2-Chat 70B 上 2.7–3.5×,且保持分布无损。序列依赖回来了,但代价从”完整 draft transformer”降到”一层自回归头”。
  • EAGLE-2 发现接受率强烈依赖上下文,静态候选树(Medusa 的 calibration 稀疏树也是静态的)天然次优,于是用 draft 模型的置信度在线动态建树,3.05–4.26×,比 EAGLE-1 再快 20–40%。
  • EAGLE-3 放弃特征预测约束,改为 training-time test + 多层特征融合直接预测 token,最高 6.5×。
  • DFlash(Chen, Liang & Liu, 2026)则动了最后一块自回归领地:起草本身还是串行的——它改用轻量 block diffusion 模型,非因果 mask 一次并行吐出整块草稿 token,报告 6× 以上无损加速、较 EAGLE-3 至多再快 2.5×。

工程侧,Medusa 至今仍是 HuggingFace TGIvLLM--spec-method medusa)支持的投机解码后端之一——但在多数 workload 下,它已不是默认最优解。

九、带走的框架:投机解码四问

回到主线。任何一个投机解码方案,拿这四个问题去拆,基本能定位它在设计空间里的坐标:

  1. 草稿从哪来? 独立小模型(Leviathan/Chen)→ target 复用多头(Stern/Medusa)→ 顺序依赖头(Hydra)→ 特征级自回归(EAGLE)→ 扩散并行(DFlash)。演化方向:越来越深地复用 target 已经算出来的东西
  2. 候选怎么组织与验证? 单链 → 静态树(SpecInfer/Medusa)→ 动态树(EAGLE-2)。演化方向:树的形状从离线标定走向在线自适应。
  3. 接受规则损不损? Greedy 匹配 / 严格 p/qp/q 无损 / Typical 有损——没有对错,只有场景匹配;但要警惕”实现里没做残差修正却宣称无损”的缝隙。
  4. 瓶颈在带宽还是算力? 投机解码的全部收益建立在 memory-bound 的前提上。batch 大了、序列长了、树宽了,前提就塌了——speedup ≈ A/r 这笔账必须在目标硬件上重算

Medusa 教会这个领域最重要的一课,或许不是那三个头,而是这个判断:加速解码不必在模型外面找帮手,target model 自己算出来的 hidden state 里,已经藏着关于未来好几个 token 的信息——问题只是你用什么结构把它取出来。从 Medusa 的边缘预测,到 EAGLE 的特征自回归,再到 DFlash 直接拿 target 的上下文特征喂扩散模型,都是对同一句话越来越高效的回答。


参考文献

arXiv 论文(按时间线)

  1. Stern, Shazeer & Uszkoreit. Blockwise Parallel Decoding for Deep Autoregressive Models. NeurIPS 2018. arXiv:1811.03115
  2. Leviathan, Kalman & Matias. Fast Inference from Transformers via Speculative Decoding. ICML 2023 (Oral). arXiv:2211.17192
  3. Chen, Borgeaud, Irving, Lespiau, Sifre & Jumper. Accelerating Large Language Model Decoding with Speculative Sampling. 2023. arXiv:2302.01318
  4. Miao et al. SpecInfer: Accelerating Generative Large Language Model Serving with Tree-based Speculative Inference and Verification. ASPLOS 2024. arXiv:2305.09781
  5. Cai, Li, Geng, Peng, Lee, Chen & Dao. Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads. ICML 2024, PMLR 235:5209–5235. arXiv:2401.10774
  6. Li, Wei, Zhang & Zhang. EAGLE: Speculative Sampling Requires Rethinking Feature Uncertainty. 2024. arXiv:2401.15077
  7. Ankner, Parthasarathy, Nrusimha, Rinard, Ragan-Kelley & Brandon. Hydra: Sequentially-Dependent Draft Heads for Medusa Decoding. 2024. arXiv:2402.05109
  8. Li, Wei, Zhang & Zhang. EAGLE-2: Faster Inference of Language Models with Dynamic Draft Trees. 2024. arXiv:2406.16858
  9. Li, Wei, Zhang & Zhang. EAGLE-3: Scaling up Inference Acceleration of Large Language Models via Training-Time Test. 2025. arXiv:2503.01840
  10. Chen, Liang & Liu. DFlash: Block Diffusion for Flash Speculative Decoding. 2026. arXiv:2602.06036

工程实践与原文