跳转至

五十六: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,再让目标模型一次性验证。 img

核心目标: - 减少目标大模型的自回归调用次数。

  • 一次 target forward 尽可能产出多个 token。

  • 在保持生成质量或分布等价的前提下提速。

它主要优化 Decode 阶段,对 Prefill 的帮助通常有限。

3. 自回归 Decode 为什么慢

标准 Decode 每生成一个 token 都要调用一次目标模型:

token 1 -> target forward
token 2 -> target forward
token 3 -> target forward
即使有 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. 从拒绝位置重新采样或回退
如果候选质量高,一次 target forward 可以接受多个 token,从而减少 target 调用次数。

5. 为什么 target 可以一次验证多个 token

Transformer 在给定一段候选序列时,可以用 causal mask 一次前向计算每个位置的 logits。

img

先分清两个场景的输入差异 · 场景 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。

接受概率通常为:

accept_prob = min(1, p(x) / q(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。

img

核心思想: - 共享 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="