一个训练场景:为什么单 Token 预测不够
训练一个 671B 参数的 MoE 语言模型时,每个训练样本都是一段文本,模型需要根据前文预测下一个 token。传统做法是只预测紧邻的下一个 token,损失函数只计算这一个位置的交叉熵。这意味着模型在每一个位置只获得一个监督信号,而文本中后续若干个 token 的语义约束完全被浪费。
以训练数据中的一句代码 def fibonacci(n): 为例,模型看到 def 后,不仅要预测 fibonacci,还应当为后续的 (n): 做好准备。单 Token 预测的训练目标只要求模型输出 fibonacci,至于它是否在内部表征中为后续 token 做了铺垫,并不直接受到监督。结果是模型需要更多样本才能学会这种长距离的依赖关系,训练效率受限。
DeepSeek-V3 在预训练阶段引入了多 Token 预测(Multi-Token Prediction,MTP)目标,让模型在每个位置同时预测多个未来 token。论文报告,这一目标显著提升了模型在评估基准上的综合表现,并且该模块可以在推理时被复用为投机解码的草稿模型,从而加速生成。本文以 DeepSeek-V3 的训练与部署为贯穿场景,拆解 MTP 的训练目标、模块设计、样本效率收益、推理加速机制,以及实现复杂度和显存开销的权衡。
MTP 的核心思想:让模型为未来做规划
多 Token 预测并不是一个新概念。早在 2024 年,Gloeckle 等人就在论文《Better & Faster Large Language Models via Multi-token Prediction》中提出,让模型使用多个独立的输出头并行预测接下来的多个 token,可以提升样本效率。其直觉是:如果模型在预测第 t+1 个 token 的同时,还要预测第 t+2、第 t+3 个 token,那么它就必须在内部表征中提前规划更长范围的语义,这种规划压力会迫使模型学习更丰富的表示。
DeepSeek-V3 采用了类似的目标,但实现方式与 Gloeckle 等人的并行多头方案不同。论文中的 MTP 模块采用串行结构:预测第 t+2 个 token 时,会把第 t+1 个 token 的预测结果也作为输入,从而保持完整的因果链。这种设计避免了并行预测中“各预测头互不知道对方预测了什么”的问题,让每个位置的预测都能利用前面位置的预测信息,理论上能产生更一致的预测序列。
MTP 的另一个动机是让训练信号更密集。在单 Token 预测中,每个位置只有一个监督信号;而 MTP 将预测范围扩展到 D 个未来 token,每个位置就有 D 个监督信号。这相当于在不增加数据量的情况下,让模型从每个样本中学习更多信息,从而可能提高数据效率。论文还提到,MTP 可能使模型能够预先规划其表示,以更好地预测未来 token,这可以理解为一种隐式的“规划”能力。
MTP 模块设计:串行预测头与共享参数
DeepSeek-V3 的 MTP 模块在主干 Transformer 之上堆叠了 D 个额外的预测头(论文中 D 取 2)。每个预测头是一个轻量级的 Transformer 块,它接收主模型在当前位置的隐状态,以及前一个预测头输出的嵌入,然后预测下一个 token。
具体来说,对于第 k 个预测头(k 从 1 到 D),它的输入是主模型在位置 i 的隐状态 h_i,以及第 k-1 个头在位置 i+1 的预测 token 的嵌入。这个嵌入通过一个线性投影与 h_i 融合,然后送入一个 Transformer 块,最后通过输出头得到第 i+k 个 token 的概率分布。
关键的设计决策是:所有预测头共享主模型的嵌入层和输出头(即 LM Head)。这意味着 MTP 模块并不引入全新的词表映射,而是复用主模型已经学好的 token 表示和输出分类器。这样做的好处是:
- 参数开销小:每个预测头只增加一个 Transformer 块和若干线性层,而不需要额外的嵌入矩阵和输出矩阵。
- 梯度共享:嵌入层和输出头同时接收主模型和 MTP 模块的梯度,这有助于这些层学习更通用的表示。
下图展示了 MTP 模块的数据流:
flowchart LR
A[输入 token 序列] --> B[主模型 Transformer]
B --> C[主模型隐状态 h_i]
C --> D[MTP 模块 1]
D --> E[预测 token t_{i+1}]
E --> F[嵌入层]
F --> G[MTP 模块 2]
C --> H[融合层]
H --> G
G --> I[预测 token t_{i+2}]
B --> J[主模型输出头]
J --> K[主模型预测 t_{i+1}]
图中,主模型在位置 i 的隐状态 h_i 被送入第一个 MTP 模块,该模块预测 t_{i+1}。然后,预测出的 t_{i+1} 的嵌入与 h_i 融合,送入第二个 MTP 模块,预测 t_{i+2}。注意,主模型本身也会预测 t_{i+1},但 MTP 模块的预测是独立的,它们共享输出头但使用不同的隐状态。
这种串行结构的一个直接后果是:MTP 模块的推理不能完全并行,因为预测 t_{i+2} 依赖于 t_{i+1} 的预测结果。但在训练时,由于我们有真实的 token 序列,可以并行计算所有位置的损失,只是每个 MTP 头的输入需要前一个头的输出,这可以通过一次前向传播完成,因为训练时我们使用的是真实 token 的嵌入,而不是预测 token。
训练目标:多位置交叉熵的加权和
MTP 的训练目标是在原有单 Token 预测损失的基础上,增加额外预测头的交叉熵损失。对于每个位置 i,总损失为:
[ \mathcal{L}{MTP} = \mathcal{L}{main}(i) + \sum_{k=1}^{D} \lambda_k \cdot \mathcal{L}_{k}(i) ]
其中 \mathcal{L}_{main}(i) 是主模型预测 t_{i+1} 的交叉熵损失,\mathcal{L}_{k}(i) 是第 k 个 MTP 头预测 t_{i+k} 的交叉熵损失,\lambda_k 是权重系数。在 DeepSeek-V3 中,D=2,且权重通常设置为 1,即所有预测头同等重要。
这里的关键是,MTP 损失只作用于主模型的隐状态,而不是独立训练一个旁路模型。也就是说,MTP 模块的梯度会通过主模型反向传播,从而影响主模型的所有参数。这不同于投机解码中常见的做法——训练一个独立的草稿模型,而是让主模型本身学会为多个未来 token 建模。
一个值得注意的细节是,MTP 模块在训练时使用的是真实 token 的嵌入作为输入,而不是预测 token。这避免了训练与推理之间的分布偏移,因为推理时 MTP 模块会使用自己的预测作为输入,如果训练时使用真实 token,模型可能无法处理推理时预测错误带来的误差累积。不过,DeepSeek-V3 的技术报告并未详细说明是否采用了类似“课程学习”或“计划采样”的策略来缓解这种偏移,但可以推断,由于 MTP 模块在推理时主要用于投机解码,其预测质量并不要求达到主模型的水平,因此这种偏移的影响可能被接受。
样本效率:为什么多预测能提升数据利用率
MTP 提升样本效率的机制可以从两个角度理解。
第一,每个训练样本提供了更多的监督信号。在单 Token 预测中,一个长度为 T 的序列只产生 T-1 个预测任务;而 MTP 将预测任务扩展到 D 个,每个位置产生 D 个预测,因此总监督信号数量约为原来的 D 倍。这意味着模型可以从同样的数据中学习更多,从而可能减少达到相同性能所需的样本量。
第二,MTP 迫使模型学习更长范围的依赖。以代码生成为例,预测 def fibonacci(n): 中的 fibonacci 可能只需要局部上下文,但要预测后面的 (n):,模型必须理解函数定义的语法结构。MTP 让模型在预测 fibonacci 的同时,也要为 (n): 做好准备,这迫使模型在内部表征中编码更全局的信息。这种“规划”压力可能使模型学习到更通用的表示,从而提升泛化能力。
DeepSeek-V3 的技术报告指出,MTP 目标经实证可显著提升模型在评估基准上的综合表现。虽然没有给出具体的样本效率提升数字,但可以推断,在相同的训练 token 数量下,MTP 模型比单 Token 预测模型表现更好,或者在达到相同性能时需要的训练数据更少。
然而,MTP 并非没有代价。额外的预测头增加了模型的计算量,训练时每个 token 的前向传播需要多计算 D 个 Transformer 块。在 DeepSeek-V3 中,由于 D=2,训练计算量大约增加了一小部分,但相对于整个 671B 模型的计算量,这个增加是可接受的。此外,MTP 模块的梯度需要反向传播到主模型,这增加了反向传播的计算量,但同样在可接受范围内。
推理加速:MTP 模块如何用于投机解码
投机解码(Speculative Sampling)是一种加速自回归生成的技术。其基本思想是:用一个较小的草稿模型快速生成多个候选 token,然后用目标大模型并行验证这些 token。如果草稿模型的预测与目标模型一致,这些 token 就可以被接受,从而减少目标模型的串行调用次数,提升生成速度。
DeepSeek-V3 的 MTP 模块天然适合作为投机解码的草稿模型,因为它与主模型共享嵌入层和输出头,且每个 MTP 头都是一个轻量级的 Transformer 块。在推理时,可以先用主模型计算当前位置的隐状态,然后依次运行 MTP 模块生成多个候选 token,最后用主模型一次前向传播验证这些 token。
具体流程如下:
- 主模型在位置
i生成隐状态h_i,并预测t_{i+1}。 - MTP 模块 1 使用
h_i预测t_{i+1}(与主模型预测可能不同)。 - 将预测的
t_{i+1}的嵌入与h_i融合,MTP 模块 2 预测t_{i+2}。 - 将
t_{i+1}和t_{i+2}作为草稿序列,连同上下文一起送入主模型,一次前向传播得到这些位置的预测分布。 - 根据接受规则(如
min(1, p_target/p_draft))逐个接受或拒绝草稿 token,被拒绝的位置重新采样。
这种方式的优势在于,MTP 模块的参数量远小于主模型(每个头只有一个 Transformer 块),因此生成草稿的速度很快。同时,由于 MTP 模块与主模型共享嵌入和输出头,其预测分布与主模型有一定的相关性,这有助于提高接受率。
DeepSeek-V3 的技术报告提到,MTP 模块可用于推测解码以实现推理加速,但未给出具体的加速比。不过,根据资料,MTP 的 token 接受率稳定在 85% 以上,训练时推理速度提升 1.8 倍。这里的“训练时推理”可能指的是在训练过程中使用 MTP 进行验证,但可以推断,在推理阶段使用 MTP 作为草稿模型也能获得可观的加速。
需要注意的是,投机解码的加速效果取决于草稿模型的接受率。如果 MTP 模块的预测与主模型差异较大,接受率会降低,加速效果减弱。因此,MTP 模块的质量是关键。由于 MTP 与主模型共享大部分参数,其预测与主模型不会偏差太大,这保证了较高的接受率。
与替代方案的比较:并行多头、独立草稿模型与 MTP
MTP 并非唯一的提升样本效率或推理速度的方案。下表对比了三种常见方案:
| 方案 | 训练目标 | 推理加速 | 参数开销 | 实现复杂度 | 适用场景 |
|---|---|---|---|---|---|
| 单 Token 预测 | 只预测下一个 token | 无 | 无 | 低 | 标准训练 |
| 并行多头预测(如 Medusa) | 多个独立头并行预测多个 token | 可作为草稿模型,但接受率较低 | 每个头需要额外的输出层,但可共享嵌入 | 中 | 训练效率提升有限,推理加速依赖草稿质量 |
| 串行 MTP(DeepSeek-V3) | 串行预测多个 token,保持因果链 | 可作为草稿模型,接受率高 | 每个头增加一个 Transformer 块,共享嵌入和输出头 | 中高 | 训练效率提升明显,推理加速效果好 |
| 独立草稿模型(如 EAGLE) | 训练一个小模型模仿大模型特征 | 加速效果好,但需要额外训练 | 独立模型,参数较多 | 高 | 推理加速,但训练成本增加 |
从表中可以看出,MTP 在训练效率和推理加速之间取得了较好的平衡。并行多头预测(如 Medusa)虽然实现简单,但由于各头独立预测,缺乏因果链,草稿质量较低,接受率不高。独立草稿模型(如 EAGLE)虽然加速效果好,但需要额外训练一个模型,增加了训练成本。MTP 则通过串行结构和参数共享,既提升了训练效率,又提供了高质量的草稿模型,且无需额外训练。
然而,MTP 的代价是增加了训练时的计算量,以及推理时的额外前向传播(MTP 模块需要运行 D 次)。在 DeepSeek-V3 中,D=2,这个开销相对可控。如果 D 更大,训练和推理的开销都会增加,但样本效率和加速效果可能进一步提升,这需要在实际中权衡。
实现细节与工程权衡
在工程实现中,MTP 模块的引入需要考虑以下几点:
- 显存开销:每个 MTP 头增加一个 Transformer 块,其参数和激活值会占用额外显存。在 DeepSeek-V3 的训练中,由于模型本身有 671B 参数,MTP 模块的显存开销相对较小,但仍需优化。论文提到,训练时使用了重采样 RMSNorm 和 MLA 上采样等技术来减少激活内存,MTP 模块也受益于这些优化。
- 计算开销:训练时,每个 token 需要额外计算
D个 Transformer 块的前向和反向。在 DeepSeek-V3 中,由于使用了 FP8 混合精度训练和 DualPipe 流水线并行,计算开销被部分隐藏。 - 推理部署:在推理时,MTP 模块可以作为草稿模型,但需要额外的推理步骤。如果 MTP 模块的预测质量不高,可能导致接受率低,反而增加延迟。因此,需要监控接受率,并根据实际情况调整是否启用 MTP 投机解码。
- 与 MoE 的交互:DeepSeek-V3 是 MoE 模型,每个 token 只激活部分专家。在投机解码的验证阶段,需要同时验证多个 token,这可能导致激活更多的专家,增加显存和计算开销。论文提到,MoE 模型与投机解码配合不佳,因为验证多个 token 时可能激活更多专家,削弱 MoE 的优势。因此,MTP 投机解码在 MoE 模型上的加速效果可能不如稠密模型。
失败模式与未解决问题
MTP 并非没有缺陷。以下是一些可能失败的模式:
- 分布偏移:训练时 MTP 模块使用真实 token 作为输入,但推理时使用预测 token,这可能导致分布偏移。如果 MTP 模块的预测错误,后续预测会进一步偏离,导致草稿质量下降。
- 长尾输入:对于罕见或复杂的输入,MTP 模块可能无法准确预测多个 token,导致接受率降低,投机解码失效。
- 评估基准偏差:MTP 提升的样本效率可能在不同任务上表现不一致。论文报告在代码和数学基准上提升明显,但在其他任务上可能有限。
未解决的问题包括:如何动态调整 D 的值以平衡效率与开销?如何进一步减少 MTP 模块的显存占用?在 MoE 模型中,如何优化投机解码以避免激活过多专家?这些问题需要进一步研究。
结论
DeepSeek-V3 的多 Token 预测通过串行预测头让模型在每个位置同时学习多个未来 token,显著提升了训练样本效率,并可作为投机解码的草稿模型加速推理。其核心在于共享嵌入和输出头,以最小化参数开销,同时通过保持因果链提高预测质量。然而,MTP 增加了训练和推理的计算量,且在 MoE 模型中的加速效果可能受限。对于训练大规模语言模型的团队,MTP 是一个值得考虑的优化方向,但需要根据具体场景权衡其收益与成本。