投机推理:小模型搭树,大模型一次验证
大模型逐 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 个,假设是 A1 和 A2。小模型分别在这两个节点后面跑一次 forward,各取 top-2:
R
┌─────┬─────┬─────┬─────┬─────┐
A1 A2 A3 A4 A5 A6
/ \ / \
B11 B12 B21 B22A1(41824) 下面:41760, 35740
A2(43744) 下面:45315, 41760A3 ~ A6 仍然是合法候选,只是不再往深处扩展——没有子节点了。
第三步:再选 2 个父节点,各展开 1 个候选
从第二层 4 个节点中,按同样的规则选出 2 个——假设选中 B11 和 B21。小模型各取 top-1:
R
┌─────┬─────┬─────┬─────┬─────┐
A1 A2 A3 A4 A5 A6
/ \ / \
B11 B12 B21 B22
| |
C111 C211B11(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 = 37539Draft Step 1:两个父节点各取 Top-2
draft_step_1_124_140 有两个有效 row。topHead[1]=2,从 6 个第一层节点里选出概率最高的 2 个继续展开——命中了 A1 和 A2:
row 0(对应 A1=41824 的下一步): [41760, 35740, ...]
row 1(对应 A2=43744 的下一步): [45315, 41760, ...]第二层:
A1 = 41824
├── B11 = 41760
└── B12 = 35740
A2 = 43744
├── B21 = 45315
└── B22 = 41760Draft 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 后面也匹配了 45315,B21 后面也匹配了 22870,但接受只能沿一条路径走。
从根开始逐级判断:
R 想生成 41824 → 命中 A1 ✓
A1 想生成 41760 → 命中 B11 ✓
B11 想生成 22870 → 命中 C111 ✓最终接受路径:
R → 41824 → 41760 → 228703 个 draft token 全部被接受。
另一条分支 A2 → 45315 → 22870 虽然局部也有匹配,但根节点已经选了 A1,A2 分支直接被丢弃——不能把不同分支上的匹配数加起来算接受数。正确逻辑必须沿一条树路径顺序判断:根节点命中哪个子节点,就只沿那个子节点的子树继续。
大模型一次验证整棵树
关键来了:大模型不是沿每条路径分别推理。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 的唯一子节点 C111。C111 被接受。
最终接受路径:
R → 41824 → 41760 → 228703 个投机 token 全部命中。如果实现支持 bonus token,还可以从叶子节点的大模型 logits 里再产出一个,本轮最多提交 4 个 token。
另一个分支 A2 → 45315 → 22870 虽然局部也有匹配,但它不在主路径上——根节点选了 A1,A2 分支就作废了。
为什么不全展开:指数爆炸
如果每一步每个节点都保留 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 重新开始一轮投机,或者回退到上一个命中点从那里分叉。具体看实现。
总结
投机推理不是”小模型猜一串、大模型逐个比”,而是:
- 小模型搭一棵宽度逐层递减的候选树。
- 大模型用 tree attention 一次 forward 并行验证所有节点。
- 沿树的一条路径顺序判断接受,不接受的分支直接丢弃。
- 树的形状由
topKSpec和topHead精确控制,本质上是在”覆盖面”和”验证开销”之间做 trade-off。