Prhub

#6432 [megatron,rollout] fix: align MTP loss and rollout metrics

原始 PR 作者 xhx1022 合并时间 2026-05-25 11:48 文件变更 13 提交数 7 评论 2 代码增减 +598 / -34

执行摘要

修复 Megatron MTP 训练与 rollout 指标对齐

原始MTP loss计算通过ColumnParallelLinear直接传递梯度到lm_head,导致验证/基础模型分布偏移;loss mask通过通用nested tensor转换,错将prompt位置标记为有效MTP位置;activation recompute因非tensor元数据PackedSeqParams传入checkpoint wrapper而崩溃;reference模型因HF配置加载MTP块浪费显存且导致labels.clone() on None失败。详见PR body Bug Background。

该PR修复了多个MTP训练关键bug(loss污染、mask错位、recompute崩溃),建议优先合并。但需关注bias处理问题,确认目标模型输出层是否使用bias;若使用,后续应跟踪补充bias加法逻辑。代码结构清晰,但测试覆盖不足是主要风险。建议团队内部补充CI测试用例。

讨论亮点

在代码审查中,gemini-code-assist[bot]指出functional_call中未处理输出层bias:ColumnParallelLinear可能返回(logits, bias),忽略bias会导致MTP loss计算错误。建议修正bias加法并优化参数字典初始化。该评论为critical级别,但合并者直接批准,未要求额外修改,问题在当前PR中仍存在。

实现拆解

  1. MTP loss梯度隔离verl/models/mcore/mtp_patch.py):在Legacy API _megatron_gptmodel_postprocess中,将compute_output_layer_and_language_model_loss替换为torch.func.functional_call,对output_layer.named_parameters()output_weight执行detach(),确保MTP辅助loss不会更新lm_head

  2. MTP loss mask对齐verl/models/mcore/model_forward.py):新增_build_mtp_loss_mask_nested函数,根据input_ids_lengths将response mask扩展为[prompt zeros; response mask]格式,与packed input_ids对齐;在gptmodel_forward_model_engine中MTP模式下调用此函数替代通用nested tensor转换。

  3. Activation recompute修复verl/models/mcore/mtp_patch.py):新增patch_mtp_layer_checkpointed_forwardunpatch_mtp_layer_checkpointed_forward,通过monkey patch MultiTokenPredictionLayer._checkpointed_forward,使checkpoint wrapper只保存tensor输入,非tensor参数在recompute closure内恢复。

  4. Reference模型MTP禁用verl/workers/engine/megatron/transformer_impl.py):在forward_only且mtp_num_layers=0时,对ref模型调用patch_postprocess覆盖MTP后处理,避免加载未使用的MTP层,同时防止labels=None的异常。

  5. Speculative decoding metric聚合verl/trainer/ppo/ray_trainer.pyverl/workers/rollout/vllm_rollout/vllm_async_server.pyverl/workers/rollout/sglang_rollout/async_sglang_server.py):新增compute_spec_decode_metrics函数,按per-request宏平均聚合draft/accept/verify统计;在vLLM和sglang的rollout后端收集原始数据;在PPO trainer的fit循环中集成metric更新。

  6. vLLM MTP drafter权重同步verl/workers/rollout/vllm_rollout/utils.py):添加_get_drafter_model_use_mtp_drafter_weight_sync_iter_all_models等辅助方法,使update_weights_from_ipcmonkey_patch_model等函数能同时作用于主模型和MTP drafter;仅在speculative method为"mtp"时同步drafter权重。

文件 模块 状态 重要度
verl/models/mcore/mtp_patch.py 模型层 modified 8.71
verl/models/mcore/model_forward.py 模型层 modified 7.44
verl/workers/rollout/vllm_rollout/utils.py 采样层 modified 7.88
verl/trainer/ppo/ray_trainer.py 训练器 modified 6.89
verl/workers/engine/megatron/transformer_impl.py 引擎层 modified 6.53
verl/utils/vllm/vllm_fp8_utils.py 工具层 modified 6.41
examples/mtp_trainer/run_mimo_7b_mtp_rl_vllm_sgl_megatron.sh 示例脚本 added 5.37

