Prhub

#5898 [model] feat: support qwen35 mtp sft/rl

原始 PR 作者 zpltys 合并时间 2026-04-24 10:43 文件变更 6 提交数 4 评论 7 代码增减 +344 / -8

执行摘要

支持 Qwen3.5 MTP 模型在 Megatron 引擎的 SFT/RL 训练

Qwen3.5模型引入了Multi-Token Prediction(MTP)特性,但现有verl框架仅支持DeepSeek/Qwen3风格的num_nextn_predict_layers字段,无法识别Qwen3.5的mtp_num_hidden_layers及其嵌套在text_config中的情况。为使SFT和RL训练能够利用MTP能力(包括训练时辅助损失和推理时投机解码),需要扩展配置解析和模型注册逻辑。

建议

  • 值得精读:如果团队计划支持更多具备MTP能力的模型,本PR中 _get_mtp_num_layers_set_mtp_num_layers 的设计模式值得参考。
  • 需跟进
    1. 将辅助函数公开化,以便跨模块复用(来自review高优先级建议)。
    2. 修复 config_converter.py 中MTP检测逻辑,复用公共函数并正确处理override。
  • 测试:建议为 megatron_utils.py 中的MTP函数编写单元测试,覆盖三种配置格式。
讨论亮点

评论区精华

gemini-code-assist[bot]:"The MTP configuration helpers _get_mtp_num_layers and _set_mtp_num_layers are useful across different modules... They should be made public by removing the leading underscore." (高优先级)
- 建议:将两个辅助函数改为公开(get_mtp_num_layers / set_mtp_num_layers)以便跨模块复用。

gemini-code-assist[bot]:"The current implementation has two issues:

  1. It duplicates the MTP layer detection logic and is less comprehensive than the get_mtp_num_layers helper (e.g., it misses num_nextn_predict_layers). 2. It potentially ignores user overrides for mtp_loss_scaling_factor." (高优先级)
  • 问题:config_converter.py 中的MTP检测没有利用_get_mtp_num_layers,且硬编码了loss scaling factor为0.1,可能忽略用户通过override传入的值。

wuxibin89:"Please reuse examples/sft/gsm8k/run_qwen3_5_megatron.sh with additional MTP option."
- 建议:SFT无需独立MTP脚本,应在原脚本中增加参数。

wuxibin89:"megatron_workers.py has been deprecated, please do not modify it."
- 约束:不应修改已废弃的megatron_workers.py。

结论:开发者在最终提交中采纳了关于不修改废弃文件的建议,并合并了SFT脚本。但公开函数命名的重构建议和config_converter中的重复逻辑问题未被完全修复。

实现拆解

实现拆解

  1. 新增通用MTP层数操作函数verl/utils/megatron_utils.py
    - 新增 _get_mtp_num_layers(hf_config) 函数,统一处理三种配置格式:num_nextn_predict_layers(DeepSeek/Qwen3风格)、mtp_num_hidden_layers(Qwen3.5风格,直接位于 hf_config)、以及 mtp_num_hidden_layers 嵌套在 hf_config.text_config 中。
    - 新增 _set_mtp_num_layers(hf_config, value) 函数,根据hf_config实际拥有的属性名设置MTP层数,保证写操作兼容性。
    - 重构原有的 check_mtp_config 函数,内部调用上述两个工具函数替代硬编码的属性访问,消除重复逻辑。

  2. 配置转换器扩展verl/models/mcore/config_converter.py
    - 在 hf_to_mcore_config_qwen3moe 函数末尾增加MTP支持:从hf_config(或text_config)中读取 mtp_num_hidden_layersmtp_loss_scaling_factor,设置到生成的 TransformerConfigmtp_num_layersmtp_loss_scaling_factor 字段。
    - 注意:当前实现未复用 _get_mtp_num_layers,且硬编码了默认的loss scaling factor(0.1),可能忽略用户通过override传入的值。

  3. 模型注册表扩展verl/models/mcore/registry.py
    - 在 SupportedModel 枚举中新增 QWEN3_5_MOE = "Qwen3_5MoeForCausalLM"
    - 在配置转换器、初始化器、前向函数、融合前向函数、权重转换器等所有注册表中,将 QWEN3_5_MOE 映射到与 QWEN3_MOE 相同的处理逻辑。

  4. 模型配置初始化优化verl/workers/config/model.py
    - 在 __post_init__ 中当MTP禁用时,同时清空可能存在的三种MTP属性(num_nextn_predict_layersmtp_num_hidden_layers、text_config中的 mtp_num_hidden_layers),避免下游模块需要单独处理。

  5. 示例脚本和Shell配置examples/sft/gsm8k/run_qwen3_5_megatron.shverl/experimental/fully_async_policy/shell/grpo_qwen35_35b_megatron_async.sh
    - 修改SFT示例脚本,添加MTP相关命令行参数(enable/enable_train/detach_encoder/loss_scaling_factor)。
    - 新增GRPO+完全异步策略的Shell脚本,展示Qwen3.5-35B-A3B的MTP配置用法,包含TP/PP/EP并行度建议。

  6. 废弃模块清理verl/workers/megatron_workers.py 的修改在最终提交中被回退或移除)
    - 根据review意见,不修改已废弃的megatron_workers.py。

