Skip to content

多 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 模块。每个模块 kk = 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_k

lambda 是一个较小的权重系数——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%(更多前向计算、更多损失项)。它的回报有两个:

  1. 更密集的训练信号。每个隐藏状态会看到 D+1 个监督目标。根据 DeepSeek-V3 的消融实验,在 MMLU、GSM8K、MATH、HumanEval 上都能带来稳定的几个百分点提升。

  2. 推理时免费的推测解码草稿器。MTP 模块本来就是为预测接下来几个 token 而训练的。把它拿来充当草稿网络后,接受率可超过 80%。在这个水平上,N=3 或 N=5 的推测解码就能带来 1.8× 的吞吐。那 10% 的训练期成本,在你第一次跑推理时就赚回来了。

与 EAGLE 的关系

EAGLE 会在预训练结束后,单独训练一个小型草稿模型。MTP 则是在预训练阶段就把草稿“烤”进模型里。这两条路线会达到相近的接受率,但走的是不同的流水线:

维度EAGLE-3MTP(DeepSeek-V3)
训练时机预训练之后预训练期间
是否向后兼容已有权重否(需要重新训练)
草稿参数1–2 层 Transformer1 个 Transformer block + 投影层
接受率0.88-0.92深度 1 时 0.80+
除加速外的收益只有推测解码更密集的训练信号 + 加速

动手实现

code/main.py 会端到端构建一个单独的 MTP 模块:共享嵌入、投影层、Transformer block、共享输出头。随后,它会在一个短小的合成序列上计算逐深度交叉熵损失,并按组件打印参数量。一个只有 32 个 token 的玩具词表,让数字更容易读。

第 1 步:共享嵌入表

一个单独的 vocab_size x hidden 表会同时被主模型以及所有深度上的 MTP 模块使用。不是复制一份,而是字面意义上的同一块张量。

第 2 步:逐深度组合

python
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 调度、内存开销,以及推理时推测解码的接线方式。

练习

  1. 运行 code/main.py。展示随着合成信号增强,逐深度损失会单调下降。再把合成数据改成固定模式,验证深度 1 和深度 2 的损失都会收敛。

  2. 对一个稠密 70B 模型(hidden 8192、80 层)且 D=1 的 MTP 模块,计算参数开销。将结果与 DeepSeek-V3 报告的 14B 开销比较。解释为什么 DeepSeek 的数字更高:MTP 的 Transformer block 继承了相同的 MoE 结构,从而抬高了每个模块的参数量。

  3. 在玩具实现中加入 D=2:再增加第二个 MTP 模块,让它接收 h^(1) 并预测 t_{i+2}。验证联合损失和参数核算与 DeepSeek 论文中的公式 19–21 一致。

  4. 把玩具切换成并行 MTP(Gloeckle 风格):在主隐藏状态顶部加入 D 个输出头,每个头预测一个不同偏移量。对同一组合成信号,比较顺序版与并行版各个深度的损失。对于 k > 1,顺序版本应当产生更低的深度-k 损失,因为它能以中间预测为条件。

  5. 把训练好的 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

延伸阅读