关键符号

compute_spec_decode_metrics _build_mtp_loss_mask_nested patch_mtp_layer_checkpointed_forward _patched_checkpointed_forward load_quanted_weights _get_drafter_model _use_mtp_drafter_weight_sync _iter_all_models

关键源码片段

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

核心修改:MTP loss 梯度隔离、activation recompute patch,影响训练正确性和稳定性

# 在 GPTModel 的 _megatron_gptmodel_postprocess 函数中,
# 当 mtp_num_layers > 0 且 labels 不为 None 且未使用新版 process_mtp_loss API 时,
# 进入旧版 Manual Rolling 分支。此段展示 MTP loss 计算的关键修改:
else:
    mtp_labels = labels.clone()
    hidden_states_list = torch.chunk(hidden_states, 1 + self.config.mtp_num_layers, dim=0)
    hidden_states = hidden_states_list[0]
    if loss_mask is None:
        loss_mask = torch.ones_like(mtp_labels)
    cp_group = getattr(self, "cp_group", None)
    for mtp_layer_number in range(self.config.mtp_num_layers):
        mtp_labels, _ = roll_tensor(
            mtp_labels, shifts=-1, dims=-1,
            cp_group=cp_group, packed_seq_params=packed_seq_params,
        )
        loss_mask, num_tokens = roll_tensor(
            loss_mask, shifts=-1, dims=-1,
            cp_group=cp_group, packed_seq_params=packed_seq_params,
        )
        # Detach output-layer params so MTP loss does not update lm_head.
        output_layer_params = {k: v.detach() for k, v in self.output_layer.named_parameters()}
        output_layer_buffers = dict(self.output_layer.named_buffers())
        # 使用 functional_call 计算 logits,参数都是 detach 状态,不会产生梯度到 lm_head
        mtp_logits, _ = torch.func.functional_call(
            self.output_layer,
            {**output_layer_params, **output_layer_buffers},
            args=(hidden_states_list[mtp_layer_number + 1],),
            kwargs={
                "weight": output_weight.detach() if output_weight is not None else None,
                "runtime_gather_output": runtime_gather_output,
            },
        )
        mtp_loss = self.compute_language_model_loss(mtp_labels, mtp_logits)
        mtp_loss = loss_mask * mtp_loss
        if self.training:
            MTPLossLoggingHelper.save_loss_to_tracker(
                mtp_loss, ... # 省略细节
            )
verl/models/mcore/model_forward.py data-contract

新增 `_build_mtp_loss_mask_nested` 函数,修正 MTP loss mask 对齐

def _build_mtp_loss_mask_nested(response_mask, input_ids_lengths, response_attention_mask):
    """Build a nested loss_mask aligned to ``input_ids = [prompt; response]`` for MTP.    ``response_mask`` is response-only data. This expands it to full packed
    input length as prompt zeros followed by valid response positions.
    """
    if isinstance(response_mask, NestedTensor):
        response_offsets = response_mask.offsets().tolist()
        response_lengths = [response_offsets[i + 1] - response_offsets[i] for i in range(len(response_offsets) - 1)]
        batch_size = len(response_lengths)
        response_values = response_mask.values()
    else:
        # 对于 padded 格式,从 response_attention_mask 提取每个样本的有效 response 长度
        assert response_attention_mask is not None, "response_attention_mask is required"
        assert not isinstance(response_attention_mask, NestedTensor)
        assert response_attention_mask.shape == response_mask.shape
        batch_size = response_mask.shape[0]
        response_lengths = response_attention_mask.to(torch.int32).sum(dim=-1).tolist()
​
    assert len(input_ids_lengths) == batch_size
