公司动态
MPT (MTP, Multi-Token Prediction) 原理详解
MPT (MTP, Multi-Token Prediction) 原理详解注在 vLLM 与 DeepSeek 体系中该方案正确缩写为MTP (Multi-Token Prediction多 Token 预测)“MPT” 是常见笔误。下文统一使用MTP。一、论文背景MTP 作为一种推测解码Speculative Decoding方案最早由DeepSeek-V3 Technical Report (arXiv:2412.19437)系统性提出并工业化落地。核心思想传统自回归解码每个 forward pass 只产生 1 个 token造成 GPU 算力浪费。MTP 在主模型Target Model之外挂载一个或多个轻量级MTP 模块与主模型共享 embedding 和 lm_head串行预测未来 D 个 token 作为草稿draft再由主模型一次性并行验证verify。被接受的草稿 token 直接产出相当于一次 forward 输出多个 token从而无损加速推理。DeepSeek-V3 的 671B 总参数中包含约14B 的 MTP 模块权重证明了其工程可行性。后续 MiMo、Qwen3、GLM4、ERNIE、Gemma4 等模型家族都引入了类似机制见 vLLM 中*_mtp.py实现。二、MTP 模块架构结合 vLLM 源码参考 [deepseek_mtp.py](file:///workspace/vllm/model_executor/models/deepseek_mtp.py)DeepSeek-MTP 的核心结构如下1. 单个 MTP 层DeepSeekMultiTokenPredictorLayerclassDeepSeekMultiTokenPredictorLayer(nn.Module):def__init__(self,vllm_config,prefix):self.enormRMSNorm(...)# 对输入 embedding 做归一化self.hnormRMSNorm(...)# 对主模型 hidden state 做归一化self.eh_projnn.Linear(hidden*2,hidden,biasFalse)# 融合投影self.shared_headSharedHead(...)# 与主模型共享的输出头self.mtp_blockDeepseekV2DecoderLayer(...)# 一个 Transformer 解码块每个 MTP 模块包含共享 embedding 层embed_tokens与主模型共享共享输出头shared_head与主模型 lm_head 共享一个 Transformer 解码块mtp_block含 MLA DeepSeekMoE投影矩阵eh_proj将[input_embed, prev_hidden]拼接后投影回 hidden_size2. forward 流程[deepseek_mtp.py#L99-L121](file:///workspace/vllm/model_executor/models/deepseek_mtp.py#L99-L121)defforward(self,input_ids,positions,previous_hidden_states,inputs_embeds,spec_step_idx):# 1. mask 掉 position 0 的输入首 token 不需要预测inputs_embedstorch.where(positions0,0,inputs_embeds)# 2. 两路归一化inputs_embedsself.enorm(inputs_embeds)previous_hidden_statesself.hnorm(previous_hidden_states)# 3. 拼接 线性投影得到新的 hiddenhidden_statesself.eh_proj(torch.cat([inputs_embeds,previous_hidden_states],dim-1))# 4. 经过一个 Transformer blockhidden_states,residualself.mtp_block(positions,hidden_states,None)returnresidualhidden_states3. 多层 MTP 的组织[deepseek_mtp.py#L124-L182](file:///workspace/vllm/model_executor/models/deepseek_mtp.py#L124-L182)DeepSeekMultiTokenPredictor持有num_nextn_predict_layers个 MTP 层DeepSeek-V3 默认 1 层。single-module MTP会把同一层重复使用通过spec_step_idx % num_mtp_layers选层multi-module MTP则每一层负责一个草稿位置。三、MTP 推测解码完整流程vLLM V1 中 MTP 复用EagleProposer见 [eagle.py](file:///workspace/vllm/v1/spec_decode/eagle.py)区别仅在于pass_hidden_states_to_modelTrue且 method 为mtp。整个流程由 [llm_base_proposer.py](file:///workspace/vllm/v1/spec_decode/llm_base_proposer.py) 的propose方法编排阶段 0主模型正常 decode 一步主模型对当前 token 做一次 forward输出target_token_ids当前 tokentarget_hidden_states最后一个 token 的隐状态关键输入next_token_ids采样得到的下一个 token即 bonus token 候选阶段 1第一轮草稿[llm_base_proposer.py#L427-L473](file:///workspace/vllm/v1/spec_decode/llm_base_proposer.py#L427-L473)把next_token_ids当作 MTP 的input_ids把主模型的target_hidden_states作为previous_hidden_states送入 MTP 模块eh_proj([embed(next_token), hnorm(hidden)]) → mtp_block → hidden对hidden调用compute_logits_greedy_sample得到第 1 个 draft token若num_speculative_tokens 1直接返回进入验证。阶段 2迭代生成剩余草稿[llm_base_proposer.py#L525-L588](file:///workspace/vllm/v1/spec_decode/llm_base_proposer.py#L525-L588)fortoken_indexinrange(self.num_speculative_tokens-1):input_idsdraft_token_ids_list[-1].int()# 上一个 draft token...model_kwargs{input_ids:input_ids,positions:...,hidden_states:self.hidden_states[...],# 上一轮 MTP 输出}ret_hidden_statesself.model(**model_kwargs)# 再次走 MTP 层draft_token_idsself._greedy_sample(last_hidden_states)draft_token_ids_list.append(draft_token_ids)每一步把上一轮的 draft token 上一轮的 hidden state 喂回 MTP串行生成 D 个草稿 token。阶段 3主模型并行验证把[bonus_token, draft_1, draft_2, ..., draft_D]拼成一个序列主模型一次 forward 同时计算这 D1 个位置的 logits由 [RejectionSampler](file:///workspace/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py) 执行验证贪婪策略草稿 token 与主模型 argmax 一致则接受遇到第一个不匹配即停止后续全部拒绝拒绝采样策略对每个草稿 token 验证P_target / P_draft ≥ UU 为均匀随机数接受则继续否则从max(0, 1 - P_target/P_draft)概率重采样替代 token阶段 4输出第一个 bonus token 必然正确来自主模型自身后续被接受的草稿 token 直接产出第一个被拒绝的位置由主模型重采样得到修正 token一次 forward 可能产出1 ~ D1个 token四、举例说明配置num_speculative_tokens 3MTP 模块为 single-module重复使用同一层。Prompt 为中国的首都是主模型已生成北京现在要继续生成。Step A主模型 decode 第 N 步输入 token北京主模型 forward → hidden_stateh_N采样得到 bonus tokennext_token_id 是Step BMTP 起草3 个草稿Draft 1input_ids 是prev_hidden h_NMTP forwardeh_proj([embed(是), hnorm(h_N)]) → mtp_block → h₁greedy sample →中国Draft 2input_ids 中国prev_hidden h₁MTP forward →h₂greedy sample →的首Draft 3input_ids 的首prev_hidden h₂MTP forward →h₃greedy sample →都草稿序列[是, 中国, 的首, 都]第 1 个是 bonus后 3 个是 MTP 草稿Step C主模型并行验证将[是, 中国, 的首, 都]作为 4 个连续位置一次性送入主模型得到每个位置的 target logits。假设主模型的 argmax 结果为位置草稿 token主模型 argmax是否匹配0是是✅ (bonus必然接受)1中国中国✅ 接受2的首的首✅ 接受3都政治❌ 拒绝Step D输出接受[是, 中国, 的首]三个 token在位置 3 用主模型 logits 重采样得到政治作为修正 token本次 decode 共产出4 个 token是中国的首政治下一步从政治继续回到 Step A反例草稿质量差时若 Draft 1 就被拒绝主模型认为是在而非是则只产出 bonus token在 修正 token相当于退化为普通 decode但因 MTP 模块很轻量开销远小于一次完整主模型 forward。五、关键设计要点总结无损性通过拒绝采样保证 MTP 输出分布与主模型自回归解码完全等价数学可证。参数共享MTP 与主模型共享 embedding 和 lm_head新增参数仅 ~14B/671B约 2%。hidden state 复用MTP 直接消费主模型最后一层 hidden state避免重复计算这是与 EAGLE 的相似之处也是 vLLM 中 MTP 复用EagleProposer的原因。串行起草 并行验证起草 D 个 token 需要 D 次轻量 MTP forward但验证只需 1 次主模型 forward主模型算力被充分利用。single-module vs multi-modulevLLM 单模块复用同一 MTP 层spec_step_idx % num_mtp_layers多模块则为每个位置配独立 MTP 层需要 re-prefill 机制修正被拒绝 token 留下的 stale KV cache见 PR #48892。Sources:DeepSeek-V3 Technical Report (arXiv:2412.19437)vLLM MTP 文档vLLM Ascend MTP 文档vLLM PR #48892: Multi-Module MTP support