在线聊天服务中,用户期望即时响应,但大模型的自回归解码却像打字机一样逐个生成 token。每一次生成都需要将全部模型参数从高带宽内存搬运到计算单元,而实际计算量却很小,导致内存带宽成为瓶颈。批处理能提升吞吐,却会进一步增加单请求延迟,并急剧扩大 KV 缓存占用。投机解码试图用一个轻量草稿模型提前生成多个候选 token,再由原始模型并行验证,但寻找一个既能快速生成又能被原始模型接受的草稿模型并不容易,双模型部署也增加了系统复杂度。
Medusa 给出了另一种思路:不引入独立草稿模型,而是在原始模型上直接附加多个轻量预测头,让模型自己充当“草稿者”。这些头并行预测未来多个位置的 token,形成一棵候选序列树,再通过树状注意力机制一次性验证所有候选,最后用典型接受方案筛选出最长可接受前缀。这样既保留了单模型部署的简洁性,又实现了与投机解码相当的加速效果。
为什么需要并行预测头
自回归解码的瓶颈在于顺序依赖:每生成一个 token,都要等待上一个 token 的隐状态,无法利用 GPU 的并行计算能力。即使硬件能同时处理上千个操作,也只能串行执行。投机解码通过“小模型猜、大模型验”来打破串行,但面临三个痛点:一是草稿模型难寻,小模型往往与目标模型分布不一致,导致大量候选被拒绝;二是系统复杂度高,需维护两个模型,在分布式部署中尤甚;三是采样效率低,重要性采样在高温下会引入额外开销,甚至比贪心解码更慢。
Medusa 的解决方案是“自投机”:在原始模型的最后一层 Transformer 之后,保留原有的语言模型头(预测下一个 token),并额外添加 K 个解码头,分别预测 t+1、t+2、…、t+K 位置的 token。每个头都是一个带残差连接的单层前馈网络,结构简单,参数少。训练时冻结骨干模型,只微调这些头,因此可以在单张 GPU 上快速完成。推理时,每个头独立输出其位置上的 top-k 候选 token,组合后形成多个候选序列,再由原始模型并行验证。由于这些头共享骨干模型的隐藏表示,预测分布与原始模型更接近,接受率更高。
多头预测与树状注意力
Medusa 的核心组件包括三个部分:Medusa 解码头、树状注意力机制和典型接受方案。
Medusa 解码头
每个解码头负责预测一个特定未来位置的 token。假设有 K 个头,则第 i 个头预测当前步之后的第 i 个 token。输入是骨干模型最后一层的隐藏状态,输出是该位置上的词表概率分布。头的结构通常为一个或多个残差块(ResBlock)接一个线性层,残差块由 LayerNorm 和前馈网络组成。这种设计让头既能利用骨干模型的深层语义,又能通过残差学习位置特定的预测能力。
训练时,可以使用原始训练数据,也可以让模型自生成数据(自蒸馏)。损失函数为交叉熵,仅更新解码头参数。实验表明,在 Vicuna-7B 上,预测 next-next token 的 top-1 准确率约 60%,但 top-5 准确率超过 80%。这意味着如果每个头保留多个候选,可以大幅提高最终序列被接受的概率。
树状注意力
每个头输出 top-k 个候选,不同位置的候选组合起来形成一棵树。例如,第一个头预测 t+1 位置的 2 个候选(It, I),第二个头预测 t+2 位置的 3 个候选(is, ’, the),则共有 2×3=6 条候选路径,构成一棵深度为 2 的树。树中的每个节点代表一个 token,从根到叶的一条路径就是一个候选序列。
为了并行验证这些序列,需要构建特殊的注意力掩码。在常规自回归解码中,每个 token 只能看到它之前的 token。在树状结构中,一条路径内的 token 遵循因果掩码,但不同路径之间不能互相看到,否则会泄露信息。因此,掩码设计为:对于每个 token,只允许它关注其前缀路径上的 token。同时,位置编码也需相应调整,确保每个 token 获得正确的位置索引。
Medusa 的树状注意力与 SpecInfer 的 token tree verification 类似,但有一个关键区别:Medusa 的树结构在推理时是规则且固定的,掩码可以预处理并缓存,进一步降低开销。而 SpecInfer 中每个小模型生成的序列长度不同,掩码需动态构建。
典型接受方案
验证阶段,原始模型对树中的所有候选 token 计算概率分布。然后需要决定接受多长的前缀。传统投机解码使用重要性采样,要求草稿模型与目标模型分布匹配,但在高温下效率下降。Medusa 提出典型接受方案:设定一个阈值,只接受那些在原始模型下概率足够高的候选。具体做法是,对于每个候选 token,计算其概率是否高于一个基于熵的阈值(取硬阈值与熵相关阈值的最小值)。第一个 token 总是用贪心解码接受,以保证每步至少生成一个 token。最终输出的是通过测试的最长前缀。
这种方案的优势在于,它不强制分布匹配,而是直接根据原始模型的置信度进行筛选。当采样温度为零时,退化为贪心解码;温度升高时,接受条件放宽,可以接受更长的序列,从而在创造性生成中也能获得加速。
训练流程:从冻结骨干到联合微调
Medusa 提供两种训练方案:Medusa-1 和 Medusa-2。
Medusa-1:冻结骨干,仅训练头
这是最轻量的方式。骨干模型完全冻结,只训练新增的 Medusa 头。训练数据可以是原始模型训练语料的子集,也可以用模型自生成的数据(自蒸馏)。自蒸馏的过程是:用原始模型对一批输入生成输出,然后用这些输出作为标签训练 Medusa 头。这样即使没有原始训练数据,也能为任意微调模型添加 Medusa 头。训练只需单张 A100-80G GPU,7B 模型几小时即可完成。
Medusa-2:联合微调
为了进一步提升预测准确率,Medusa-2 同时微调解码头和骨干模型。但这会带来风险:骨干模型的能力可能退化。为此,Medusa-2 采用一种特殊训练策略:在训练头的同时,对骨干模型施加正则化或使用较小的学习率,以保持其原始语言建模能力。实验表明,Medusa-2 可将加速比从 2.2 倍提升至 2.3~3.6 倍,但需要更多计算资源和更仔细的超参调整。
下表对比了两种训练方案的关键差异:
| 特性 | Medusa-1 | Medusa-2 |
|---|---|---|
| 训练对象 | 仅 Medusa 头 | 头 + 骨干模型 |
| 骨干模型是否冻结 | 是 | 否(但需保护原始能力) |
| 训练数据需求 | 可自蒸馏,无需原始数据 | 通常需要原始训练数据或高质量数据 |
| 训练成本 | 低,单 GPU 数小时 | 较高,需多 GPU 或更长训练时间 |
| 预测准确率 | 中等,top-1 约 60% | 更高,top-1 可提升 |
| 加速比 | 约 2.2 倍 | 2.3~3.6 倍 |
| 适用场景 | 快速部署,资源有限 | 追求极致加速,可接受再训练成本 |
在线聊天场景中的推理流程
考虑一个在线聊天服务,用户输入“今天天气真好,适合”,模型需要补全。使用 Medusa 的推理过程如下:
- Prefill 阶段:对输入 prompt 进行一次性前向计算,生成 KV 缓存,并得到最后一个 token 的隐藏状态。
- 候选生成:隐藏状态送入 K 个 Medusa 头,每个头输出 top-k 个候选 token。例如,K=2,k1=2,k2=3,则生成 6 条候选序列。
- 树验证:将候选序列组织成树,构建树掩码,与原始模型一起进行前向计算,得到每个候选 token 的概率。
- 接受决策:应用典型接受方案,从根节点开始遍历,接受概率高于阈值的 token,直到遇到拒绝或叶节点。假设接受的前缀为“出去走走”,则这些 token 被正式输出。
- 下一轮:以接受的最后一个 token 的隐藏状态作为起点,重复步骤 2~4,直到生成结束符。
下图展示了这一流程:
flowchart TD
A[输入 prompt] --> B[Prefill: 计算 KV 缓存和最后隐藏状态]
B --> C[Medusa 头并行预测多个候选 token]
C --> D[构建候选树和树掩码]
D --> E[原始模型并行验证所有候选]
E --> F[典型接受方案筛选最长可接受前缀]
F --> G{是否生成结束符?}
G -- 否 --> H[以接受的最后一个 token 为起点]
H --> C
G -- 是 --> I[输出完整回复]
关键点在于,除了第一次 Prefill,后续每个解码步骤都同时完成“生成候选”和“验证候选”,即边生成边验证。因为验证阶段已经计算了所有候选的 logits,这些 logits 可以直接用于下一轮的候选生成,无需额外前向传播。
与投机解码的权衡对比
Medusa 与投机解码的根本区别在于草稿来源。投机解码依赖独立的小模型,而 Medusa 的草稿来自同一模型的多头预测。这带来一系列工程权衡:
- 模型部署:Medusa 只需维护一个模型,减少内存占用和系统复杂度;投机解码需要加载两个模型,显存压力更大。
- 训练成本:Medusa 仅需微调解码头,成本低;投机解码可能需要从头训练或蒸馏一个草稿模型,成本更高。
- 预测质量:Medusa 头共享骨干模型的隐藏表示,分布偏移小,接受率高;独立草稿模型可能与目标模型存在分布差异,导致更多拒绝。
- 加速潜力:Medusa 的加速比受限于头的数量和每个头的 top-k 设置,通常为 2~3.6 倍;投机解码理论上可以达到更高加速,但需要极强且对齐良好的草稿模型。
- 采样灵活性:Medusa 的典型接受方案在高温下仍能加速;投机解码的重要性采样在高温下可能失效。
下表总结了关键对比维度:
| 维度 | Medusa | 传统投机解码 |
|---|---|---|
| 草稿模型 | 无独立模型,使用附加头 | 独立小模型 |
| 模型数量 | 1 | 2 |
| 训练成本 | 低(微调头) | 高(训练或蒸馏草稿模型) |
| 系统复杂度 | 低,单模型部署 | 高,需协调两个模型 |
| 预测准确率 | 较高,共享隐藏表示 | 依赖草稿模型质量 |
| 加速比 | 2.2~3.6 倍 | 可达 2.5 倍以上(理想情况) |
| 高温采样加速 | 有效 | 可能退化 |
| 适用场景 | 单模型服务,资源受限 | 可获取强草稿模型时 |
候选接受率如何影响加速比
加速比取决于每个解码步骤平均接受的 token 数。设每步平均接受长度为 L,则理想加速比约为 L。L 由两个因素决定:每个头的 top-k 设置和典型接受的阈值。
增大 top-k 会增加候选数量,提高接受长前缀的概率,但也增加验证阶段的计算量。例如,第一个头取 top-2,第二个头取 top-3,共 6 条候选;若都取 top-5,则候选数增至 25,验证成本显著上升。因此需要权衡。实验表明,适度扩大 top-k(如 2~3)可在不显著增加延迟的情况下提升加速比。
典型接受的阈值控制着接受严格度。阈值越高,越保守,接受长度越短,加速比越低,但生成质量更接近原始模型;阈值越低,越激进,接受长度增加,加速比提升,但可能引入低概率 token,影响质量。在创造性任务中,适当降低阈值可获得额外 10% 的加速,而质量下降可忽略。
当采样温度升高时,原始模型的概率分布更平坦,候选 token 的概率普遍降低,若阈值不变,则接受率下降。但典型接受方案使用熵相关阈值,温度升高时熵增,阈值自动降低,从而维持接受率。这使得 Medusa 在非贪心生成中仍能有效加速,而传统投机解码可能因重要性采样失效而减速。
失败模式与部署边界
Medusa 并非万能,在以下情况下可能失效或收益有限:
- 长尾生成任务:如代码生成或数学推理,需要精确的 token 序列,候选接受率可能骤降,导致加速比接近 1。因为 Medusa 头难以预测这些领域中高度结构化的 token。
- 小模型或窄领域模型:骨干模型的隐藏表示可能不够丰富,Medusa 头预测准确率低,加速效果差。
- 极高吞吐场景:Medusa 主要针对 batch size=1 的低延迟场景优化。当 batch size 很大时,内存带宽已充分利用,并行候选带来的额外计算可能成为负担,加速效果减弱。
- 分布式推理:虽然 Medusa 不引入新模型,但树状注意力需要特殊的掩码和位置编码,在张量并行或流水线并行框架中可能需要额外适配。
- 阈值敏感:典型接受的阈值需要针对具体任务调整。设置不当可能导致质量下降或加速不明显。
部署时,应监控以下指标:
- 每步接受 token 数:直接反映加速效果,若持续接近 1,需检查头预测准确率或阈值设置。
- 生成质量:使用自动评估指标(如 BLEU、ROUGE)或人工评估,确保加速未损害输出。
- 延迟分布:观察 P50/P99 延迟,确保长尾请求未被过度拒绝。
- 显存占用:Medusa 头增加少量参数,但树验证可能增加峰值激活内存,需确保不超出 GPU 显存。
未解决的问题与展望
Medusa 在简化投机解码方面迈出了重要一步,但仍有一些开放问题:
- 动态头数量:当前头数量固定,能否根据输入难度动态调整?简单输入用更少的头,复杂输入用更多头,可能进一步优化计算。
- 头结构的改进:目前头是简单的残差网络,能否设计更强大的预测结构,如引入注意力或条件计算?
- 与量化、蒸馏的结合:Medusa 与模型量化、蒸馏等技术结合时的效果和稳定性尚需更多验证。
- 多轮对话中的状态管理:在多轮对话中,KV 缓存不断增长,树状注意力的掩码和位置编码如何处理长上下文?
- 标准化支持:目前 Medusa 需要修改模型代码,能否将其集成到主流推理框架(如 vLLM、TensorRT-LLM)中,提供开箱即用的加速?
Medusa 的核心贡献在于证明了“自投机”的可行性:通过极低的训练成本,在不改变原有模型能力的前提下,将单模型推理速度提升 2 倍以上。对于在线聊天这类延迟敏感的应用,Medusa 提供了一条简单有效的加速路径。