​
    pieces = []
    for i in range(batch_size):
        actual_total = int(input_ids_lengths[i])
        actual_response = int(response_lengths[i])
        actual_prompt = actual_total - actual_response
        assert actual_prompt >= 0
        # 构建 [prompt zeros; response mask] 的 loss_mask
        prompt_pad = torch.zeros(actual_prompt, dtype=response_mask.dtype, device=response_mask.device)
        if isinstance(response_mask, NestedTensor):
            response_piece = response_values[response_offsets[i] : response_offsets[i + 1]]
        else:
            response_piece = response_mask[i, :actual_response]
        full = torch.cat([prompt_pad, response_piece], dim=0)
        assert full.shape[0] == actual_total
        pieces.append(full)
​
    return torch.nested.nested_tensor(pieces, layout=torch.jagged)
verl/workers/rollout/vllm_rollout/utils.py core-logic

添加 drafter weight sync 支持,使 MTP drafter 也能接收 actor 权重

class VllmWorker:
    # ... 现有方法 ...
​
    def _get_drafter_model(self):
        """Return the drafter's model object, or None if unavailable."""
        drafter = getattr(self.model_runner, "drafter", None)
        return drafter.model if drafter is not None and hasattr(drafter, "model") else None
​
    def _get_draft_model_config(self):
        """Return the draft model config from speculative_config, or None."""
        spec = self.model_runner.vllm_config.speculative_config
        return spec.draft_model_config if spec is not None and spec.draft_model_config is not None else None
​
    def _use_mtp_drafter_weight_sync(self):
        """Return whether the vLLM MTP drafter should receive actor weights."""
        spec = self.model_runner.vllm_config.speculative_config
        # 仅当 speculative method 为 mtp 且 drafter 存在时才同步
        return spec is not None and spec.method == "mtp" and self._get_drafter_model() is not None
​
    def _iter_all_models(self):
        """Yield models that need weight updates.
        Only vLLM MTP drafter sync is supported for now.
        """
        yield self.model_runner.model
        if self._use_mtp_drafter_weight_sync():
            yield self._get_drafter_model()
​
    def _iter_all_models_with_config(self):
        """Yield (model, model_config) for models that need post-processing."""
        yield self.model_runner.model, self.model_runner.vllm_config.model_config
        if self._use_mtp_drafter_weight_sync():
            draft_cfg = self._get_draft_model_config()
            if draft_cfg is not None:
                yield self._get_drafter_model(), draft_cfg

评论区精华

MTP loss functional_call 忽略 bias 可能导致错误 正确性

gemini-code-assist[bot] 在 mtp_patch.py line 179 评论指出,`torch.func.functional_call` 返回的 `mtp_logits` 可能是一个 `(logits, bias)` 元组,当前代码只取了第一个元素,忽略 bias,导致当输出层使用 bias 时 MTP loss 计算错误。此外,建议优化参数字典初始化以减少循环内开销。

结论:PR 合并时未修改该问题,合并者直接批准。问题仍存在,需要后续跟进。 · 待处理

风险与影响

1) MTP loss梯度隔离依赖functional_call,但未处理输出层bias(如review指出的),若模型使用bias则MTP auxiliary loss错误。
2) MTP loss mask对齐函数改变原数据流,可能影响非packed场景的兼容性。
3) Activation recompute patch为monkey patch,可能与其他自定义checkpoint wrapper冲突。
4) Reference模型MTP禁用通过patch_postprocess覆盖,未来Megatron版本如果调整MTP接口可能导致不兼容。
5) Spec decode metrics聚合函数为新增代码,未经大规模验证,可能边界异常。

直接使用Megatron MTP训练的团队需要更新代码以应用此修复;vLLM/sglang speculative decoding用户将获得正确的accept rate和accept length metric。影响范围主要为训练代码和rollout metric,无API breaking change。建议用户配合使用新增的run_mimo_7b_mtp_rl_vllm_sgl_megatron.sh脚本进行验证。

核心路径变更 缺少测试覆盖 Bias 未处理

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论