
verl MTP 多 Token 预测训练与推理加速实战指南【免费下载链接】verlverl/HybridFlow: A Flexible and Efficient RL Post-Training Framework项目地址: https://gitcode.com/GitHub_Trending/ve/verlMulti-Token PredictionMTP多 Token 预测通过在解码器主干之外挂载一个轻量级草稿头让模型在预测下一个 Token 的同时预测后续多个 Token既能作为辅助训练目标提升样本利用效率又能作为投机解码Speculative Decoding的草稿来源加速推理。本文以 docs/advance/mtp.md 为主体结合 verl 仓库中的配置定义、Megatron 桥接补丁与可运行脚本完整讲解在 verl 中开启 MTP 训练RL/SFT与 Rollout 推理加速的适用范围、依赖版本、核心参数、源码原理与实测结论。读完本文你将能够为 MiMo-7B-RL、Qwen-next、DeepSeek 系列等携带 MTP 模块的模型配置一套可复现的 MTP 训练流程并了解何时应该以及何时不应该在推理阶段启用 MTP 加速。一、支持范围训练与推理引擎的兼容矩阵MTP 属于对模型结构有硬性要求的高级特性需要模型原生携带 MTP 头因此 verl 对其支持范围有明确约束训练引擎仅支持mbridge/Megatron-Bridge Megatron组合其他训练引擎如 FSDP、VeOmni当前不兼容推理引擎兼容所有引擎但模型必须位于对应引擎的兼容列表中即 SGLang、vLLM 均可用作 MTP 模型的 Rollout/推理后端前提是模型受引擎支持可训练模型当前支持基于 MTP 架构的 mimo-7B-RL、Qwen-next 与 Deepseek 系列模型。依赖版本与补丁要求由于 MTP 涉及 Megatron 前向、损失计算与权重同步等多处改动官方文档给出了精确到 commit 的依赖清单组件版本要求说明mbridge应用 PR #62 的补丁与评审建议已合入 main 分支Megatron-Bridge 侧的 MTP 支持Megatron-Bridge尝试 mimo-7B-RL 时需应用 PR #2387 的补丁未来合入 main桥接层对 MTP 结构的支持Megatron-LM使用最新 dev 版本commit23e092f41ec8bc659020e401ddac9576c1cfed7e该版本支持 MTP CP 训练方式Megatron-LM启用recompute_granularityfull时需要包含 PR #3457commitffd66a3e62026-06-03的 dev 版本该 PR 将padding_mask穿透到MultiTokenPredictionLayer._checkpointed_forward缺少它时MultiTokenPredictionLayer.forward会传入未声明关键字参数导致第一步训练即抛TypeError。已发布的megatron-core0.18.0 / 0.18.2 也未携带该修复跟踪于 issue #4933sglang使用分支fix_mtp_update_weights_from_tensor对应 PR #17870修复 MTP 从 tensor 更新权重时的 OOM 问题注上表提到的 PR 号与 commit 哈希均来自原文档记录其中#3457对应的 dev commit 比默认 pin 的23e092f41晚约半年两者选其一需按是否启用 full recompute 决定。SFT 示例脚本 examples/sft/gsm8k/run_mimo_7b_mtp_megatron.sh 中给出了可复现的依赖准备方式克隆Megatron-LM的 dev 分支并 checkout 到23e092f41ec8bc659020e401ddac9576c1cfed7e克隆 mbridge 的feature/verl_mtp分支并 checkout 到6bf2d45a15dc4fb52d2f0c38ff546bee33447d10随后将两者加入PYTHONPATH即可。另外下载 MiMo-7B-RL 后需手动在其config.json中设置max_position_embeddings: 32768见 examples/mtp_trainer/README.md否则长序列训练会出问题。二、MTP 训练配置核心参数全解verl 中所有 MTP 相关配置统一挂在actor_rollout_ref.model.mtp前缀之下其数据结构定义在 verl/workers/config/model.py 的MtpConfig类中。源码注释与字段默认值如下class MtpConfig(BaseConfig): enable: bool False # 加载/保存 MTP 参数但不使用 enable_train: bool False # 训练时使用 MTP 参数 enable_rollout: bool False # Rollout 时使用 MTP 参数推理加速 # 训练参数 detach_encoder: bool False # 训练时是否分离冻结编码器 mtp_loss_scaling_factor: float 0.1 # MTP 损失缩放系数 # SGLang Rollout 参数 speculative_algorithm: str EAGLE speculative_num_steps: int 3 speculative_eagle_topk: int 1 speculative_num_draft_tokens: int 4 # vLLM Rollout 参数 method: str mtp num_speculative_tokens: int 1四种典型配置场景原文档将 MTP 训练归纳为四种场景完整参数组合如下配置场景核心参数说明仅加载 MTP 参数enableTrueVRAM 占用会增加但导出的参数包含 MTP 模块可直接用于线上部署全参数 MTP 训练enableTrue、enable_trainTrue、mtp_loss_scaling_factor0.1MTP Loss 会作用于全部模型参数MTP 参数专属训练enableTrue、enable_trainTrue、detach_encoderTrue冻结 Encoder 层只更新 MTP 模块参数MTP Loss 仅作用于 MTP 参数MTP 加速 RolloutvLLMenableTrue、enable_rolloutTrue、methodmtp、num_speculative_tokens1基于 MTP 实现 Rollout 阶段推理加速MTP 加速 RolloutSGLangenableTrue、enable_rolloutTrue、speculative_algorithmEAGLE、speculative_num_steps2、speculative_eagle_topk2、speculative_num_draft_tokens4同上走 EAGLE 投机解码路径注意 SGLang 场景下源码默认值是speculative_num_steps3、speculative_eagle_topk1而原文档推荐组合为speculative_num_steps2、speculative_eagle_topk2使用时请以实际需求为准覆盖默认值。与模型配置的联动HFModelConfig在 verl/workers/config/model.py 中对mtp.enable做了联动处理若enableFalse会把 HF 配置中的mtp_num_hidden_layers或text_config.mtp_num_hidden_layers即 Qwen3.5 风格的嵌套配置置为 0避免加载冗余的 MTP 层enableTrue时则保留原始 MTP 结构用于加载与训练。三、完整 RL 训练脚本MiMo-7B Megatron 实战examples/mtp_trainer/run_mimo_7b_mtp_megatron.sh 是 MTP RL 训练的标准入口脚本SGLang Rollout Megatron 训练同步混合引擎模式其 MTP 相关配置为actor_rollout_ref.model.mtp.enableTrue actor_rollout_ref.model.mtp.enable_trainTrue actor_rollout_ref.model.mtp.mtp_loss_scaling_factor${mtp_loss_scaling_factor:-0.1} actor_rollout_ref.model.mtp.detach_encoderTrue即官方推荐的冻结编码器、只训练 MTP 头方案detach_encoderTrue。其余关键工程参数包括模型与并行actor_tp2、actor_pp2、actor_cp2Megatron 侧 TP×PP×CP8 卡rollout_tp4并开启param_offloadTrue、optimizer_offloadTrue、use_mbridgeTrue数据与算法GRPOalgorithm.adv_estimatorgrpomax_prompt_length2048、max_response_length8192ppo_max_token_len_per_gpu20480奖励reward.reward_manager.namedapo并配置 overlong buffer 惩罚len4096、penalty_factor1.0启动方式GPU 环境下默认通过uv run --frozen --all-packages --extra sglang --extra megatron python3 -m verl.trainer.main_ppo启动同时把 Ray worker 的py_executable指向同一 uv 环境NPU 环境则回退到系统 Python。训练 400 步、batch size 128、响应上限 8k 的完整命令可参考脚本顶部的环境变量覆盖如MODEL_PATH、TRAIN_BATCH_SIZE、TOTAL_TRAINING_STEPS等。四、源码级原理verl 如何把 MTP 接进 MegatronMTP 在 verl 中的落地依赖对 MegatronGPTModel的一组运行时补丁全部集中在 verl/models/mcore/mtp_patch.py并由 verl/workers/engine/megatron/transformer_impl.py 在训练引擎初始化时调用check_mtp_config、patch_engine_mtp、get_megatron_mtp_loss等。核心机制如下1. 前向流程补丁patch_postprocess_megatron_gptmodel_postprocess源码中注明拷贝自 Megatron-LM commit23e092f41...的gpt_model.py重写了GPTModel._postprocess当mtp_in_postprocess且存在labels时先调用self.mtp(...)得到 MTP 层的输出 hidden states再走主输出层补丁会按MultiTokenPredictionLayer.forward是否声明padding_mask参数signature探测决定是否透传从而规避上文提到的#3457兼容问题同时兼容mhc_multistream多流等新特性。2. MTP 损失计算新老两套 API 分支_megatron_gptmodel_postprocess对 MTP Loss 做了版本分支新版 Megatron含process_mtp_loss直接调用_process_mtp_loss由 Megatron 内部完成 chunking、rolling、loss scalingverl 额外解析 CP group优先使用 Dynamic CP 提供的 per-microbatch 分组并传入cp_group、tp_group旧版 Megatronlegacy 路径手动对labels与loss_mask执行roll_tensor(shifts-1)用torch.func.functional_call以detach 后的 output layer 参数计算 MTP logits保证 MTP Loss 不回传 lm_head最后通过MTPLossAutoScaler.apply按mtp_loss_scaling_factor / mtp_num_layers缩放梯度。3.detach_encoder的实现patch_mtp_layer_get_embeddingsMTP 参数专属训练的关键是detach_encoder。_patched_get_embeddings_for_detach对MultiTokenPredictionLayer._get_embeddings打补丁对input_ids、position_ids以及可选padding_mask执行roll_tensor后计算 token embedding随后执行decoder_input decoder_input.detach()并令主解码器 hidden statesrequires_gradTrue——主模型Encoder的梯度路径被切断只有 MTP 模块可更新与文档冻结 Encoder、只训 MTP的语义完全对应。配套的unpatch_mtp_layer_get_embeddings在训练结束后还原。4. 激活重计算recompute支持patch_mtp_layer_checkpointed_forwardMTP 层做激活重计算时_patched_checkpointed_forward会把张量参数与非张量参数常量分类打包重计算时只保存张量并重建 kwargs避免非张量参数进入 checkpoint 的 saved tensors。约束如下recompute_method uniform时要求recompute_num_layers 1recompute_method block暂不支持 MTP会告警并跳过重计算。5. 推理侧与损失归一化verl/models/mcore/patch.py 的apply_mtp_inference_patch在推理无 labels时临时把mtp_num_layers置为None再调用原始_postprocess避免 MTP 层在纯生成路径产生额外开销在 MoE 路由场景下transformer_impl.py 会计算mtp_loss_normalization_factor routed_num_tokens / batch_num_tokens通过_get_mtp_loss_config乘到mtp_loss_scaling_factor上使 MTP Loss 的逐 token 归一化与主 Loss 对齐calculate_per_token_lossTrue时生效。五、Rollout 阶段启用 MTPvLLM 与 SGLang 配置MTP 头天然可作为投机解码的草稿模型。启用方式即第一节表中的两组参数# vLLM 后端 actor_rollout_ref.model.mtp.enableTrue actor_rollout_ref.model.mtp.enable_rolloutTrue actor_rollout_ref.model.mtp.methodmtp actor_rollout_ref.model.mtp.num_speculative_tokens1 # SGLang 后端EAGLE 算法 actor_rollout_ref.model.mtp.enableTrue actor_rollout_ref.model.mtp.enable_rolloutTrue actor_rollout_ref.model.mtp.speculative_algorithmEAGLE actor_rollout_ref.model.mtp.speculative_num_steps2 actor_rollout_ref.model.mtp.speculative_eagle_topk2 actor_rollout_ref.model.mtp.speculative_num_draft_tokens4需要注意的是SGLang 后端需使用官方文档指定的fix_mtp_update_weights_from_tensor分支否则从 tensor 更新 MTP 权重时会 OOM对应上游 PR #17870。在训练器侧verl/trainer/ppo/v1/trainer_base.py 会在mtp.enable and mtp.enable_rollout时从 Rollout 产出的extra_fields中抓取投机解码统计spec_num_draft_tokens、spec_num_accepted_tokens、spec_num_verify_steps用于日志与指标统计若后端不报告逐请求投机统计这三个值保持为None不影响训练流程。此外examples/mtp_trainer 目录还提供了两种变体run_mimo_7b_mtp_rl_vllm_sgl_megatron.shSGLang/vLLM 双后端可选对齐 slime 的 RL/EAGLE 设置run_mimo_7b_mtp_fully_async_megatron_multinode.sh完全异步、训练/推理分离部署的多节点方案基于verl.experimental.fully_async_policy.fully_async_main默认单节点 44 切分trainer rollout用于冒烟测试扩展示例TRAIN_NNODES4 TRAIN_NGPUS_PER_NODE8 \ ROLLOUT_NNODES4 ROLLOUT_NGPUS_PER_NODE8 \ bash examples/mtp_trainer/run_mimo_7b_mtp_fully_async_megatron_multinode.sh六、实验结论哪些配置有效哪些无效原文档基于mimo-7B-math、max_response_length8k的实验给出如下结论完整图表见原文档引用的 wandb 链接无明显效果的场景基座模型本身不携带 MTP 参数基座模型携带 MTP 参数但训练时未训练 MTP 模块基座模型携带 MTP 参数且训练 MTP但mtp_loss_scaling_factor0基座模型携带 MTP 参数训练 MTP 并 detach encoder且mtp_loss_scaling_factor0.1。有明显效果的场景唯一基座模型携带 MTP 参数MTP Loss 作用于全部模型参数且mtp_loss_scaling_factor0.1。官方推荐虽然全参数 MTP 训练效果最显著官方仍推荐采用detach_encoderTrue的 MTP 参数专属训练方案——它在显著降低显存/计算开销的同时能够获得接近的收益且 MTP 模块参数独立更新、更易调试。从源码角度可以佐证MTP Loss 作用于全部参数的机制legacy 路径中mtp_loss_scale config.mtp_loss_scaling_factor / config.mtp_num_layers通过MTPLossAutoScaler把梯度回传至主模型 hidden states而detach_encoderTrue时decoder_input与主解码器之间的梯度链路被.detach()切断见第四节第 3 点损失只能驱动 MTP 模块自身参数更新。七、MTP 推理加速的性能注意事项文档特别强调了 MTP 加速 Rollout 的两面性启用 MTP 后Rollout 的投机接受率acceptance rate提升约14%但在 H20 GPU 上整体吞吐不仅没有提升反而略微下降。原因在于接受率收益会被模型规模与推理硬件的算力差距抵消。原文档给出的硬件 FP16 Tensor Core 算力参考硬件型号FP16 性能TFLOPSH20148H8001,671H2001,979实测案例将 mimo-7B 单独部署在 H20 上、使用 SGLang 启用 MTP 投机解码后Rollout 吞吐下降约50%。MTP 加速的收益高度依赖草稿生成与验证的算力比当草稿头的额外计算在低算力卡H20 的 FP16 算力约为 H800 的 1/11上占比过大时投机解码省下的解码步数不足以覆盖草稿开销。因此官方给出两条结论当前优先建议暂不在推理阶段启用 MTP 加速未来规划持续优化 Rollout 阶段的投机解码逻辑以提升吞吐。八、SFT 训练同一套配置、同一套流程MTP 同样支持 SFT 训练配置与 RL 训练完全一致同样以model.mtp.*前缀传入。参考脚本为 examples/sft/gsm8k/run_mimo_7b_mtp_megatron.sh其 MTP 核心配置为model.mtp.enableTrue engine.use_mbridgeTrue engine.override_transformer_config.recompute_methoduniform engine.override_transformer_config.recompute_granularityfull engine.override_transformer_config.recompute_num_layers1注意该脚本启用了recompute_granularityfull因此按第一节的版本要求Megatron-LM 需使用包含 PR #3457 修复的 dev commit否则第一步训练就会因MultiTokenPredictionLayer._checkpointed_forward收到未声明的padding_mask关键字而抛TypeError。该脚本基于mimo-7B-math gsm8k 数据集的实验结果原文档引用的 wandb 链接显示MTP 层的存在对主 Loss 影响有限但当 MTP 层被 detachdetach_encoderTrue时mtp_loss会收敛到更高的值——这符合直觉MTP 头无法借助 Encoder 的梯度信息只能靠自身参数拟合下一 Token 分布收敛目标相对更难。九、总结与决策清单你的目标推荐配置注意事项导出带 MTP 模块的部署权重mtp.enableTrue显存占用增加权重可直接线上部署RL/SFT 训练 MTPenableTrueenable_trainTruemtp_loss_scaling_factor0.1效果最显著的是全参数训练推荐detach_encoderTrue平衡开销Rollout 加速vLLMenable_rolloutTruemethodmtpnum_speculative_tokens1需模型支持H20 等低算力卡上吞吐可能下降Rollout 加速SGLangenable_rolloutTrue EAGLE 系列参数需使用上游修复 OOM 的分支深入阅读建议配置定义见 verl/workers/config/model.py补丁实现见 verl/models/mcore/mtp_patch.py引擎侧接线见 verl/workers/engine/megatron/transformer_impl.py可直接运行的 RL/SFT 脚本见 examples/mtp_trainer 与 examples/sft/gsm8k。【免费下载链接】verlverl/HybridFlow: A Flexible and Efficient RL Post-Training Framework项目地址: https://gitcode.com/GitHub_Trending/ve/verl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考