多 Token 预测(MTP) 本节摘要:GPT-2 到 Llama 3 的每个自回归 LLM 每位置只训一个损失:预测下一 token。DeepSeek-V3 加了第二个损失:预测再下一个。额外 14B 参数(671B 模型上)通过梯度流蒸馏回主模型,训好的 MTP 头在推理时改作投机解码草稿,接受率 80%+,1.8 倍生成吞吐白送。本节按 DeepSeek 技术报告从零搭顺序 MTP 模块,算损失与共享头参数布局,解释为什么 MTP 保因果链而 Gloeckle 等的原始并行 MTP 打破了它。 学习目标 阅读完本节,你应当能够: 陈述 MTP 训练目标,推导跨预测深度的联合损失。
本节摘要:GPT-2 到 Llama 3 的每个自回归 LLM 每位置只训一个损失:预测下一 token。DeepSeek-V3 加了第二个损失:预测再下一个。额外 14B 参数(671B 模型上)通过梯度流蒸馏回主模型,训好的 MTP 头在推理时改作投机解码草稿,接受率 80%+,1.8 倍生成吞吐白送。本节按 DeepSeek 技术报告从零搭顺序 MTP 模块,算损失与共享头参数布局,解释为什么 MTP 保因果链而 Gloeckle 等的原始并行 MTP 打破了它。
阅读完本节,你应当能够:
下一 token 预测是标准 LLM 训练目标,每个隐藏态被监督预测恰好一个东西:紧随其后的 token。这信号出奇地弱——序列里多数信息延伸超一个 token:结构、连贯性、事实性、算术流。模型得通过万亿 token 上累积许多单 token 信号来学这些。
MTP 问:要是每个隐藏态被监督同时预测多个未来 token 呢?Gloeckle 等(Meta,2024)证明有帮助——他们在骨干上放几个独立输出头,各预测不同偏移。并行、简单,但头看同一隐藏态无层级细化,预测不因果链,无法用于投机解码。
DeepSeek-V3(2024 年 12 月)把 MTP 重设计为保每个预测深度因果链的顺序模块。模型从 h_i^(0) 预测 t+1,再从组合了 h_i^(0) 与 E(t+1) 嵌入的新隐藏态 h_i^(1) 预测 t+2,依此类推。每深度是自己的小 Transformer 块,共享嵌入与共享输出头保参数开销适中。DeepSeek-V3 规模下,671B 主权重上加 14B MTP 模块——2% 开销换来更密训练信号 加 一个现成的推理时投机解码草稿。
DeepSeek-V3 在主模型上加 D 个 MTP 模块。每模块 k(k=1..D)预测深度 k 的 token——即给定到位置 i 的前缀,预测 t_{i+k}。模块 k 含:自己的 Transformer 块 T_k(含注意力与 MLP)、投影矩阵 M_k(把上一深度隐藏态与下一深度真值 token 嵌入组合)、共享嵌入 E(同主模型)、共享输出头 Out(同主模型)。
训练时,给定到位置 i 的前缀,逐深度隐藏态:
h_i^(0) = 主模型骨干在位置 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))(k=1..D),逐深度损失是对真值 t_{i+k} 的交叉熵,总损失是各深度加权和。
Gloeckle 等的并行头都看 h_i^(0),各预测不同偏移——但 h_i^(0) 没含它预测的中间 token 信息,预测间无因果依赖,不能链式用于投机解码。DeepSeek 的顺序模块:深度 k 的隐藏态 h_i^(k) 显式含了 E(t_{i+1})...E(t_{i+k})(经 M_k 投影喂入),所以预测深度 k+1 时模型已「见过」深度 1..k 的真值——因果链保持。这让训好的 MTP 模块推理时可作投机解码草稿:生成 t+1 后,把它的嵌入喂入深度 1 模块产 t+2 预测,接受率因链式条件而高(80%+)。
每 MTP 模块约一个 Transformer 块(T_k:注意力+MLP)+ 投影 M_k。嵌入 E 与输出头 Out 共享故零额外。DeepSeek-V3:D 个模块、每模块约 2~3B、总 14B(671B 上约 2%)。训练时这些参数提供更密监督,梯度流回主模型骨干使其更强;推理时 MTP 模块可丢弃(只用主模型)或保留作投机草稿(1.8 倍吞吐)。
class MTPModule: def __init__(self, embed_dim, num_heads, ff_dim, shared_embed, shared_head): self.T = TransformerBlock(embed_dim, num_heads, ff_dim) # 自己的块 self.M = np.random.randn(embed_dim, 2*embed_dim) * 0.02 # 投影 self.embed = shared_embed # 共享嵌入 self.head = shared_head # 共享输出头 def forward(self, h_prev, next_token_id): # 组合上一深度隐藏态与下一深度真值嵌入 e = self.embed.token_embed[next_token_id] combined = self.M @ np.concatenate([rmsnorm(h_prev), rmsnorm(e)]) h = self.T.forward(combined) # 经自己的块 return h, self.head @ h # 新隐藏态,预测 logits
def mtp_loss(main_model, mtp_modules, tokens, D=2): h0 = main_model.backbone(tokens) # 主骨干隐藏态 total_loss = ce_loss(main_model.head @ h0, tokens[1:]) # 深度 0:标准下一 token h_prev = h0 for k in range(1, D+1): losses_k = [] for i in range(len(tokens)-k-1): h_prev_i, logits = mtp_modules[k-1].forward(h_prev[i], tokens[i+k]) # 喂真值 token losses_k.append(ce(logits, tokens[i+k+1])) total_loss += (1.0/D) * np.mean(losses_k) h_prev = ... # 更新为深度 k 隐藏态 return total_loss
训练后,MTP 模块可作投机解码草稿:主模型产 t+1,把 E(t+1) 喂入 MTP 模块 1 产 t+2 预测,再喂入模块 2 产 t+3……链式条件使接受率高。这是 DeepSeek-V3 报告 1.8 倍生成吞吐的来源——MTP 模块白送了草稿。
def mtp_overhead(main_params_b, D, per_module_b): return main_params_b + D * per_module_b # DeepSeek-V3: 671B + 2*7B ≈ 685B, MTP 开销约 2%
DeepSeek-V3 是首批大规模部署 MTP 的,其开源权重含训好的 MTP 头,推理引擎(vLLM、SGLang)可加载作投机草稿。对比独立草稿模型(第 15 节):MTP 头共享主模型嵌入与头,额外内存约 2%,远低于独立小模型;且因在主模型表示空间上训,接受率更高。对比 Gloeckle 并行 MTP:并行版只能加训练信号,推理时不能链式用作草稿;DeepSeek 顺序版两者皆可。
本节产出 outputs/prompt-mtp-integrator.md——一个提示,接收模型 config 与目标(更强训练信号?推理加速?),推荐 MTP 深度 D、是否保留推理时模块、参数预算,给出训练与投机解码集成方案。
(Easy) 算不同 D(1、2、3)下 7B 模型的 MTP 参数开销,量化训练信号密度提升与内存代价的权衡。
(Medium) 实现并行 MTP(Gloeckle 版):骨干上 3 个独立输出头各预测 t+1、t+2、t+3,对比顺序版的推理时能否作投机草稿。
(Medium) 在小语料上训有/无 MTP(D=2)的小模型,对比同步数下的损失下降速度,验证 MTP 的更密训练信号。
(Hard) 实现推理时 MTP 投机解码:主模型产 t+1,MTP 模块链式产 t+2、t+3 草稿,测接受率与加速比,验证 DeepSeek 的 80%+ 接受率与 1.8 倍吞吐。
(Hard) 实验梯度流:训练时冻结主模型骨干只训 MTP 模块,对比 MTP 模块梯度流回骨干后主模型的提升,验证「MTP 参数蒸馏回主模型」的机制。
下一节,DualPipe 并行:DeepSeek-V3 的双向流水线,把前向反向计算与 MoE all-to-all 通信重叠,气泡近零。