投机推理:小模型搭树,大模型一次验证

12 min

大模型逐 token 生成太慢,投机推理(Speculative Decoding)用一个小模型先猜几个 token,大模型再一次性验证。但这不只是”猜一串 token 然后逐个比对”那么简单——真正有效的做法是搭一棵树,让大模型用一次 forward 并行验证整棵树上的所有候选。

投机推理要解决什么问题

大模型自回归生成,每出一个 token 都要做一次完整的 forward。一个 forward 和生成 100 个 token 的 forward,计算量差不了太多——瓶颈在显存带宽,不在计算。所以一次 forward 只出一个 token,GPU 大部分时间在等数据搬来搬去。

投机推理的思路很直接:让一个轻量的小模型先跑几步,把猜出来的 token 打包送给大模型,大模型一次 forward 验证所有候选,命中的直接收下,没命中的从大模型自己的分布里采样。理想情况下,大模型跑一次能产出多个 token,吞吐直接翻倍。

但问题来了:小模型的 top-1 经常猜错。如果只猜一条链,第一个 token 错了后面全废。所以要猜多条路径——也就是一棵树。

树形投机怎么配置

以这份配置为例:

{
  "gamma": 3,
  "topKSpec": [
    [6],
    [2, 2],
    [1, 1]
  ],
  "topHead": [1, 2, 2],
  "treeType": 0,
  "gammaThreshold": 0
}

三个参数控制树的形状:

  • gamma: 3 — 小模型连续跑 3 步,树的最大深度是 3。
  • topHead: [1, 2, 2] — 第 i 步有几个父节点需要继续展开。第一步只有根节点 1 个,第二步从上一层选 2 个节点继续展开,第三步再选 2 个。
  • topKSpec: [[6], [2, 2], [1, 1]] — 第 i 步的每个父节点保留几个候选子节点。第一步根节点取 top-6,第二步两个父节点各取 top-2,第三步两个父节点各取 top-1。

总结成一句话:

第一步:1 个父节点 × 6 个候选 = 6 个节点 第二步:2 个父节点 × 2 个候选 = 4 个节点 第三步:2 个父节点 × 1 个候选 = 2 个节点 总共 12 个投机节点

第一步:根节点展开 6 个候选

小模型在根节点 R(已经确认的上下文)跑一次 forward,取 logits 的 top-6:

                   R
     ┌─────┬─────┬─────┬─────┬─────┐
     A1    A2    A3    A4    A5    A6

比如小模型 top-6 是:

41824, 43744, 44298, 34051, 28380, 37539

这里是 6 个候选。但要不要全部继续向下展开,由下一轮的 topHead[1] = 2 决定——只有 2 个会被选中继续走。

第二步:选 2 个父节点,各展开 2 个候选

从 6 个第一层候选中,按照路径概率选出最值得继续展开的 2 个,假设是 A1A2。小模型分别在这两个节点后面跑一次 forward,各取 top-2:

                   R
     ┌─────┬─────┬─────┬─────┬─────┐
     A1    A2    A3    A4    A5    A6
    /  \  /  \
  B11 B12 B21 B22
A1(41824) 下面:41760, 35740
A2(43744) 下面:45315, 41760

A3 ~ A6 仍然是合法候选,只是不再往深处扩展——没有子节点了。

第三步:再选 2 个父节点,各展开 1 个候选

从第二层 4 个节点中,按同样的规则选出 2 个——假设选中 B11B21。小模型各取 top-1:

                   R
     ┌─────┬─────┬─────┬─────┬─────┐
     A1    A2    A3    A4    A5    A6
    /  \  /  \
  B11 B12 B21 B22
   |       |
  C111    C211
B11(41760) 下面:22870
B21(45315) 下面:22870

投机阶段结束。总共 12 个候选节点,分布在三层。

一个真实例子:从 dump 看三轮投机

上面用的是示意数字,下面是一份真实 dump 数据。小模型是三层的,每层输出一个矩阵文件。

Draft Step 0:根节点取 Top-6

draft_step_0_123_139 的 row 0(根节点输出)为:

[41824, 43744, 44298, 34051, 28380, 37539, ...]

第一层 6 个候选:

A1 = 41824
A2 = 43744
A3 = 44298
A4 = 34051
A5 = 28380
A6 = 37539

Draft Step 1:两个父节点各取 Top-2

draft_step_1_124_140 有两个有效 row。topHead[1]=2,从 6 个第一层节点里选出概率最高的 2 个继续展开——命中了 A1A2

row 0(对应 A1=41824 的下一步): [41760, 35740, ...]
row 1(对应 A2=43744 的下一步): [45315, 41760, ...]

第二层:

A1 = 41824
├── B11 = 41760
└── B12 = 35740

A2 = 43744
├── B21 = 45315
└── B22 = 41760

Draft Step 2:两个父节点各取 Top-1

draft_step_2_124_140 两个 row 的 top-1 相同:

row 0(对应 B11=41760 的下一步): 22870
row 1(对应 B21=45315 的下一步): 22870

第三层:

