执行摘要
- 一句话:修复Megatron CP+TP序列对齐与新MTP API适配
- 推荐动作:该PR修复了重要的并行对齐bug并适配了新API,值得对Megatron引擎开发者和经验丰富的用户精读。关键设计决策:使用条件导入实现API版本兼容,以及对齐公式背后的推导逻辑(乘积而非lcm)值得记录。建议为修改的路径补充示例或测试。
功能与动机
根据PR body描述,需要修复两个问题:1)当同时启用CP和TP时,seq_len应该填充到(2 * cp) * tp;2)适配新的mtp_loss API以兼容Megatron dev分支。这些修复确保在CP+TP配置下序列填充正确,并保持与新版Megatron MTP计算接口的兼容性。
实现拆解
- 序列长度对齐修正:在
verl/models/mcore/util.py的preprocess_bshd_engine和build_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操作。
- 新版MTP API适配:在
verl/models/mcore/mtp_patch.py中,通过try/except导入process_mtp_loss并定义_HAS_PROCESS_MTP_LOSS标志。在_megatron_gptmodel_postprocess函数中,当_HAS_PROCESS_MTP_LOSS为True时,调用_process_mtp_loss一次性完成hidden_states的chunk、rolling、loss scaling和输出层计算,替代原先的手动循环。同时为函数签名新增is_spec_decode参数以对齐新版接口。
- 向后兼容保留:若导入失败(旧版Megatron),
_HAS_PROCESS_MTP_LOSS为False,继续执行原有的逐层MTP损失计算逻辑,不破坏现有功能。
- 辅助调整:修正
cp_group的获取方式,优先从self.cp_group获取,否则回退到self.pg_collection.cp(若存在)。
关键文件:
verl/models/mcore/mtp_patch.py(模块 模型层;类别 source;类型 data-contract;符号 _HAS_PROCESS_MTP_LOSS, _process_mtp_loss, _megatron_gptmodel_postprocess, patch_postprocess): 核心文件:适配新版Megatron MTP损失API,增加新/旧API自动选择路径,调整函数签名。
verl/models/mcore/util.py(模块 模型层;类别 source;类型 data-contract;符号 preprocess_bshd_engine, build_vlm_attn_mask_bshd): 修改序列长度对齐规则,保证CP+TP配置下填充正确。
关键符号:_megatron_gptmodel_postprocess, preprocess_bshd_engine, build_vlm_attn_mask_bshd
关键源码片段
verl/models/mcore/mtp_patch.py
核心文件:适配新版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 计算 ...
评论区精华
本PR无实质review讨论。Gemini-code-assist[bot]仅做了评论摘要并声明无反馈,维护者wuxibin89直接批准合并。
风险与影响
- 风险:
- 向后兼容风险:新API路径通过
try/except保护,旧路径保留,风险可控。但若新版Megatron的process_mtp_loss接口签名与预期不符,可能引发运行时错误。
- 序列对齐变更风险:对齐从
lcm(tp_size, 2*cp_size)改为tp_size * cp_size * 2,在cp=1时结果相同(均为tp_size),但cp>1时新值更严格(可能更大),可能导致更多的padding和显存占用,但不会出错。
- 缺少测试覆盖:未新增单元测试或E2E测试验证新API路径,存在回归隐患。
- 参数变更:函数签名增加
is_spec_decode参数,调用处若未更新可能遗漏参数(但patch可见调用处已处理)。
- 影响:直接影响使用Megatron引擎并启用Multi-Token Prediction(MTP)的用户,特别是同时使用CP和TP的用户。修复了之前因对齐不足可能导致的隐藏状态错位或计算错误,不影响纯TP或纯CP场景。对新版Megatron dev分支的用户,PR提供了MTP计算路径的平滑迁移。整体影响范围较小,但正确性关键。
- 风险标记:核心路径变更, 缺少测试覆盖, 向后兼容风险
关联脉络
参与讨论