多 token 预测(Multi-Token Prediction,MTP)
从 GPT-2 到 Llama 3,所有自回归 LLM 在每个位置上都只训练一个损失:预测下一个 token。DeepSeek-V3 在每个位置又加了第二个损失:预测再后面的那个 token。额外的 14B 参数(相对于 671B 模型)通过梯度流被蒸馏回主模型,而训练好的 MTP 头又在推理时被改造成推测解码(speculative decoding)的草稿器,接受率超过 80%。1.8× 的生成吞吐就这样“白送”了。本课会按照 DeepSeek 技术报告构建顺序式 MTP 模块,计算损失和共享头的参数布局,并解释为什么 MTP 保留了因果链,而 Gloeckle 等人最早提出的并行式 MTP 却打断了它。
类型: Build 语言: Python(stdlib) 前置课程: 第 10 阶段·04(预训练一个 mini GPT),第 10 阶段·15(推测解码) 耗时: ~60 分钟
学习目标
- 陈述 MTP 的训练目标,并推导跨不同预测深度的联合损失。
- 解释 Gloeckle 等人的并行式 MTP 头(2024)与 DeepSeek-V3 的顺序式 MTP 模块之间的差异,以及为什么顺序设计能保留因果链。
- 计算在预训练过程中加入 MTP 模块带来的参数与内存开销。
- 从零实现一个 MTP 模块:共享嵌入、按深度划分的 Transformer block、投影层,以及共享输出头。
问题
下一个 token 预测是标准的 LLM 训练目标。每个隐藏状态只会被监督去预测一件事:紧接着的下一个 token。这个信号其实相当弱。一个序列中的大部分信息都不止延伸一个 token——比如结构、一致性、事实性、算术流程。模型只能在数万亿个 token 上,靠不断累加这种单 token 信号,慢慢学会这些能力。
MTP 会问:如果每个隐藏状态都被监督去同时预测多个未来 token,会怎样?Gloeckle 等人(Meta,2024)证明这会有帮助。他们的实现是在骨干网络顶部挂上多个彼此独立的输出头,每个头预测一个不同偏移量的位置。并行、简单,但这些头看到的是同一个隐藏状态,没有任何分层细化;而且它们的预测之间不存在因果链,所以不能拿来做推测解码。
DeepSeek-V3(2024 年 12 月)把 MTP 重设计成了顺序式模块,在每个预测深度都保留了因果链。模型先用 h_i^(0) 预测 t+1,再用一个新的隐藏状态 h_i^(1) 去预测 t+2;这个新隐藏状态会把 h_i^(0) 与 E(t+1) 的嵌入结合起来,依此类推。每个深度都对应一个自己的小 Transformer block。共享嵌入与共享输出头让额外参数保持在可控范围内。在 DeepSeek-V3 的规模下,MTP 模块只是在 671B 主模型参数之上增加了 14B。这个 2% 的开销,不仅换来了更密集的训练信号,还顺手带来一个可在推理时直接使用的推测解码草稿器。
本课会从零构建单个 MTP 模块,以及深度为 D 的损失。数学很整洁。实现大约 150 行。
概念
顺序式 MTP 配方
DeepSeek-V3 在主模型顶部增加了 D 个 MTP 模块。每个模块 k(k = 1..D)负责预测深度为 k 的 token,也就是在已知截至位置 i 的前缀条件下预测 t_{i+k}。
模块 k 包含:
- 一个拥有自己注意力和 MLP 的 Transformer block
T_k。 - 一个投影矩阵
M_k,用于把前一深度的隐藏状态和下一深度真实 token 的嵌入结合起来。 - 共享嵌入
E(与主模型相同)。 - 共享输出头
Out(与主模型相同)。
训练时,对于截至位置 i 的前缀,逐深度隐藏状态为:
h_i^(0) = main model backbone at position i
h_i^(k) = T_k( M_k * concat(RMSNorm(h_i^(k-1)), RMSNorm(E(t_{i+k}))) ) for k >= 1逐深度预测为:
logits_{i+k} = Out(h_i^(k-1)) for k = 1..D逐深度损失是针对真实值 t_{i+k} 的交叉熵:
L_k = CE(logits_{i+k}, t_{i+k})跨深度的联合损失为:
L_MTP = (lambda / D) * sum_{k=1..D} L_klambda 是一个较小的权重系数——DeepSeek-V3 在训练前 10% 使用 0.3,之后使用 0.1。总训练损失为 L_main + L_MTP。
为什么是顺序式,而不是并行式
Gloeckle 最初提出的并行 MTP 有 D 个输出头,每个头都直接作用于 h_i^(0)。每个头都从同一个骨干隐藏状态预测 t_{i+k}。这样训练没有问题,但这些预测彼此不以对方为条件。你无法利用 head_1 的输出去帮助 head_2——因为所有头都是并行触发的。
DeepSeek-V3 的顺序设计会从 h_i^(k-1) 和真实下一 token 嵌入 E(t_{i+k}) 构造 h_i^(k)。这保留了因果链:若要预测 t_{i+k+1},深度 k+1 的模块会看到位置 t_{i+k} 上发生了什么。这在结构上与自回归解码器消费自己输出的方式完全一致——因此 MTP 模块可以直接作为推测解码的草稿器使用。
在推理时:把 h_i^(k-1) 和草拟出的 t_{i+k} 一起喂给模块 k+1,即可得到对 t_{i+k+1} 的预测。不断重复。这正是 EAGLE 风格的草稿流程,只不过这里使用的是训练好的 MTP 模块作为草稿网络。DeepSeek-V3 报告称,第一个 MTP 模块的接受率超过 80%,整体可实现约 1.8× 的加速。
参数核算
对于一个隐藏维度为 h、词表大小为 V 的模型:
- 主模型:数十亿参数,再加一个大小为
V * h的输出头。 - 共享输出头:复用主模型的头。无需额外参数。
- 共享嵌入:复用主模型的嵌入。无需额外参数。
- 每个 MTP 模块:
- 投影
M_k:(2h) * h = 2h^2。 - Transformer block
T_k:注意力(MHA 约为4h^2)加上 MLP(SwiGLU 在比例为 8/3 时通常约为8h^2)。每个 block 大约12h^2。
- 投影
每个模块的总额外参数:约 14h^2。对于 DeepSeek-V3 的 h = 7168,且 D = 1:纸面参数量约为 ~14 * 7168^2 = ~720M。而 DeepSeek-V3 报告的是 14B——差异主要来自 MTP 模块里的 expert 层同样也是 MoE。
推测解码的回报
在预训练阶段,MTP 模块会让训练变慢大约 10%(更多前向计算、更多损失项)。它的回报有两个:
更密集的训练信号。每个隐藏状态会看到 D+1 个监督目标。根据 DeepSeek-V3 的消融实验,在 MMLU、GSM8K、MATH、HumanEval 上都能带来稳定的几个百分点提升。
推理时免费的推测解码草稿器。MTP 模块本来就是为预测接下来几个 token 而训练的。把它拿来充当草稿网络后,接受率可超过 80%。在这个水平上,N=3 或 N=5 的推测解码就能带来 1.8× 的吞吐。那 10% 的训练期成本,在你第一次跑推理时就赚回来了。
与 EAGLE 的关系
EAGLE 会在预训练结束后,单独训练一个小型草稿模型。MTP 则是在预训练阶段就把草稿“烤”进模型里。这两条路线会达到相近的接受率,但走的是不同的流水线:
| 维度 | EAGLE-3 | MTP(DeepSeek-V3) |
|---|---|---|
| 训练时机 | 预训练之后 | 预训练期间 |
| 是否向后兼容已有权重 | 是 | 否(需要重新训练) |
| 草稿参数 | 1–2 层 Transformer | 1 个 Transformer block + 投影层 |
| 接受率 | 0.88-0.92 | 深度 1 时 0.80+ |
| 除加速外的收益 | 只有推测解码 | 更密集的训练信号 + 加速 |
动手实现
code/main.py 会端到端构建一个单独的 MTP 模块:共享嵌入、投影层、Transformer block、共享输出头。随后,它会在一个短小的合成序列上计算逐深度交叉熵损失,并按组件打印参数量。一个只有 32 个 token 的玩具词表,让数字更容易读。
第 1 步:共享嵌入表
一个单独的 vocab_size x hidden 表会同时被主模型以及所有深度上的 MTP 模块使用。不是复制一份,而是字面意义上的同一块张量。
第 2 步:逐深度组合
def combine(prev_hidden, next_token_embed, M_k):
# concat along feature dim, then project down to hidden
concat = rms_norm(prev_hidden) + rms_norm(next_token_embed) # vector addition stand-in
projected = matvec(M_k, concat)
return projected真实的 DeepSeek-V3 会把两个经 RMSNorm 的向量拼接成 [2h],再用一个 h x 2h 矩阵做投影。这个玩具版本为了保持 stdlib 的简洁,只用向量加法来代替。
第 3 步:深度 k 上的 Transformer block
由自注意力加 MLP 组成。在这个玩具里,一个单层线性 attention block 和一个 SwiGLU MLP 足以在没有 numpy 的情况下把结构展示清楚。
第 4 步:共享输出头
复用主模型的输出投影。输出对整个词表的 logits。
第 5 步:逐深度损失
对偏移量为 k 的真实 token,计算 softmax(logits) 的交叉熵。再用 lambda / D 的缩放系数,在多个深度上做聚合。
第 6 步:参数核算
打印总参数量、共享部分(嵌入、头)的参数量,以及每个模块的额外参数量。展示 MTP 额外开销占主模型大小的比例。
使用
MTP 已集成在 DeepSeek-V3(2024 年 12 月)与 DeepSeek-R1 系列中。推理时:
- DeepSeek 自家的服务栈可以开箱即用地把 MTP 模块作为推测解码器使用。
- 截至 2026 年 4 月,vLLM 与 SGLang 都已有对 DeepSeek-V3 MTP 的集成路径。
- AMD 的 ROCm SGLang 教程展示了一个针对 V3 checkpoint 的 MTP 推测解码配置,并测得 1.8× 加速。
在新的预训练运行中,适合使用 MTP 的情形:
- 你掌控完整的预训练流水线,并希望提前积累更密集的训练信号。
- 你知道模型未来会被大规模部署,并希望免费获得推测解码。
- 你的隐藏维度至少有 4096。在 1B 级规模上,额外开销往往大于收益。
不适合的情形:
- 你在微调一个现成的稠密预训练模型。MTP 模块并没有被训练过。
- 你在做研究模型,希望有一个足够干净的基线可供比较。MTP 会改变架构本身。
交付
本课会产出 outputs/skill-mtp-planner.md。给定一份预训练运行规格(模型大小、数据、算力),它会返回一个集成 MTP 的方案:深度数 D、lambda 调度、内存开销,以及推理时推测解码的接线方式。
练习
运行
code/main.py。展示随着合成信号增强,逐深度损失会单调下降。再把合成数据改成固定模式,验证深度 1 和深度 2 的损失都会收敛。对一个稠密 70B 模型(hidden 8192、80 层)且 D=1 的 MTP 模块,计算参数开销。将结果与 DeepSeek-V3 报告的 14B 开销比较。解释为什么 DeepSeek 的数字更高:MTP 的 Transformer block 继承了相同的 MoE 结构,从而抬高了每个模块的参数量。
在玩具实现中加入 D=2:再增加第二个 MTP 模块,让它接收 h^(1) 并预测
t_{i+2}。验证联合损失和参数核算与 DeepSeek 论文中的公式 19–21 一致。把玩具切换成并行 MTP(Gloeckle 风格):在主隐藏状态顶部加入 D 个输出头,每个头预测一个不同偏移量。对同一组合成信号,比较顺序版与并行版各个深度的损失。对于
k > 1,顺序版本应当产生更低的深度-k 损失,因为它能以中间预测为条件。把训练好的 MTP 模块当作 EAGLE 风格草稿器使用:在推理时调用模块
k提出t_{i+k}。把这些草稿 token 与主模型在保留序列上的预测进行比较,测量接受率。如果你在玩具实验中做到 50%+,就复现了“用 MTP 当草稿器”的经验性质。
关键术语
| 术语 | 人们常说什么 | 它真正的含义 |
|---|---|---|
| MTP 模块 | “额外损失块” | 一个小型 Transformer block 加上投影层,用来预测主模型之后第 k 个位置上的 token |
| 预测深度 | “哪个偏移量” | 整数 k,表示模块 k 会在已知截至位置 i 的前缀时预测 t_{i+k} |
| 并行 MTP | “Gloeckle 风格” | 在同一个骨干隐藏状态上放置 D 个独立头,不形成条件链 |
| 顺序 MTP | “DeepSeek-V3 风格” | 每个模块都以更浅一层的隐藏状态和下一 token 的嵌入为条件;保留因果链 |
| 共享输出头 | “复用主头” | MTP 模块调用的是主模型的 LM head,而不是单独的输出投影 |
| 共享嵌入 | “复用主表” | 同一张词表嵌入表在各处共用;没有重复参数 |
| 投影矩阵 M_k | “组合 hidden + next-token” | 一个 h x 2h 的线性层,把前一深度的隐藏状态和目标 token 嵌入折叠成下一深度输入 |
| 联合损失 L_MTP | “平均后的额外损失” | 各深度交叉熵损失的算术平均,再乘以 lambda |
| 深度 1 接受率 | “MTP 草稿猜对的频率” | D=1 的 MTP 模块 top-1 预测与主模型 top-1 预测相同的频率;在 DeepSeek-V3 上超过 80% |
| Lambda 权重 | “额外损失的重要性” | 每个深度的缩放系数;DeepSeek-V3 在训练早期用 0.3,后期用 0.1 |
延伸阅读
- DeepSeek-AI — DeepSeek-V3 Technical Report (arXiv:2412.19437) —— 完整的顺序式 MTP 描述(第 2.2 节),包括联合损失公式与推理时的 1.8× 加速
- Gloeckle et al. — Better & Faster Large Language Models via Multi-token Prediction (arXiv:2404.19737) —— DeepSeek 设计所改进的并行 MTP 基线
- DeepSeek-V3 model card on Hugging Face —— 总计 685B(671B 主体 + 14B MTP),以及部署说明
- Leviathan et al. — Fast Inference from Transformers via Speculative Decoding (arXiv:2211.17192) —— MTP 所嵌入的推测解码框架
- Li et al. — EAGLE-3 (arXiv:2503.01840) —— 2025 年 EAGLE 的草稿架构,也是 MTP 所竞争的对照路线