Prhub

#6206 [megatron] fix: fix seq_len pad len, and adapt to new mtp_loss api (for megatron dev brance)

原始 PR 作者 zpltys 合并时间 2026-05-06 11:57 文件变更 2 提交数 2 评论 0 代码增减 +82 / -52

执行摘要

修复 Megatron CP+TP 序列对齐与新 MTP API 适配

根据PR body描述,需要修复两个问题:1)当同时启用CP和TP时,seq_len应该填充到(2 * cp) * tp;2)适配新的mtp_loss API以兼容Megatron dev分支。这些修复确保在CP+TP配置下序列填充正确,并保持与新版Megatron MTP计算接口的兼容性。

该PR修复了重要的并行对齐bug并适配了新API,值得对Megatron引擎开发者和经验丰富的用户精读。关键设计决策:使用条件导入实现API版本兼容,以及对齐公式背后的推导逻辑(乘积而非lcm)值得记录。建议为修改的路径补充示例或测试。

讨论亮点

本PR无实质review讨论。Gemini-code-assist[bot]仅做了评论摘要并声明无反馈,维护者wuxibin89直接批准合并。

实现拆解

  1. 序列长度对齐修正:在verl/models/mcore/util.pypreprocess_bshd_enginebuild_vlm_attn_mask_bshd函数中,将对齐规则从math.lcm(tp_size, 2 * cp_size)改为tp_size * cp_size * 2,并更新注释说明原因:zigzag CP下每rank持有seq_len/cp_size个token,之后还需被tp_size整除以支持Sequence Parallel的scatter操作。
  2. 新版MTP API适配:在verl/models/mcore/mtp_patch.py中,通过try/except导入process_mtp_loss并定义_HAS_PROCESS_MTP_LOSS标志。在_megatron_gptmodel_postprocess函数中,当_HAS_PROCESS_MTP_LOSSTrue时,调用_process_mtp_loss一次性完成hidden_states的chunk、rolling、loss scaling和输出层计算,替代原先的手动循环。同时为函数签名新增is_spec_decode参数以对齐新版接口。
  3. 向后兼容保留:若导入失败(旧版Megatron),_HAS_PROCESS_MTP_LOSSFalse,继续执行原有的逐层MTP损失计算逻辑,不破坏现有功能。
  4. 辅助调整:修正cp_group的获取方式,优先从self.cp_group获取,否则回退到self.pg_collection.cp(若存在)。
文件 模块 状态 重要度
verl/models/mcore/mtp_patch.py 模型层 modified 7.7
verl/models/mcore/util.py 模型层 modified 6.23

关键符号

_megatron_gptmodel_postprocess preprocess_bshd_engine build_vlm_attn_mask_bshd

关键源码片段

verl/models/mcore/mtp_patch.py data-contract

核心文件:适配新版 Megatron MTP 损失 API,增加新 / 旧 API 自动选择路径,调整函数签名。

# verl/models/mcore/mtp_patch.py (head)
# 从 Megatron 尝试导入新版 process_mtp_loss
try:
    from megatron.core.transformer.multi_token_prediction import (
        process_mtp_loss as _process_mtp_loss
    )
    _HAS_PROCESS_MTP_LOSS = True
except ImportError:
    _HAS_PROCESS_MTP_LOSS = False# ... 省略其它导入和辅助函数 ...def _megatron_gptmodel_postprocess(
    self,
    hidden_states,
    input_ids,
    # ... 其它参数 ...
    is_spec_decode=None, # 新增参数,对齐新版 Megatron 接口
):
    # ... 前置逻辑 ...
​
    if self.config.mtp_num_layers and labels is not None:
        if _HAS_PROCESS_MTP_LOSS:
            # 新版 API 路径:一次调用完成 chunk、roll、loss 计算及缩放
            # 获取 cp_group(优先使用属性,否则从 pg_collection 回退)
            cp_group = getattr(self, "cp_group", None) or (
                self.pg_collection.cp if hasattr(self, "pg_collection") else None
            )
            # 如果需要 μP 缩放
            scale_logits_fn = (
                self._scale_logits
                if (hasattr(self, "_scale_logits") and self.config.use_mup)
                else None
            )
            hidden_states = _process_mtp_loss(
                hidden_states=hidden_states,
                labels=labels,
                loss_mask=loss_mask,
                output_layer=self.output_layer,
                output_weight=output_weight,
                runtime_gather_output=runtime_gather_output,
                is_training=self.training,
                compute_language_model_loss=self.compute_language_model_loss,
                config=self.config,
                cp_group=cp_group,
                packed_seq_params=packed_seq_params,
                scale_logits_fn=scale_logits_fn,
            )
        else:
            # 旧版 API 路径(保留原始逐层手动计算)
            # 此处在 head 中未展示,但功能保持原状
            pass
    # ... 后续 logits 及 loss 计算 ...

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

  1. 向后兼容风险:新API路径通过try/except保护,旧路径保留,风险可控。但若新版Megatron的process_mtp_loss接口签名与预期不符,可能引发运行时错误。
  2. 序列对齐变更风险:对齐从lcm(tp_size, 2*cp_size)改为tp_size * cp_size * 2,在cp=1时结果相同(均为tp_size),但cp>1时新值更严格(可能更大),可能导致更多的padding和显存占用,但不会出错。
  3. 缺少测试覆盖:未新增单元测试或E2E测试验证新API路径,存在回归隐患。
  4. 参数变更:函数签名增加is_spec_decode参数,调用处若未更新可能遗漏参数(但patch可见调用处已处理)。

直接影响使用Megatron引擎并启用Multi-Token Prediction(MTP)的用户,特别是同时使用CP和TP的用户。修复了之前因对齐不足可能导致的隐藏状态错位或计算错误,不影响纯TP或纯CP场景。对新版Megatron dev分支的用户,PR提供了MTP计算路径的平滑迁移。整体影响范围较小,但正确性关键。

核心路径变更 缺少测试覆盖 向后兼容风险

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论