B11 = 41760
└── C111 = 22870

B21 = 45315
└── C211 = 22870

完整投机树

把三轮拼起来,就是这棵树:

                                 R
                     prefill top-1 = 43107

     ┌───────────┬───────────┬───────────┬──────────┬──────────┬──────────┐
     │           │           │           │          │          │
 A1=41824    A2=43744    A3=44298    A4=34051   A5=28380   A6=37539
    /  \        /  \
B11=41760  B12=35740  B21=45315  B22=41760
   |                    |
C111=22870          C211=22870
第一层:6 个节点
第二层:4 个节点
第三层:2 个节点
总计:12 个投机节点

Validation:大模型一次验证,命中一条路径

大模型跑一次 forward(validate_124_140),在各父节点位置产出的 top-1 为:

R       → 41824    (验证第一层 A1~A6)
A1      → 41760    (验证 B11、B12)
A2      → 45315    (验证 B21、B22)
B11     → 22870    (验证 C111)
B21     → 22870    (验证 C211)

一共 5 个父节点位置的 logits,覆盖全部 12 个候选。注意 A2 后面也匹配了 45315B21 后面也匹配了 22870,但接受只能沿一条路径走。

从根开始逐级判断:

R      想生成 41824  →  命中 A1        ✓
A1     想生成 41760  →  命中 B11       ✓
B11    想生成 22870  →  命中 C111      ✓

最终接受路径:

R → 41824 → 41760 → 22870

3 个 draft token 全部被接受。

另一条分支 A2 → 45315 → 22870 虽然局部也有匹配,但根节点已经选了 A1A2 分支直接被丢弃——不能把不同分支上的匹配数加起来算接受数。正确逻辑必须沿一条树路径顺序判断:根节点命中哪个子节点,就只沿那个子节点的子树继续。

大模型一次验证整棵树

关键来了:大模型不是沿每条路径分别推理。12 个节点被打包成一次输入,配合 tree attention mask 和 position id,大模型跑一次 forward,同时产出所有父节点位置的 logits。

逻辑上等价于:

R 的 logits       → 用来验证 A1~A6
A1 的 logits      → 用来验证 B11、B12
A2 的 logits      → 用来验证 B21、B22
B11 的 logits     → 用来验证 C111
B21 的 logits     → 用来验证 C211

一共只需要 5 个父节点位置的 logits(1 + 2 + 2),就能覆盖全部 12 个候选。

接受过程:沿一条路径走

有了大模型的 logits,接受判断是沿树的一条路径顺序进行的,不是把所有匹配的行加起来。

举个例子。假设某次实际运行中,大模型在各个父节点位置想生成的 token 是:

R     → 41824
A1    → 41760
A2    → 45315
B11   → 22870
B21   → 22870

从根节点开始:

第一级:大模型在 R 后想生成的是 41824,这在小模型猜的 6 个候选中(命中了 A1)。A1 被接受,沿着 A1 的子树继续。

第二级:大模型在 R → A1 后想生成 41760,在 A1 的子节点 [41760, 35740] 中。B11 被接受,沿 B11 子树继续。

第三级:大模型在 R → A1 → B11 后想生成 22870,等于 B11 的唯一子节点 C111C111 被接受。

最终接受路径:

R → 41824 → 41760 → 22870

3 个投机 token 全部命中。如果实现支持 bonus token,还可以从叶子节点的大模型 logits 里再产出一个,本轮最多提交 4 个 token。

另一个分支 A2 → 45315 → 22870 虽然局部也有匹配,但它不在主路径上——根节点选了 A1A2 分支就作废了。

为什么不全展开:指数爆炸

如果每一步每个节点都保留 top-6:

第 1 层:6
第 2 层:6 × 6 = 36
第 3 层:36 × 6 = 216
总计:258

大模型要一次验证 258 个节点,attention mask 的构造开销和显存占用会吃掉所有收益。当前配置 12 个节点,用很小的开销换最多 3 个 token 的连续命中。

配置的取舍逻辑

[6] → [2,2] → [1,1] 这个漏斗形状有明确的意图:

  • 第一层宽(6):第一个 token 的不确定性最高,覆盖面要大,提高根节点命中率。
  • 第二层收窄(2+2):只有最有希望的两个分支继续展开,控制规模。
  • 第三层极窄(1+1):已经走了两步正确路径,第三个 token 大概率就是 top-1,不需要额外的候选。

如果某一步大模型的 top-1 不在小模型的候选中,这条路径就断了。断了之后有两种处理:直接从大模型 logits 采样新 token 重新开始一轮投机,或者回退到上一个命中点从那里分叉。具体看实现。

总结

投机推理不是”小模型猜一串、大模型逐个比”,而是:

  1. 小模型搭一棵宽度逐层递减的候选树。
  2. 大模型用 tree attention 一次 forward 并行验证所有节点。
  3. 沿树的一条路径顺序判断接受,不接受的分支直接丢弃。
  4. 树的形状由 topKSpectopHead 精确控制,本质上是在”覆盖面”和”验证开销”之间做 trade-off。