← 返回首页
16_speculative_decoding_viz.html
投机解码与 MTP:先便宜地猜,再一次验收
第 5 章的自回归生成,每吐一个字都要把主模型完整跑一遍。
投机解码换一种分工:
便宜的草稿员先猜 γ 个字,主模型一次 forward 把这 γ 个位置全算出来,再按一条接受规则逐个验收。
草稿员可以是一个小模型,也可以是主模型自带的
MTP 头。
本关两种都真训一遍,再在 RTX 4090 上量账。配套代码 phase4-efficiency/14_mtp.py。
这一关
主模型 4 层 · 0.743M n_embd 128,第 12 章 LLaMA 骨架
草稿模型 1 层 64 维 0.055M,主模型的 7%
γ 1–8 每轮草稿长度
MTP 深度 1 多预测 1 个字
λ 0.3 MTP loss 权重
block_size 64 上下文窗口
RTX 4090 · CUDA · PyTorch 2.4.1 · tiny shakespeare 字符级 · 单种子
① 一次 forward 验收一串
第 2 章的因果掩码让第 k 个位置只看得见前 k 个字。
所以把"上下文 + γ 个草稿"一起喂给主模型,每一行输出都是"只看前面"时的结果,一次 forward 就拿到 γ+1 个位置的分布 p。
拖 γ、切模式,点格子里任意一行看它看得见什么:
3
灰色 8 个字取自验证集;橙色草稿字符是示意,不是模型输出。
最右一格读数来自 14_mtp.py:同一段 64 字符的验证集文本,只喂前 56 个字 vs 喂满 64 个字,前 56 个位置的 logits 逐个比较,最大差 0.0。多喂的 8 个字没有改动前面任何一个位置的输出。
↳ 代码:14_mtp.py 的 spec_draft_model 【步骤 2】main_model(crop(seq + drafts)) 取 ml[0, -(gamma + 1):]
草稿中途猜错了,后面那几行不就白算了?
是白算了。第一个被拒绝的位置之后,草稿全部作废,那几行的 p 也一起丢掉。
但这次 forward 本来就要跑:哪怕只接受 0 个草稿,主模型也会在拒绝位置重抽出 1 个字,不比普通生成少。
错位置之前的那些行,因为因果掩码看不见后面的错字,算出来的 p 仍然有效。
真实推理系统里,这一步怎么和 KV cache 配合?
上下文的 K/V 已经在 KV cache 里,验收时只需要把"最后一个已确认的字 + γ 个草稿"这一小段送进去,
算出它们的 γ+1 行输出,同时把它们的 K/V 追加进 cache。被拒绝的草稿对应的 K/V 截掉即可。
本关的教学脚本没有 KV cache,每次 forward 都把最近 64 个字整段重算,机制一样,只是更慢。
↳ 下一步:一次 forward 拿到了 γ+1 个 p,草稿员当时给的是它自己的分布 q。p 和 q 不一样时,收还是不收?
② 接受规则:min(1, p/q)
草稿员从 q 里抽出一个字 x。主模型给它的概率是 p(x)。规则(Leviathan et al. 2023,arXiv 2211.17192 §2.3):
以 min(1, p(x)/q(x)) 的概率接受;拒绝时从 max(0, p−q) 归一化后的分布里重抽一个。
下面是验证集里一个真实位置,上下文结尾 ,主模型与草稿模型对下一个字的前 8 名。点一行,假设草稿抽到了它:
p · 主模型
q · 草稿模型
投机采样 20 万次的频率
–
20 万次里的接受率(理论 Σmin(p,q))
亲手抽在页面里跑同一套规则
总变差距离:投机采样离 p 是 0.0005,草稿 q 自己离 p 是 0.053。经过验收的输出分布回到了 p。
↳ 代码:14_mtp.py 的 Counter.check(接受 / 残差重抽)· 20 万次模拟在脚本里"正确性(采样)"那一段
为什么这样验收之后,输出分布和 p 完全相同?
对任意一个字 x,它被输出有两条路:
① 草稿抽到 x 并被接受:q(x) · min(1, p(x)/q(x)) = min(p(x), q(x));
② 草稿被拒绝(总概率 1 − Σ min(p,q)),再从残差里抽到 x:残差是 max(0, p−q) 除以它的总和,
而这个总和恰好等于 Σ (p − min(p,q)) = 1 − Σ min(p,q),两者约掉,这条路的概率是 max(0, p(x) − q(x))。
两条路相加:min(p,q) + max(0, p−q) = p(x)。上面的计算卡对每个字都把这两项算了出来。
证明见 Leviathan 2023 附录 A.1;平均接受率 α = Σ min(p,q)(Corollary 3.6)。
主模型比草稿更确定(p > q)的字,为什么一定接受?
p(x) ≥ q(x) 时 p/q ≥ 1,接受概率取 1。草稿抽到 x 的频率 q(x) 还不够 p(x),全部收下也只是"不超额";
差的那部分 p(x) − q(x) 由拒绝后的残差分布补上。
反过来 q(x) > p(x) 的字(本例里的 e:q=0.0489,p=0.0062)草稿抽得太多,只按 p/q 的比例收下,多出来的被拒绝。
↳ 下一步:每个草稿有一定概率被接受。那一轮该让草稿员猜几个?猜多了,后面的更容易作废。
③ γ 猜几个
设每个草稿被接受的概率都是 α(相互独立)。一轮里主模型 forward 1 次,期望吐出
(1 − α^(γ+1)) / (1 − α) 个字(Leviathan 2023 式 (1))。
但草稿员每轮要跑 γ 次,每次耗时是主模型的 c 倍,墙钟加速的期望是上式再除以 γc + 1(Theorem 3.8)。
拖 α 和 c,橙线是公式,蓝点是本关小草稿模型的 6 次真实测量:
0.70
0.27
–
本关 c = 草稿 / 主模型单次 forward
c 决定 γ 能开多大。本关草稿模型单次 forward 0.225 ms、主模型 0.836 ms,c ≈ 0.27;Leviathan 论文实验里草稿比目标模型小两个数量级,c 始终 小于 0.05(Definition 3.7 下的说明)。
↳ 代码:14_mtp.py 主测量那段 for gamma in [1, 2, 3, 4, 6, 8],每行记下 tokens_per_main 与 expected_tokens_per_main
γ=8 每次吐的字比 γ=6 还少,公式不对吗?
公式里的 α 是一个固定值,真实测量里每一行的 α 各不相同:γ=6 那次实测接受率 0.717,γ=8 那次 0.676。
把各自的 α 代进公式,γ=6 期望 3.192 个、γ=8 期望 2.993 个,实测 3.093 和 2.956,方向一致。
每行只生成 600 个字、跑一次,α 在 0.67–0.75 之间浮动属于这个样本量下的波动。
公式还假设每个位置的接受相互独立,真实文本里难猜的字往往扎堆出现。
↳ 下一步:小草稿模型要单独训练、单独部署,每轮还要自回归跑 γ 次。能不能让主模型训练时顺手学会打草稿?
④ MTP:主模型自带草稿头
DeepSeek-V3 的 Multi-Token Prediction(arXiv 2412.19437 §2.2):主模型最后一层的隐状态 ht 已经编码了前文,
再拼上"下一个字"的嵌入,过一个小 Block,就能预测下下个字。嵌入表和输出头与主模型共用。
推理时它每次猜 1 个字,正好当草稿员。点图里的部件看它是什么、多少参数:
真实样本MTP 自草稿生成的 240 个字
主模型每次 forward 做两件事:验收上一轮 MTP 头给的草稿,再顺手采样一个字;MTP 头接着给下一轮的草稿。
接受时这次 forward 吐 2 个字,拒绝时吐 1 个。拖滑块逐次看:
被接受的草稿
主模型采样
拒绝后重抽
1
对比同样每轮猜 1 个:MTP 头 vs 小草稿模型
DeepSeek-V3 用顺序 MTP、深度 D=1,报告第二个 token 的接受率在 85%–90%,解码 TPS 为原来的 1.8 倍(§5.4.3)。MTP loss 权重 λ 在前 10T token 取 0.3,之后的 4.8T 取 0.1。
↳ 代码:14_mtp.py 的 GPT.mtp_logits · spec_mtp · train 里 loss = main + args.lam * mtp
MTP 头预测的是隔一个字,val loss 为什么和主模型差不多?
它的输入里已经有 xt+1 的嵌入,预测 xt+2 时和主模型预测下一个字面对的是同一类问题,只是多了主模型 4 层算好的 ht 可用。
结束时 MTP 头 val loss 1.5332,主模型 1.5276。小草稿模型只有主模型 7% 的参数、从原始字符算起,val loss 1.8862。
推理时 xt+1 是主模型刚采样出来的真字,所以 MTP 头的草稿建立在正确的前一个字上。
本关 MTP 接受率落在 DeepSeek-V3 报告的区间里,算复现吗?
不算。DeepSeek-V3 是 671B 总参数、每 token 激活 37B 的 MoE,在 14.8T token 上预训练,评测的是多领域生成;
本关主模型 0.74M 参数、字符级 tiny shakespeare、单种子。规模和任务差了几个数量级,两个数字接近只能说明机制在两端都能工作,不能互相印证。
Meta 2024 年的 MTP 和 DeepSeek-V3 的有什么不同?
Gloeckle et al.(arXiv 2404.19737,ICML 2024)在共享主干上接 n 个独立输出头,各自是一层 Transformer,并行预测后面 n 个 token;
4-token 预测的模型用自投机解码,推理最多快 3 倍。
DeepSeek-V3 改成顺序模块:第 k 个模块吃第 k−1 个的隐状态和第 i+k 个 token 的嵌入,
原文是 "sequentially predict additional tokens and keep the complete causal chain at each prediction depth"。本关代码只做深度 1:模块输入里除了 ht,还有 xt+1 的嵌入,这是它和独立并行头最直观的区别。
↳ 下一步:forward 次数确实少了。钟表上快了多少?把所有方法摆进一张账本。
⑤ 真实账本
每种方法都从 ROMEO:\n 开始,temperature 1.0 采样生成 600 个字,记下主模型和草稿的 forward 次数与墙钟秒数。
RTX 4090 · CUDA · PyTorch 2.4.1 · batch=1 · 每行只跑一次。点表头按钮排序,点一行看理论值:
forward 次数最多省到 1/3(γ=6 每次吐 3.09 个字),墙钟加速最高 1.27×(MTP),γ=6、γ=8 比普通生成还慢。
正确性输出没有被改动
forward 省了三分之二,墙钟为什么只快 1.0–1.3 倍,γ=6、8 还变慢?
① 草稿不够便宜:草稿模型单次 forward 0.225 ms,是主模型 0.836 ms 的约 27%。γ=8 时每轮 8 次草稿 forward 加起来比一次主 forward 还贵。
② γ 次草稿逐个跑:草稿员自己也是自回归,γ 个字要 γ 次 forward,每次都有 Python 调度、kernel 启动的固定开销。
这类开销不随模型变小而缩小:草稿模型参数只有主模型的 7%,单次 forward 却要主模型 27% 的时间。
③ 每个字都有采样开销:batch=1,脚本每抽一个字都把分布拷回 CPU 再 multinomial,草稿的字也要抽,验收时还有 .item() 同步。
普通生成 600 个字 0.53 s,折合每个字约 0.88 ms,比单次 forward 的 0.836 ms 多出来的主要是这类开销。
④ 每行只测一次:几百毫秒量级的计时,相邻两行差 0.02–0.03 s 可能就是抖动。
点表里任意一行,可以看到用该行 α 和 c ≈ 0.27 代进 Theorem 3.8 的理论墙钟加速,实测都低于它。式子只计 forward 的耗时,③ 不在式子里。
600 个字早就超出 64 的上下文窗口了,这几行的输出还和主模型一致吗?
一致的是分布,口径是"验收时的截断规则"。总长超过 block_size=64 以后,普通生成每步把窗口右移 1 格;
验收时却是把"已生成 + γ 个草稿"整体截成最后 64 个字,靠后的草稿位置看到的历史更短。
两边其实是窗口规则不同的两台生成器。所以这几行保证的是:输出分布等于"主模型在验收截断规则下"的分布,不是和普通滑窗生成逐字对应。
这也是贪心一致性检查只在窗口以内做的原因:8 个字的提示 + 生成 48 个字 + 最多 8 个草稿 = 64。
真实系统用 KV cache,上下文远长于生成长度,不会碰到这个问题。
贪心解码时,接受规则变成什么?
temperature=0 时脚本把 p、q 都换成 argmax 位置为 1 的 one-hot(dist 函数)。
草稿抽到的字 x 是 q 的 argmax:若它也是 p 的 argmax,p(x)/q(x) = 1,接受;否则 p(x) = 0,必拒绝,
残差 max(0, p−q) 只剩 p 的 argmax 那一格,重抽出来的就是主模型自己会选的字。
所以贪心下投机解码的输出应当和主模型贪心逐字相同,下面的检查就是验证这一点。
带 MTP loss 训练,主模型本身变好了吗?
本关看不出来。同样 5000 步,带 MTP loss(λ=0.3)的主模型 val loss 1.5276,不带的 1.5357,差 −0.008。
两次都是单种子,这个差距在种子噪声的量级内,不足以说明 MTP 让主模型变好。
Gloeckle 2024 报告的收益出现在 13B 规模的代码任务上,和本关的规模、数据都不同。
读原论文时,为什么 Chen 2023 和 Leviathan 2023 的 p、q 对不上?
两篇论文同期独立提出同一套规则,字母含义相反。
Leviathan et al.(arXiv 2211.17192,ICML 2023)里 p 是目标模型、q 是草稿,接受概率 min(1, p/q);
Chen et al.(DeepMind,arXiv 2302.01318)里 q 是目标模型、p 是草稿,写作 min(1, q/p),残差 (q − p)₊。
本页统一用 Leviathan 的记号:p = 主模型,q = 草稿。Chen 2023 在 Chinchilla 70B 上报告分布式环境下 2–2.5 倍解码加速。
Medusa、EAGLE 又是怎么打草稿的?
Medusa(Cai et al.,arXiv 2401.10774):在主模型上加几个额外的解码头,并行预测后面多个 token,再用树注意力一次验收多条候选。Medusa-1 加速超过 2.2 倍,Medusa-2 为 2.3–3.6 倍。
EAGLE(Li et al.,arXiv 2401.15077):草稿在特征层(倒数第二层)做自回归,原文认为这比在 token 层自回归更直接;LLaMA2-Chat 70B 上 2.7–3.5 倍。
DeepSeek-V3 论文提到自己保持因果链的做法和 EAGLE 相近,区别是 EAGLE 的目标是投机解码,MTP 首先用来改进训练。
↳ 跑法:python 14_mtp.py(主模型 + MTP 头 + 草稿模型 + 全部测量)· python 14_mtp.py --mtp 0 --no-spec(不带 MTP loss 的对照)
🎉 投机解码与 MTP · 通关
你亲手验过了接受规则为什么不改分布、γ 为什么不能无限加、MTP 头为什么接受率更高,也看清了玩具规模下省 forward 不等于省时间。