跳转至

mtp_patch.py — MTP (Multi-Token Prediction) 补丁

文件路径

verl/models/mcore/mtp_patch.py

文件概述

为 Megatron-Core 的 MTP (Multi-Token Prediction) 功能提供 PPO 训练适配补丁。核心问题是:mcore 默认的 MTP 后处理直接计算 loss,但 PPO 需要的是 logits(用于后续计算 log_probs)。

关键代码讲解

patch_postprocess() / unpatch_postprocess()

def patch_postprocess(model):
    """替换 GPTModel._postprocess 为 PPO 兼容版本"""
    model._postprocess_backup = model._postprocess
    model._postprocess = _megatron_gptmodel_postprocess.__get__(model, model.__class__)

def unpatch_postprocess(model):
    """恢复原始 _postprocess"""
    model._postprocess = model._postprocess_backup

PPO 兼容的 _postprocess

def _megatron_gptmodel_postprocess(self, hidden_states, labels, ...):
    # 推理路径:委托给原始实现
    if labels is None:
        return self._postprocess_backup(...)

    # 训练路径:MTP 处理 + 返回 logits(而非 loss)
    if self.config.mtp_num_layers:
        hidden_states_list = torch.chunk(hidden_states, 1 + self.config.mtp_num_layers, dim=0)
        hidden_states = hidden_states_list[0]  # 主 hidden_states

        for mtp_layer_number in range(self.config.mtp_num_layers):
            # 对每个 MTP 层计算 loss
            mtp_labels = roll_tensor(mtp_labels, shifts=-1, ...)
            mtp_loss = self.compute_output_layer_and_language_model_loss(
                hidden_states_list[mtp_layer_number + 1], labels=mtp_labels, ...)
            # 通过 MTPLossAutoScaler 将 MTP loss 附加到主梯度
            hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss_scale * mtp_loss)

    # 关键区别:返回 logits 而非 loss
    logits, _ = self.output_layer(hidden_states, weight=output_weight)
    return logits.transpose(0, 1).contiguous()  # [s,b,h] -> [b,s,h]

核心区别:标准 mcore 在最后直接计算 cross-entropy loss 并返回。但 PPO 需要 logits 来计算 log_probs、entropy 等量,因此这个补丁改为返回 logits。

patch_mtp_layer_get_embeddings()

def patch_mtp_layer_get_embeddings(model):
    """修改 MTP 层的 _get_embeddings 方法,detach embedding"""
    for layer in target_layers:
        layer._get_embeddings = _patched_get_embeddings_for_detach.__get__(layer, ...)

def _patched_get_embeddings_for_detach(self, input_ids, position_ids, embedding, hidden_states, ...):
    # ... 正常计算 embedding
    decoder_input = embedding(input_ids=input_ids, position_ids=position_ids)
    # 关键:detach decoder_input 和 hidden_states
    decoder_input = decoder_input.detach()
    hidden_states = hidden_states.detach()
    return input_ids, position_ids, decoder_input, hidden_states

为什么 detach:在 PPO 训练中,MTP 层的梯度不应该回传到主 decoder,因此需要 detach 切断梯度流。

核心函数列表

函数名 作用
patch_postprocess() 替换 _postprocess 为 PPO 版本
unpatch_postprocess() 恢复原始 _postprocess
_megatron_gptmodel_postprocess() PPO 兼容的后处理(返回 logits)
patch_mtp_layer_get_embeddings() Detach MTP 层的 embedding
unpatch_mtp_layer_get_embeddings() 恢复 MTP 层的 embedding

与其他模块的关系

  • 被训练流程中的 actor 模型前后处理调用
  • 使用 mcore 的 MTPLossAutoScaler、roll_tensor 等工具
  • 与 model_forward.py 配合使用

小结

MTP 补丁解决了 mcore MTP 功能与 PPO 训练的不兼容问题。两个核心修改:1) 后处理返回 logits 而非 loss;2) MTP 层的 embedding detach 防止梯度泄漏。