文件 模块 状态 重要度
verl/utils/megatron_utils.py 工具层 modified 7.54
verl/models/mcore/config_converter.py 模型配置 modified 6.68
verl/workers/config/model.py 工作节点 modified 6.18
verl/experimental/fully_async_policy/shell/grpo_qwen35_35b_megatron_async.sh 实验性 added 5.46
verl/models/mcore/registry.py 模型配置 modified 5.4

关键符号

_get_mtp_num_layers _set_mtp_num_layers check_mtp_config

关键源码片段

verl/workers/config/model.py data-contract

配置初始化:在 __post_init__ 中统一清空禁用 MTP 时的所有字段,减少下游适配负担。

# 在 __post_init__ 方法中,位于 per model patch 之后
# When MTP is disabled, zero out MTP layer counts from hf_config so that
# downstream engine/worker code does not need to handle each MTP field format
# individually.
if not self.mtp.enable:
    if hasattr(self.hf_config, "num_nextn_predict_layers"):
        self.hf_config.num_nextn_predict_layers = 0
    if hasattr(self.hf_config, "mtp_num_hidden_layers"):
        self.hf_config.mtp_num_hidden_layers = 0
    if hasattr(self.hf_config, "text_config") and hasattr(self.hf_config.text_config, "mtp_num_hidden_layers"):
        self.hf_config.text_config.mtp_num_hidden_layers = 0

评论区精华

MTP 辅助函数应公开化 设计

gemini-code-assist[bot] 建议将 _get_mtp_num_layers/_set_mtp_num_layers 改为公开函数,以便跨模块复用并避免 lint 问题。

结论:未采纳。最终提交保持了下划线前缀。ArronHZG 要求遵循建议,但未强制修改。 · 待处理

config_converter 中 MTP 逻辑重复且不完整 正确性

gemini-code-assist[bot] 指出 config_converter.py 中的 MTP 检测未覆盖 num_nextn_predict_layers,且硬编码 loss scaling factor 忽略用户 override。

结论:未采纳。最终提交的代码仍存在这两个问题。 · 待处理

SFT 脚本整合建议 设计

wuxibin89 建议复用现有 SFT 脚本而非创建新脚本。

结论:已采纳。在最终提交中合并在 examples/sft/gsm8k/run_qwen3_5_megatron.sh 中。 · 已解决

不应修改已废弃的 megatron_workers.py other

wuxibin89 指出 megatron_workers.py 已废弃,不应修改。

结论:已采纳。最终提交回退了相关修改。 · 已解决

风险与影响

风险分析

  1. 配置兼容性风险verl/models/mcore/config_converter.py
    - hf_to_mcore_config_qwen3moe 中MTP检测未覆盖 num_nextn_predict_layers,若Qwen3.5未来版本采用该字段名将无法识别。
    - mtp_loss_scaling_factor 硬编码为0.1,若用户通过override指定则被忽略,可能导致训练损失计算错误。

  2. 废弃模块污染verl/workers/megatron_workers.py
    - 虽然最终提交回退了修改,但中途的修改可能残留不一致。审查确认最终版本无异。

  3. 测试覆盖缺失
    - 本次改动涉及多个核心模块(配置转换、模型注册、utils),但无对应的单元测试或集成测试。建议补充测试覆盖MTP配置的读写场景。

  4. API破坏风险
    - 新增 supportedModel.QWEN3_5_MOE 枚举值,但未删除任何枚举,无破坏性变更。
    - _get_mtp_num_layers_set_mtp_num_layers 以下划线开头(内部),理论上不会影响外部API。但review建议公开化,若未来修改命名则影响内部调用。

影响分析

  • 用户影响:需要升级transformers至5.3.0+,并合并mbridge PR #98才能使用MTP。新增的Shell脚本提供了开箱即用的配置。SFT脚本通过额外参数支持MTP,不影响已有用法。
  • 系统影响:MTP配置仅在启用时激活(mtp.enable=True),默认关闭,不增加现有训练的额外开销。配置转换器中的MTP逻辑仅在Qwen3.5模型上生效。
  • 团队影响:代码主要改动集中在megatron_utils.py、config_converter.py、registry.py,构成了未来支持更多MTP模型的参考模式。review建议部分未完全落实(如公开函数命名),需要后续跟进。
配置兼容性风险 忽略用户覆盖参数 缺少测试覆盖

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论