五十六:Speculative Decoding-投机解码¶
来源:http://mp.weixin.qq.com/s?__biz=MzYyNTk3Njg1NA==&mid=2247484729&idx=1&sn=9213acd67980fc3d4d16c4d6e6c0631e&chksm=f01eb640c7693f565ce4666ce00efe2e72b569cc81ec88fc1125ef615bc8f8235eb061eec9d6#rd
1. 学习范围¶
本日主题是 Speculative Decoding,也就是投机解码,重点覆盖经典 draft-verify 机制和 Medusa。 本日覆盖: - 自回归 Decode 为什么慢。
-
投机解码的提议-验证思想。
-
Draft model 与 target model。
-
Greedy 和 sampling 下的接受机制。
-
为什么投机解码可以保持输出分布。
-
加速比来源和失败条件。
-
Medusa 的多头候选与 tree verification。
-
常见面试问题和生产化注意事项。
2. 投机解码的定位¶
投机解码是一类加速 LLM Decode 的方法。它不改变目标模型本身的概率分布,而是用更便宜的方式提前猜多个 token,再让目标模型一次性验证。
核心目标: - 减少目标大模型的自回归调用次数。
-
一次 target forward 尽可能产出多个 token。
-
在保持生成质量或分布等价的前提下提速。
它主要优化 Decode 阶段,对 Prefill 的帮助通常有限。
3. 自回归 Decode 为什么慢¶
标准 Decode 每生成一个 token 都要调用一次目标模型:
即使有 KV cache,每一步仍然需要访问模型权重和历史 KV。由于每步只生成少量 token,GPU 可能无法充分利用,且串行依赖很强。 投机解码试图把多个 token 的验证合并到一次 target forward 中。4. Draft-Verify 基本流程¶
经典投机解码包含两个模型: - Draft model:小模型或便宜模型,用来快速猜候选 token。
- Target model:原始大模型,用来验证候选 token。
流程:
1. draft model 连续生成 k 个候选 token
2. target model 对这 k 个 token 一次性并行验证
3. 接受连续通过验证的 token
4. 从拒绝位置重新采样或回退
5. 为什么 target 可以一次验证多个 token¶
Transformer 在给定一段候选序列时,可以用 causal mask 一次前向计算每个位置的 logits。
先分清两个场景的输入差异 · 场景 A:正常自回归生成(串行,慢) 每一步只输入已生成前缀,单步只预测下一个 token: 输入 [x₁] → 算出位置 1 的 logit(预测 x₂) 输入 [x₁,x₂] → 算出位置 2 的 logit(预测 x₃) 输入 [x₁,x₂,x₃] → 算出位置 3 的 logit(预测 x₄)每轮只新增一个 token,必须串行循环 N 次,N 是序列长度。 · 场景 B:整条候选序列一次性前向(并行,投机解码验证用) 一次性把完整候选 [x₁,x₂,…,xₖ] 全部塞进 Transformer。依靠因果掩码 causal mask,矩阵注意力天然支持同时计算全部 k 个位置的输出 logits。 - 如果候选 token 等于 Target 最大概率 token:接受,继续下一个
- 一旦不匹配:丢弃后面所有 draft,只用 Target 当前分布重新采样,截断
因此,虽然生成本身是自回归的,但验证一个已经提出的候选序列可以并行完成。 这是投机解码的关键: - draft 负责串行猜测。
-
target 负责并行验证。
-
最终输出仍由 target 分布控制。
6. Greedy 场景¶
在 greedy decoding 中,验证逻辑比较直观。 Draft 提出若干 token,target 计算对应位置的 argmax: - 如果 draft token 等于 target argmax,则接受。
-
一旦某个位置不一致,就停止接受。
-
后续从 target 的正确 token 继续生成。
这种方式容易理解,但生产中还常见 sampling 场景。
7. Sampling 场景¶
在 sampling 下,不能简单用“是否等于 argmax”判断。经典 speculative sampling 使用接受-拒绝机制保持目标分布。 设: - target 分布为 p。
-
draft 分布为 q。
-
draft 采样 token x。
接受概率通常为:
如果拒绝,则从修正后的剩余分布中采样,保证最终分布等价于 target model。8. 分布保持为什么重要¶
投机解码的一个重要卖点是可以在理论上保持目标模型分布,而不是用小模型替代大模型。 这意味着: - 输出质量由 target model 保证。
-
draft model 只影响速度,不应改变最终分布。
-
适用于需要严格复现目标模型采样语义的场景。
如果实现没有保持分布,就变成近似加速,质量和一致性要重新评估。
9. 加速比来源¶
投机解码的加速来自: - target model 一次 forward 验证多个 token。
-
draft model 比 target model 便宜。
-
draft token 被接受的平均长度较高。
-
验证 batch 能高效利用 GPU。
粗略地说,平均每轮接受 token 越多,target 调用次数越少,加速越明显。
10. 接受率¶
接受率是投机解码的核心指标。 它受以下因素影响: - draft 和 target 的分布相似度。
-
temperature、top-p、top-k 等采样参数。
-
任务领域是否匹配。
-
当前上下文是否容易预测。
-
draft 长度 k。
接受率低时,draft 做了很多无效工作,投机解码可能变慢。
11. Draft Model 选择¶
理想 draft model 应该: - 比 target 明显更快。
-
和 target 分布足够接近。
-
支持同样 tokenizer。
-
适配同一业务领域。
-
显存和部署成本可接受。
太小的 draft 快但不准,太大的 draft 准但成本高。需要找平衡点。
12. Draft Length¶
Draft length k 表示每轮 draft 尝试生成多少候选 token。 k 太小: - 潜在加速有限。
k 太大: - 后面 token 更难被接受。
-
draft 成本变高。
-
target 验证计算和缓存压力变大。
实际系统通常需要根据接受率和延迟调 k。
13. 投机解码的代价¶
代价包括: - 需要额外 draft model 或额外预测头。
-
系统调度更复杂。
-
KV cache 管理更复杂。
-
低接受率时可能变慢。
-
batching 与 serving 调度更难。
-
采样语义实现容易出错。
因此投机解码不是无条件提速,需要实测。
14. 与 KV Cache 的关系¶
投机解码仍然依赖 KV cache。 需要考虑: - target model 的 KV cache 如何更新。
-
draft model 是否维护自己的 KV cache。
-
被拒绝的候选 token 的 KV 是否需要丢弃。
-
tree-based 方法如何管理候选分支 KV。
KV 管理错误可能导致结果不一致或显存浪费。
15. 与 batching 的关系¶
投机解码在单请求上容易理解,但在 serving 中会和 batching 产生复杂交互。 问题包括: - 不同请求接受长度不同。
-
target verify 的候选长度不同。
-
batch shape 更动态。
-
有些请求需要回退采样。
-
continuous batching 调度更复杂。
所以生产系统要同时优化算法和调度。
16. 什么时候效果好¶
投机解码效果好的条件: - draft 明显快。
-
draft 和 target 足够一致。
-
平均接受长度高。
-
target 验证能高效并行。
-
Decode 是主要瓶颈。
-
采样温度不太高。
例如格式化文本、代码、低温问答和强模式输出常更容易预测。
17. 什么时候效果差¶
效果差的情况: - high temperature 导致候选不稳定。
-
draft 与 target 差异大。
-
任务分布不匹配。
-
draft 运行成本接近 target。
-
serving batch 已经很满,验证收益被调度开销抵消。
-
输出很短,启动开销占比高。
此时投机解码可能不如普通 decode。
18. Medusa 的定位¶
Medusa 是投机解码的一种代表方法。它不一定需要独立的小 draft model,而是在目标模型上增加多个预测头,用这些头同时预测未来多个 token。
核心思想: - 共享 target model 的 backbone。
-
增加多个 Medusa heads。
-
每个 head 预测不同未来位置的 token。
-
用 tree attention 验证多条候选路径。
这样可以减少额外 draft model 的部署成本。
19. Medusa Heads¶
Medusa heads 是附加在模型上的多个预测头。 例如: - head 1 预测下一个 token。
-
head 2 预测下下个 token。
-
head 3 预测更远 token。
这些头可以生成多个候选 token 或候选分支。然后目标模型用特殊的验证方式检查哪些候选路径可接受。
20. Tree Verification¶
Medusa 往往不是只提出一条线性候选,而是提出一个候选树。 Tree verification 通过 tree attention mask 在一次或少数几次 forward 中验证多个候选路径。 收益: - 同时探索多个可能 token。
-
提高至少一条路径被接受的概率。
-
减少目标模型重复计算。
代价: - 验证逻辑更复杂。
- 候选树太大会增加计算和显存开销。
21. Medusa 的优势¶
优势包括: - 不需要单独部署一个 draft model。
-
共享大模型 backbone。
-
可通过少量 fine-tuning 训练预测头。
-
对服务部署更简单。
-
能结合 tree candidate 提高接受长度。
它适合不方便维护额外小模型的场景。
22. Medusa 的风险¶
风险包括: - 需要训练或适配 Medusa heads。
-
heads 预测质量影响接受率。
-
tree size 过大可能抵消收益。
-
对模型结构和推理框架有侵入。
-
采样和验证实现复杂。
Medusa 不是纯调度优化,而是模型结构和推理算法结合。
23. 常见变体¶
投机解码有多种变体: - 小 draft model。
-
n-gram/prompt lookup draft。
-
self-speculative decoding。
-
early-exit/layer skipping draft。
-
Medusa 多头预测。
-
EAGLE 等基于特征预测的方案。
它们的共同点都是“便宜地产生候选,昂贵模型负责验证或纠正”。
24. 评估指标¶
评估投机解码应关注: - speedup。
-
accepted tokens per target forward。
-
acceptance rate。
-
TTFT。
-
TPOT/ITL。
-
tokens/s。
-
质量一致性。
-
draft latency。
-
verify latency。
-
额外显存。
不能只看单请求 demo 的加速倍数。
25. 生产化注意事项¶
生产中要注意: - 是否保持 target 分布。
-
draft 和 target tokenizer 是否一致。
-
sampling 参数是否支持。
-
与 batching、KV cache、流式输出是否兼容。
-
低接受率时是否自动降级。
-
是否按业务场景做 A/B 和压测。
投机解码要有 fallback。否则在不适合的负载上可能拖慢服务。
26. 面试表达要点¶
面试中可以这样表达: - 标准 Decode 是逐 token 调 target,串行且慢。
-
投机解码用便宜 draft 先猜多个 token。
-
target 一次 forward 并行验证这些候选。
-
接受通过验证的连续 token,拒绝时按目标分布修正采样。
-
加速取决于 draft 成本和接受率。
-
Medusa 用多个预测头替代独立 draft,并通过候选树验证。
27. 核心总结¶
投机解码的本质是用便宜计算换取减少昂贵 target forward 的次数。它的正确性关键在于验证和接受-拒绝机制,性能关键在于 draft 成本、接受长度和 serving 调度。Medusa 是重要代表,它通过多预测头和 tree verification 生成并验证候选,减少对独立 draft model 的依赖。
28. 参考资料¶
-
Fast Inference from Transformers via Speculative Decoding: https://arxiv.org/abs/2211.17192
-
Accelerating Large Language Model Decoding with Speculative Sampling: https://arxiv.org/abs/2302.01318
-
Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads: https://arxiv.org/abs/2401.10774
-
Medusa GitHub: https://github.com/FasterDecoding/Medusa
-
EAGLE: Speculative Sampling Requires Rethinking Feature Uncertainty: https://arxiv.org/abs/2401.15077
预览时标签不可点<div class="