普通自回归模型一次前向计算只确定下一个 Token。这个 Token 出来后,模型才能继续算后一个。单请求生成常受显存带宽限制,每走一步都要读取一遍大模型权重,计算单元却没有吃满。Medusa 想让一次读取换回多个可接受的 Token。
我查看了 Medusa 论文的方法、训练方案和消融实验。它省掉了单独的小型草稿模型,在原模型最后的隐藏状态上增加多个预测头。第一个头猜紧接着的 Token,后面的头分别猜更远的位置,每个头还能提出若干高概率候选。
多条候选用一棵树一起验证
几个预测头的候选直接排列组合,数量会很快增加。Medusa 把共享前缀的候选并成树,再调整注意力掩码,让原模型在一次前向计算里检查多条延续。通过验证的最长前缀进入正式输出,剩余候选丢弃。这些计算白做了。下一轮随后继续生成。
这一步仍以原模型的判断为准。使用标准拒绝采样时,输出分布可以与原模型保持一致。论文还提出 typical acceptance,允许接受处在合理概率范围内的候选,用温度和阈值控制偏离。差别不小。它可以提高接受率,也会放宽与原始采样分布完全一致的要求,质量需要单独评测。
候选越多,命中的机会通常越大,验证开销也会上升。树里加入一条很少命中的分支,仍会扩大注意力矩阵和线性层工作。这会拖慢整轮验证。论文通过搜索常见候选模式选择稀疏树结构。消融结果显示,只有多个预测头而没有树注意力时,加速约 1.5 倍;加入树注意力后约 1.9 倍,优化树结构后约 2.2 倍。多猜几个词并不会直接按头数加速。
两种训练方式动到的参数不同
Medusa 1 冻结原模型,只训练新增预测头。原模型能力不会被训练过程改动,所需显存也较少,论文将它用于已有模型的无损加速。实验报告超过 2.2 倍速度提升。若原训练数据拿不到,还可以让原模型生成数据,再用自蒸馏教这些预测头。
Medusa 2 会同时训练预测头与原模型。预测头能更准确地猜到后续 Token,论文报告的加速提高到 2.3 倍至 2.8 倍。代价也很清楚,训练流程要保护原来的下一词预测能力,模型质量必须重新验。论文主要测试批次为一的本地使用场景,服务端大批次下的收益不能照搬。
评估 Medusa 时,我会把一次前向平均接受多少 Token、每轮验证耗时和任务质量放在一起看。只报每秒 Token 数,会掩盖接受规则已经改变输出的情况。代码补全等后续结构较容易预测的任务,多个头更可能连续命中;开放式写作有更多合理分支,接受率可能降低。
Medusa 的价值来自减少完整模型前向计算次数。它没有消除验证,也没有保证每次猜测都能留下。预测头、候选树和接受规则三处一起调好,显存带宽才可能真正换成更快的可见输出。