万字综述 10+ 种 LLM 投机采样推理加速方案
来源:互联网
时间:2026-08-16 14:06:16
# 投机解码方案综述:从原理到实践
## 一、背景
咱们之前聊过不少关于投机解码(Speculative Decoding)的方案,从 Parallel Decoding 到 Google 的那篇经典工作,再到 SpecInfer、Medusa、Lookahead Decoding,以及阿里最近提出的 Lookahead。最近看到一篇不错的综述文章 [2401.07851],专门梳理了投机解码领域的研究进展,下图 Figure 3 就是其中一张总结图。
今天这篇文章,我们也来做一个系统性的汇总——主要聚焦在我们实际接触过、或者已经在开源社区见到落地的方案,同时横向对比一下各自的优劣势。
## 二、自回归解码(Autoregressive Decoding)
### 2.1 概述
当前主流的大语言模型基本都是 Decoder-only 的 Transformer 架构。它的推理过程可以拆成两个阶段:
- **Prefill 阶段**:根据输入的 tokens(比如 "Recite, the, first, law, of, robotics")生成第一个输出 token(比如 "A")。这个阶段只需要一次 forward,而且输入 token 之间可以并行计算,效率很高。
- **Decoding 阶段**:从生成第一个 token 开始,模型就进入了自回归模式——一次只生成一个 token,一直生成到终止符或满足特定条件。假设最终输出有 N 个 token,那 Decoding 阶段就需要做 N-1 次 forward,而且这些 forward 只能串行执行,效率极低。更麻烦的是,随着生成进行,每个新 token 都需要关注之前所有的 token,计算量也在逐步增加。
### 2.2 KV Cache
在 LLM 推理中,最核心的计算单元就是 Multi-Head Attention。从计算图上看,主要的计算集中在两部分:一是灰色的 Linear(矩阵乘),二是 Scaled Dot-Product Attention 中的 MatMul 矩阵乘法。
这里的 Mask 是一个下三角矩阵,它正是实现 LLM Decoder 特性的关键:每个 token 只能看到当前位置及之前的 token。
把 QKT 理解成一个相关性矩阵会更容易。假设有 4 个 token,对应 4 个 step:
- Step 2 依赖 Step 1 的结果,矩阵的第一行不用重复算;
- Step 3 依赖 Step 1 和 Step 2,矩阵的前两行不用重复算;
- Step 4 依赖前三个 step,前三行都不用重复算。
Decoding 阶段 token 是逐个生成的,如果每次都重新计算之前的结果,那效率简直没法看。最直接的优化思路就是:把之前生成过程中计算好的中间结果缓存起来,下次直接用。下面这张图就直观地展示了有无 Cache 的区别。
拿 GPT2 模型在 T4 GPU 上做个测试,效果很明显:
| GPT2/T4 | 无 KV Cache | 有 KV Cache | 加速比 |
|---------|-------------|-------------|--------|
| Output Token 1000 | 52.53s | 9.12s | 5.76x |
KV Cache 当然也不是没有代价。它本质上是用空间换时间,占用的显存会显著增加。举个具体的例子:Vicuna-13B 模型(FP16 推理,Transformer layer num=40,embedding size=5120),在 Sequence Length=1024、batch size=8 的情况下,光 KV 缓存就需要 2 * 2 * 40 * 5120 * 1024 * 8 = 6.7G 的显存。
### 2.3 访存瓶颈
Decoding 阶段 token 逐个处理,加上 KV Cache 之后,Multi-Head Attention 里的矩阵乘矩阵操作,全部降级成了矩阵乘向量。
Transformer 里另一个关键组件 FFN 也是如此,它主要包含两个矩阵乘法,但 token 之间本来就不会交叉融合,所以也用不着 Cache。不过一样的问题是:矩阵乘矩阵也变成了矩阵乘向量。
矩阵乘向量操作是典型的访存 bound。LLM Decoding 阶段的这些核心计算,本质上都受制于内存带宽。
基于 V100 GPU、FP16 精度来画一个简化的 Roofline Model:
- Prefill 阶段(三角符号):假设 batch size 为 1,Sequence Length 越大,计算强度越大,通常落在 Compute Bound 区域;
- Decoding 阶段(圆形符号):batch size 越大,计算强度越大,理论性能峰值也越高,但通常落在 Memory Bound 区域。
### 2.4 Decoding 阶段优化
从下图(来自 [2308.16369] SARATHI)可以看得很清楚:
- Prefill 阶段即使在比较小的 batch size 下,也能获得可观的计算强度,吞吐自然很高;
- Decoding 阶段则需要比较大的 batch size 才能有好的计算强度。
针对 Decoding 阶段计算强度低的问题,主要有两种优化思路:
1. **不同请求间 Batching(Continuous Batching)**:通过增大 Decoding 阶段的 batch size 来充分发挥 GPU 算力。但缺点也很明显——它无法降低单个请求的延迟,而且如果请求数量很少(比如只有一个用户),那这种方案基本派不上用场。
2. **单请求内一次验证多个 Decoding Step(Speculative Decoding)**:虽然一次 Decoding Step 的计算量会线性增加,但因为 GPU 本来就处在访存瓶颈状态,计算时间不会明显增加。如果 Decoding Step 的数量能减少更多,那整体时延就有望降下来。
来看一组数据对比:batch size 为 1 和 512 时,LLM 中几个主要 OP 的计算耗时。batch size 从 1 增加到 512,计算量增加了 512 倍,但整体时间只增加了约 3 倍(数据来自 openppl-public)。如果平均每次验证能通过 5 个 token,那总的 Decoding Step 就降为原来的 1/5,总时间变成原来的 3/5——这就是单请求加速的逻辑。
当然,在实际的生产环境中,需要综合考虑这两种优化思路。如果 Continuous Batching 已经让 batch size 足够大,各种矩阵运算接近或超过了 Roofline 的交叉点,那再用 Speculative Decoding 就不会有什么收益了。这一点,恰恰是当前很多讨论投机采样时最容易忽略的。
## 三、并行解码(Parallel Decoding)
### 3.1 方案概述
Parallel Decoding 的开山之作是 2018 年发表在 NIPS 上的 *Blockwise Parallel Decoding for Deep Autoregressive Models*。作者提出了一种分块并行解码方案,在保证精度不变的情况下可以获得 2x 加速,如果愿意牺牲少量精度,甚至能到 7x。
假设输出序列长度为 m,Autoregressive Decoding 需要执行 m 步才能拿到最终结果。模型越大,每一步的时延越大,整体时延就会被放大至少 m 倍。而 Blockwise Parallel Decoding 的目标是:用 l 步完成整个预测,并且 l 远小于 m。
### 3.2 实现细节
具体方案分为三步:
**Step 1:Predict 阶段**
令 p₁ = p(p 是原始模型),然后额外训练一系列辅助模型 p₂, p₃, ..., pₖ。其中 pᵢ(yⱼ₊ᵢ|y≤ⱼ, x) 表示:第 i 个模型根据之前 1 到 j 个输出,预测第 j+i 个输出。
举个例子:假设已经生成了 "I", "saw", "a", "dog", "ride" 这 5 个 token,后续要生成 "in", "the", "bus"。那么:
- 第 1 个模型(原始模型)负责生成 "in" 所在位置的 token;
- 第 2 个模型负责生成 "the" 所在位置的 token;
- 第 3 个模型负责生成 "bus" 所在位置的 token。
**Step 2:Verify 阶段**
上一步得到了 K 个新的输出,但还不能确认它们都正确,需要进一步验证。验证的方法是:分别将前 j-1 个新 token 与输入合并,用原始模型 p₁ 预测下一个 token。如果预测结果与第 j 个 token 相同,就接受这个 token。
还是用上面的例子:假设上一步得到了 "in", "the", "car",用原始模型分三组验证(这三组可以并行执行):
- 第一组:输入 "I saw a dog ride",待验证 token 为 "in"。预测结果是 "in",匹配,接受。
- 第二组:输入 "I saw a dog ride in",待验证 token 为 "the"。预测结果是 "the",匹配,接受。
- 第三组:输入 "I saw a dog ride in the",待验证 token 为 "car"。预测结果是 "bus",不匹配,不接受。
**Step 3:Accept 阶段**
假设上一步生成了 10 个 token,在第 5 个 token 处发现不一致,那就只接受前 4 个 token,和输入合并,然后开始下一次生成。沿用上面的例子:因为第三组的 "car" 和 "bus" 不一致,所以只接受 "in" 和 "the",下一次迭代的输入就是 "I saw a dog ride in the"。
模型的实现方式是在最后一个 Transformer Decoder 层额外加几个 head,分别对应 p₂, ..., pₖ:
- Predict 阶段:原始模型 p₁ 和辅助模型 p₂,...,pₖ 相互独立,可以并行执行,耗时和生成一个 token 差不多;
- Verify 阶段:从生成的 K 个 token 里挑选最长匹配前缀。因为一次可以生成多个 token(≤K),所以能减少整体所需的 Decoding 步数;
- Accept 阶段:只接受第一个不一致之前的 token。由于验证用的是原始模型 p₁,这保证了最终结果和原始序列的预测结果完全一致(和 Greedy Decoding 的结果一致)。
并行验证过程可以一次 forward 完成。理想情况下(每次生成的 K 个 token 全被接受),总的解码次数可以从 m 降低到 2m/K。
### 3.3 评估结果
作者在 WMT 2014 英德翻译数据集上做了实验。Baseline 模型在 8 个 P100 GPU 上训练了 1,000,000 steps,使用 Greedy Decoding,在 newstest2023 development 数据集上 BLEU 得分 25.56。
针对不同的猜测 token 数,作者分别用原始数据和 Beam Search 生成的蒸馏数据训练辅助模型,每项都在相同的硬件上额外训练 1,000,000 steps。结果如下:
- **Regular ●**:冻结 backbone,使用原始数据,平均接受 token 数很小,最大只有 1.76;
- **Distillation ◼**:冻结 backbone,使用蒸馏数据,平均接受 token 数略有提升,最大 1.91,BLEU 也随之提高;
- **Fine Tuning ▲**:不冻结 backbone,使用原始数据,平均接受 token 数增大,最大 3.01;
- **Both ◆**:不冻结 backbone,使用蒸馏数据,平均接受 token 数明显增大,最大达到 4.95,BLEU 也相应提高。
在机器翻译任务上,当 k=8 时,平均接受 token 数为 4.7,整体加速比达到 3.3 倍。
## 四、投机解码(Speculative Decoding)
### 4.1 方案概述
Google 和 Deepmind 在 2022 年提出了投机采样方案 *Fast Inference from Transformers via Speculative Decoding*。思路其实很简单:用一个高效的小模型先生成多个候选 token,然后再让大模型去验证。
### 4.2 实现细节
设 Mₚ 为目标模型(大模型),给定前缀输入 x_{
作者用一个包含不同 γ 取值的简单示例做了说明:紫色是执行目标模型 Mₚ 的 decoder,蓝色是执行近似模型 M_q 的 decoder,黄色和橙色是调用 encoder。
### 4.3 评估结果
作者基于 T5X 代码库验证了 T5-XXL 模型的加速效果。实验设置如下:
- **模型**:标准的 encoder-decoder T5 1.1 版本
- 目标模型 Mₚ:T5-XXL(11B)
- 近似模型 M_q:T5-Large(800M)、T5-Base(250M)、T5-Small(75M)
- **任务**:
- 英语到德语翻译(WMT EnDe 数据集微调)
- 文本总结(CCN/DM 数据集微调)
- **硬件**:TPU-v4
- **推理参数**:batch-size = 1,分别测试 argmax sampling(temp=0)和 standard sampling(temp=1)
结果如下表所示。最小的近似模型 T5-Small(75M)获得了最高的加速比——模型越小,推理越快,而且生成质量相比 T5-Base 没有下降太多。在 EnDe 任务上,temp=0 时获得 3.4 倍加速,temp=1 时获得 2.6 倍加速。
Huggingface 官方的测试结果(Assisted Generation)也印证了这一点:
- Assistant Model:facebook/opt-125m
- 目标模型:facebook/opt-1.3b、6.7b、30b、66b
- 数据集:C4 (en, validation set)
## 五、SpecInfer
### 5.1 方案概述
SpecInfer([2305.09781])的核心思路是:通过一系列小模型 SSM(Small Speculative Model)联合预测 LLM 的输出,并把它们的预测结果组织成一棵 Token 树,树中每个分支对应一个候选 token 序列。然后 LLM 用基于树的并行解码(Tree-Based Parallel Decoding)机制来一次性验证整棵树所有 token 的正确性。
和传统方案不同,SpecInfer 把 LLM 当作 token 树的验证器,而不是增量式的解码器。这种方式能显著降低端到端延迟,同时保持模型质量。作者评估结果显示,相比现有的 LLM 服务框架(截至 2023 年 5 月),SpecInfer 的分布式推理性能可以提升 1.3-2.4 倍,如果使用 offload 机制,甚至可以提升 2.6-3.5 倍。
### 5.2 实现细节
SpecInfer 的一个核心是预测器(Speculator)的设计。这里面临一个两难:一方面,更准确的预测器能生成更长的匹配序列,有助于减少 Decoding 步数;另一方面,句子中的某些短语容易推测,有些却很有挑战性。固定的配置(比如 beam search 的宽度和深度)会导致性能不足——窗口太小可能错过更长的匹配序列,窗口太大又会生成大量无效 token。
SpecInfer 采用两个关键技术来解决这个挑战:
1. **Collective Boost-Tuning**:一种新的微调技术,通过自适应 boosting 使一组 SSM 的聚合预测结果与 LLM 输出对齐;
2. **可学习的推测调度器**:学习给定输入 token 序列,找到最优的一组 SSM 及其推测配置。
SpecInfer 的另一个重要工作是 Tree Based Parallel Decoding 机制。简单来说:
- **Sequence-based Decoding**:最直观的方式是并行处理多个子序列,但这种方式会造成极大浪费——每个子序列都要计算并保留各自的 key-value Cache;
- **Tree-based Parallel Decoding**:通过精心设计的 Attention Mask,一次性验证所有 token,计算量更小,空间占用也更少。
### 5.3 评估结果
作者在 OPT-30B、LLaMA-30B 和 LLaMA-65B 上对比了 SpecInfer 与 vLLM、Huggingface TGI、FasterTransformer。所有模型都部署在 2 台 4×A10 GPU 的机器上,机间使用 Pipeline 并行,机内使用 Tensor 并行,采用半精度计算。vLLM 和 TGI 不支持流水线并行和多机部署,所以它们在单机上运行。
结果提升很明显:单节点相比现有系统提升 1.3-2 倍,2 节点提升 1.4-2.4 倍。
在 token 接受率方面,使用 LLaMA-160M 作为 SSM、LLaMA-7B 作为 LLM 进行验证。Token 树包含 16 个 token,随着 SSM 数量从 1 增加到 5,平均接受 token 数从 2.92 提升到 3.58,平均接受率约为 1/5。
## 六、Medusa
### 6.1 方案概述
Medusa([2401.10774])可以看作是 Blockwise Parallel Decoding 和 SpecInfer 的结合体,并对前者的多 Head 做了升级。原来的方案是一个 Head 生成一个 token,而 Medusa 让它变成一个 Head 生成多个候选 token——因为作者观察到,预测 next next token 时 top1 的准确率可能只有 60%,但 top5 有可能超过 80%。然后根据这些 Head 生成的 token 构建笛卡尔积,形成多个候选 token 序列,后续采用 SpecInfer 的 Token 树验证机制来完成验证。
在早期的 Medusa-1 版本中,LLM 的 Backbone 和 LM Head 都是固定的,只微调新增的 Medusa Head。到了 Medusa-2 版本,LLM 的 Backbone 也会参与联合微调,进一步提升了速度。
### 6.2 实现细节
Medusa 在 LLM 的最后一个 Transformer Layer 之后保留了原始的 LM Head,额外增加多个 Medusa Head,从而获得多个候选 token 序列,再通过一次 Decoding Step 完成验证。
来看 Attention Mask 矩阵的设计:假设 Head 1 在下一个位置生成 2 个可能的 token("It" 和 "I"),Head 2 在下下个位置生成 3 个可能的 token("is", "'", "the")。这样下一个位置和下下个位置就有 2 × 3 = 6 种可能的候选序列。对应的 Attention Mask 矩阵也有相应的结构。
Token 树中有多少个 token?假如有 3 个 head,第一个 head 有 3 个候选,第二个有 5 个,第三个有 7 个,那就有 3 × 5 × 7 = 105 个候选 token 序列。合并成 token 树后,总共有 3 + 3×5 + 3×5×7 = 123 个 token。这会把计算量放大到 123/3 ≈ 41 倍。虽然 LLM 推理在 batch size 较小时是 IO bound,可以利用空闲算力,但并不意味着可以无限制地增加计算量。
### 6.3 评估结果
使用 MT-bench 进行模拟真实场景的评估。结果如下:
- Medusa-1 通过相对简单的设置可以获得 2.18x 和 2.33x 的加速,33B 模型的速度和原始方案的 13B 模型相当;
- Medusa-2 进一步提升,加速比达到 2.83x。
作者还验证了不同验证 token 数的影响:当待验证 token 数为 50-100 时获得最大速度,但加速比也只有 2.5-3x。
## 七、REST(基于检索)
### 7.1 方案概述
REST([2311.08252])和 Medusa 出自同一个团队。与 Medusa 需要草稿模型不同,REST 借鉴了 RAG 的思路——根据当前上下文从知识库中检索相关 token,然后基于这些 token 来生成草稿序列。它可以无缝插入到现有 LLM 中,无需额外训练。在 batch size 为 1 的情况下,7B 和 13B 模型分别实现 1.62x 和 2.36x 的加速。
### 7.2 实现细节
思路相当直接:先用输入的 prompt 去检索数据库(可以自定义),然后用检索结果构建草稿 token 树,最后和 Medusa 一样使用 Tree Attention 通过一次 Decoding 验证。
### 7.3 评估结果
作者在 HumanEval 和 MT-Bench 上做了评估:
- CodeLlama 模型上获得 2.26-2.36x 的加速;
- Vicuna 模型上获得 1.62-1.77x 的加速。
知识库大小对平均接受 token 数的影响:知识库越大,接受的 token 越多,但 25GB 的知识库平均也只接受约 2.6 个 token。
草稿 token 数的影响:草稿 token 越多,平均接受数越高。但 50 个草稿 token 只接受约 2.6 个,200 个只接受约 2.8 个,性价比相对较低。
## 八、前向解码(Lookahead Decoding)
### 8.1 方案概述
Lookahead Decoding 利用 Jacobi 迭代法,直接使用 LLM 同时提取和验证 n-grams,打破了自回归解码的顺序依赖性。相比之前的并行解码方案,它不需要草稿模型就能降低解码次数。
一个示例:标准 Autoregressive Decoding 生成速度为 34.83 tokens/s,而 Lookahead Decoding 为 60.69 tokens/s,几乎是 2 倍,而且生成的结果完全一致。
一个 Decoding Step 大致包含以下步骤:
1. **Parallel Decoding**:一次 forward,生成候选 token 对应的待验证 token 序列;
2. **Verify**:对比待验证 token 与候选 token,确定最长的正确序列;
3. **Collect N-Grams**:将未验证通过的候选 token 和对应生成的 token 组成 N-Gram 序列,添加到 N-Gram Pool 中;
4. **Update**:用生成的待验证 token 序列更新候选序列;
5. **Match N-Grams**:用候选序列中的 token 依次从 N-Grams 中匹配,并替换候选序列。
### 8.2 实现细节
#### 8.2.1 Lookahead
Lookahead 的目的是生成新的 N-Grams,它由两个参数定义的二维窗口来操作:
- **window size W**:在未来的 token 位置上向前多远,以进行并行解码(计算量随 W 线性增加);
- **N-Gram size N**:回顾多少步之前的 Jacobi 迭代来检索 N-Gram。
举个例子:回顾 4 个 Step,展望 5 个 Token。在当前步骤 t,使用前 3 个步骤形成的轨迹做一次 Jacobi 迭代,为所有 5 个位置生成新 token。然后收集 4-gram(比如位置 1 的橙色、位置 2 的绿色、位置 3 的红色 token,加上当前 step 新生成的黄色 token "4")。随着解码进行,轨迹中最旧的 token 会被删除,以维持 N 和 W 的恒定。当 N=2 时,Lookahead 解码等价于 Jacobi 解码。
#### 8.2.2 Verify
每个解码步骤除了 Lookahead 分支,还有 Verify 分支,目的是识别和确定有希望的 N-Gram。在 Verify 分支中,用 N-Gram 的第一个 token 去匹配输入的最后一个 token,一旦匹配上,就把对应的 N-Gram 添加到当前输入,然后通过 LLM forward 来验证。随着 N-Gram 缓存增加,会有多个相同 token 开头的 N-Gram 出现,这增加了验证成本。为了控制成本,作者将验证分支中候选 N-Gram 数量的上限设为 G,通常与 W 成正比,比如 G=W。
#### 8.2.3 Lookahead 和 Verify 同时完成
LLM 解码主要受内存带宽限制,因此可以在一个 Step 内合并 Lookahead 和 Verify,利用 GPU 的并行处理能力来隐藏开销。作者通过设计特殊的注意力掩码来实现,这个掩码遵循两个原则:
1. 每个 token 只能看到它之前的 token(因果掩码);
2. Lookahead 分支中的 token 看不到 Verify 分支中的 token,反之亦然。
以 W=5, N=4, G=5 为例:
- Lookahead 分支:一次 decoding 原本只需执行 "0",现在多了虚线框内 W×(N-1)-1 个 token,计算量是原来的 W×(N-1) 倍。新生成的 token 可用于构建下一次 Verify 分支的候选序列;
- Verify 分支:对应 G 个 N-Gram 候选,其中的 4-Gram 第一个 token 都对应 "0",所以计算量是原来的 G×(N-1) 倍,用于更新下一次 Lookahead 分支的序列。
### 8.3 评估结果
#### 8.3.1 缩放法则
作者验证了不同 W 和 N 对效率的影响。当 N 足够大(比如 11)时,随着 W 增加,解码步数几乎可以线性降低。
#### 8.3.2 代价和局限性
前面提到过,Lookahead 部分计算量是原来的 W×(N-1) 倍,Verify 部分是原来的 G×(N-1) 倍,所以一个 Step 的总计算量是原来的 W×(N-1) + G×(N-1) 倍。作者在 A100 上测试了比较好的配置:对于 7B、13B、33B 模型,每个 step 的计算量分别是原来的 120 倍、80 倍和 56 倍。
- **适合的场景**:算力存在极大浪费的情况,比如 batch size 为 1 时明显的 IO bound。只有 batch size 达到 64 甚至 100 以上才能充分发挥算力,此时增加的计算量正好可以利用起来;
- **不适合的场景**:当推理服务本身已经通过 Continuous Batching 达到了比较大的 batch size(比如 16 或 32),留给 Lookahead Decoding 的空间就很小了。
## 九、EAGLE(基于特征)
### 9.1 方案概述
EAGLE(Extrapolation Algorithm for Greater Language-model Efficiency,[2401.15077])由北京大学和微软等团队提出,是一种无损的投机采样加速方案。与传统方案不同,EAGLE 利用最后一个 Transformer Block 的输出特征来自回归地生成草稿 token,并通过提前一步集成 token 来解决下一个特征预测中的采样不确定性问题。
EAGLE 提供的是无损加速,不需要微调原始 LLM(但会增加一个小型的草稿模型),生成的文本保持与普通自回归解码相同的分布。使用 gpt-fast,EAGLE 在单个 RTX 3090 上就能让 LLaMA2-Chat 13B 达到平均 160 tokens/s,而 Huggingface 的实现只有 24 tokens/s。
### 9.2 实现细节
EAGLE 没有用一个小 LLM 去预测后续 token,也没有用多个 LM head,而是借助之前三个 token 的输入和输出 Embedding 来预测下一个 token 的 Embedding,再通过 LM Head 生成 token。
假设已经生成的 token 为 "How", "can", "I":
**Draft model 第一次 Forward**(用 "How", "can", "I" 预测下一个):
- 将 "How" 在原始 LLM 最后一个 Transformer Layer 的输出 Embedding f_how 与 "can" 对应的第一个 Transformer Layer 的输入 Embedding e_can 组合成第一个 Embedding;
- 将 "can" 的输出 Embedding f_can 与 "I" 的输入 Embedding e_I 组合成第二个 Embedding;
- 用这两个 Embedding 经过 One Auto-regression Head 预测 "I" 对应的输出 Embedding f_I;
- 用 f_I 经过原始 LLM 的 LM Head 和 Sampling 过程预测下一个 token(比如 "make"/"help")。
**Draft model 第二次 Forward**(用 "How", "can", "I", "make"/"help" 预测下一个):
- 和第一次类似,在上一步基础上扩增了新生成的 token;
- 从 "How", "can", "I", "make" 预测得到 "a"/"out";
- 从 "How", "can", "I", "help" 预测得到 "with"/"you"。
**Draft model 第三次 Forward**:
- 继续扩增,预测下一层的 token。
Draft model 的 Embedding 层、LM Head 和 Sampling 都来自原始 LLM,而 Auto-regression Head 只包含一个 FC Layer(用于降维)和一个 Transformer Decoder Layer。
和其他投机采样方案相比,EAGLE 的待验证树更稀疏(可以避免一些无效生成),待验证的 token 也更少。
### 9.3 评估结果
在 MT-bench 上,EAGLE 比基线 Decoding 快 3 倍,比 Lookahead 快 2 倍,比 Medusa 快 1.6 倍。
EAGLE 采用固定的 3 个 Draft Forward,即待验证的 token 最多 3 层,一次验证通过的最大 token 数为 4。作者验证了序列模式(without Tree Attention)和树模式(with Tree Attention)的差异。
树模式因为每次验证的 token 是序列模式的超集(3 个 → 10 个),加速比更高(访存瓶颈下增加的延迟小),但提升相对有限。
在 MT-bench 上:
- 树模式的平均 token 接受数接近 4;
- 序列模式也有约 3 的接受数。
当然,树模式会明显增加验证阶段的计算量——由 Chain 模式的 3 个 token 变为 10 个,而接受的 token 数只增加了 0.6-0.7。需要根据 Roofline 模型来确定该用哪种方案。即使只用 Chain 模式,也能获得 2.2-2.7x 的加速。
作者也验证了 EAGLE 与 Batching 的结合:随着 Batch Size 增加,加速比会适当降低。这和我们之前的结论一致(这里用的是 Sampling Temperature 为 0,如果增大,加速比会进一步降低)。
## 十、阿里 Lookahead(针对 RAG 等)
### 10.1 方案概述
LLM 受训练语料的限制,无法感知最新内容(比如训练之后发生的新闻),也容易产生幻觉。为了解决这个问题,业界提出了 RAG(检索增强生成)范式:当用户发起生成请求时,先根据 prompt 检索外部知识库,把检索到的信息组合到 prompt 中,然后再输入 LLM 生成结果。
在 RAG 系统中,LLM 生成的内容很可能就来自 prompt 中之前检索到的内容。这天然适合作为投机采样方案中猜测的 token 序列,避免了额外模型或额外 Head 来生成待验证 token。阿里提出的 Lookahead 方案([2312.12728])正是基于这个思路设计的。
### 10.2 实现细节
整体思路和之前的投机采样方案类似,关键区别在于待验证 token 的来源——从 prompt 中直接构建。与单序列相比,多序列能提升接受率,token 前缀树则能进一步降低成本。
具体实现是通过设计特殊的 Mask,一次验证多个 token 序列或 token 前缀树。这种思路在 SpecInfer 和 Medusa 中也用过。
### 10.3 评估结果
推理加速结果如下:
- 在 AntRAG 上效果最明显,达到 5 倍;
- 在 Dolly 上稍差一些,也有 2 倍。
Baseline 用的是 Huggingface 的 Transformer 库(性能可能偏低,如果用 vLLM 或 TensorRT-LLM 可能会有不同结果)。LLMA 是微软发布的方案,Lookahead(Parallel) 是多分支方案,Lookahead(Hierarchical) 是前缀树方案。
作者在 GitHub 上还提供了其他模型在 Dolly-15k 和 GSM-8k 上的测试结果,提升同样在 2 倍左右。其中 decoding length(生成 token 长度)为 64,branch length(并行验证的 token 数)为 8。
不同 decoding length 和 branch length 下的 EDL(Effective Decoding Length,接受的 token 数):branch length 越长,接受的 token 数越多。当 branch length 为 30-40 时,接受的 token 数在 9-12 之间,基本能达到 1/4,冗余计算相比之前方案少了很多。
## 十一、其他优化方案
### 11.1 Staged Speculative Decoding
在基于草稿模型的方案中,草稿模型和原始 LLM 的差异会直接影响接受 token 数。草稿模型越大,猜测越准,但效率也越不可忽略。通常草稿模型是原始 LLM 的 1/20-1/15 大小,生成草稿时的效率也可能偏低。
[2308.04623] 提出了在草稿模型中也使用投机采样的思路,进一步优化效率。相比传统投机采样,可以再加速 40%-50%。
### 11.2 DistillSpec
[2310.08461] 采用蒸馏方案来对齐草稿模型和原始 LLM。在 Greedy Decoding 等采样方案下,相比传统投机解码(Standard SD)可以获得 10-45% 的加速。
如果可以接受有损 Decoding,蒸馏方案可以进一步将解码延迟降低 6-10 倍,且损失相对较小。
## 十二、总结
### 12.1 概述
我们把这一系列试图降低 Decoding Step 的方式统称为投机解码。大体思路都分为两步:
1. 通过某种方式生成一系列候选 token 序列(可以是序列、树、多头树等);
2. 并行地通过一次 Forward 验证这些候选序列。
只要平均每个 Step 验证的 token 数 > 1,就能降低总的 Decoding 步数。
### 12.2 生成序列的方式
极端情况下,假设词表大小为 1000,窗口大小为 5,那窗口内 5 个 token 所有可能的组合有 1000⁵ 种。如果算力无限,一次 Decoding Step 就能验证所有组合。当然实际词表可能有几万甚至十几万,算力也有限,关键在于如何把组合数降到合理范围。
#### 12.2.1 多头方式
既然可以根据输入序列预测下一个 token,那也可以预测下下一个、下下下一个,只是准确率会低一些。这样在 Decoding Step 的同时就能额外生成候选序列,下次 Decoding Step 来验证。
#### 12.2.2 小模型生成
小模型也能做序列生成,只是效果没那么好。比如 LLM 是 100B,小模型是 100M,在 LLM 的每个 Decoding Step 先用小模型生成 10 步,再用 LLM 并行验证。
#### 12.2.3 利用历史记录
对于确定的上下文,后续生成中很可能包含某些历史短语序列或错误生成序列中的短语。比如 "He is a very very famous computer scientist",在 "He is a" 处如果能额外生成一些序列,后续生成时就有可能用上。
RAG 任务中,生成结果很容易包含检索出来的子序列,这相当于增加了猜中的概率。
#### 12.2.4 利用知识库
有的方案通过外部知识库生成候选序列,和 RAG 类似,但检索到的语料是用于构建草稿 token 树,而不是加到 prompt 里。
### 12.3 方案对比
从以下几个维度对上述方案进行汇总:
- **是否有损**:能否与原始 LLM 保持完全一致(受采样策略影响,这里主要指 Greedy Decoding 场景);
- **LLM 微调**:是否需要微调原始 LLM(多头不归为 LLM 微调);
- **草稿**:生成待验证草稿 token 的方式;
- **Token 接受率**:平均接受的 token 数/平均验证的 token 数,接受率越高,无效计算越少。有损方案通常能提升接受率;
- **加速对比**:加入投机采样后的整体生成速度/原始生成速度。有损方案通常能提升加速比。
### 12.4 约束和限制
投机采样方案的迭代还在继续,选型时要充分考虑以下因素:
#### 12.4.1 加速比和计算量的平衡
投机采样本质上是牺牲计算来换取解码步数减少。留给它的空间有多大,取决于计算强度所处的位置:
- **黄球区域**:明显的 Memory Bound,适当增加计算量不会明显增加延迟,留给投机解码的空间大,比如 batch size 为 1 的场景;
- **绿球区域**:仍属 Memory Bound,但计算强度已经较大,离算力峰值不远,留给投机解码的空间相对小,比如 Online Inference 场景使用 Continuous Batching 后;
- **红球和蓝球**:位于 Compute Bound,增加计算量会导致时延明显增加,投机解码基本没有收益。
#### 12.4.2 实现复杂度
有些方案需要添加多头或微调 LLM,这些是一次性成本,但不能忽略数据收集、训练、评估等环节的投入。
#### 12.4.3 与现有 LLM 推理框架的适配
当前主流推理框架(Huggingface TGI、vLLM、TensorRT-LLM、LMDeploy 等)集成了 PagedAttention、FlashDecoding、KV-Cache 等高级特性,性能远高于原生的 Huggingface Transformer。如果投机解码方案与这些框架不兼容,带来的收益可能不如直接用高效框架。
#### 12.4.4 应用场景
同一方案在不同场景、甚至不同模型上的表现差异都很大。真正落地时,一定要针对自己的场景充分测试,以实测结果为准。