Medusa 多头投机解码:LLM 推理加速的工程演进与深度实战

2026 年,LLM 推理优化进入"后 PagedAttention"时代。当 KV Cache 优化和 Continuous Batching 成为标配,Speculative Decoding(投机解码)成为下一个性能突破口。本文从经典投机解码理论出发,深入 Medusa 多头投机架构的工程实现,涵盖 Draft Model 选择策略、投机宽度动态调节、与 vLLM/TensorRT-LLM 的集成,以及 2026 年最新研究的 Medusa-2 与 LayerSkip 变体。


1. 投机解码的核心原理:为什么可以"白嫖"性能?

大语言模型自回归生成(Autoregressive Generation)的瓶颈在于内存带宽墙:每个 token 的生成只需执行一次矩阵乘法,但需要从 HBM 中加载全部模型参数。对于 70B 模型(140GB BF16),生成一个 token 所需计算量不大,却要读 140GB 数据——A100 80GB HBM 带宽仅 2TB/s,理论极限约 14 tokens/s。

Speculative Decoding 的核心洞察:验证比生成便宜。用一个小型 Draft 模型(或投机头)一次性预测 K 个 token,然后用大型 Target 模型并行验证这 K 个 token。由于 Target 模型单次前向计算可以处理多个 prompt speculative tokens,验证 K 个 token 的代价远低于自回归生成 K 次。

数学表达:设自回归生成的接受率为 α(0 < α < 1),每次投机尝试产出 K 个 token,则加速比近似为:


加速比 ≈ (1 - α^(K 1)) / (1 - α)

当 α = 0.8, K = 5 时,加速比约为 3.6x。这意味着推理成本直接降到原来的 ~28%。

关键性质:Speculative Decoding 保证输出分布与标准自回归采样完全相同——它不是近似方法,而是在不改变任何输出质量前提下减少前向传播次数的技术。


2. 经典架构 vs Medusa:两条路线的工程博弈

2.1 经典 Speculative Decoding:双模型架构

经典方案使用独立的 Draft Model(如 LLaMA-7B 作为 LLaMA-70B 的 Draft),流程如下:


1. Draft Model 自回归生成 γ 个 token: [t₁, t₂, ..., tγ]
2. Target Model 单次前向处理 [prompt   t₁   t₂   ...   tγ]
3. 逐位比较 Target 与 Draft 的概率分布:
   - 若 Target 接受 tᵢ:继续比较 tᵢ₊₁
   - 若 Target 拒绝 tᵢ:从调整后的分布中重采样 tᵢ
4. 最终产出: [accepted tokens]   [1 new token]

工程痛点:

  • 双模型显存开销:需要同时加载 Draft Target,对于单卡部署极其不友好
  • Draft Model 选择困难:小模型接受率低,大模型显存开销高
  • 参数同步开销:Draft 和 Target 在同一个 batch 中调度需要复杂协调

2.2 Medusa 多头投机:复用 Target Model 本身

Medusa(arXiv: 2401.10774)的天才想法:不必另起炉灶训练 Draft Model,直接在 Target Model 头部附加 K 个投机头(Head)即可。

每个投机头 Hᵢ 负责预测序列中第 i 个未来位置的 token:


输入: Target Model 最后一层隐藏状态 h
投机头 H₁ → 预测第  1 个 token
投机头 H₂ → 预测第  2 个 token
投机头 H₃ → 预测第  3 个 token
...
投机头 Hₖ → 预测第  K 个 token

每个投机头就是一个简单的 MLP:


class MedusaHead(nn.Module):
    def __init__(self, hidden_size, vocab_size):
        super().__init__()
        self.mlp = nn.Sequential(
            nn.Linear(hidden_size, hidden_size),
            nn.SiLU(),
            nn.Linear(hidden_size, vocab_size)
        )
    
    def forward(self, hidden_states):
        return self.mlp(hidden_states)

训练策略极为简洁:冻结 Target Model 全部参数,仅训练 K 个 Medusa Head。使用 rejection sampling loss:


def medusa_loss(heads_logits, target_logits, target_ids, alpha=0.8):                        
                    
点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部