跳转至

model_forward_fused.py — 融合 kernel 前向函数

文件路径

verl/models/mcore/model_forward_fused.py

文件概述

为 Megatron-Core 模型提供使用融合 kernel 的前向函数。与标准前向的区别在于:不输出 logits,而是直接使用 linear_cross_entropy 融合计算 log_probs 和 entropy,避免实例化巨大的 logits 矩阵。

关键代码讲解

GPTModel 前向替换

def patch_fused_forward(model):
    """将 GPTModel 的 forward 替换为融合版本"""
    model.forward_backup = model.forward
    model.forward = _fused_GPTModel_forward.__get__(model, model.__class__)

def unpatch_fused_forward(model):
    """恢复原始 forward"""
    model.forward = model.forward_backup

融合前向函数

def _fused_GPTModel_forward(self, input_ids, position_ids, attention_mask, ...):
    """在 GPTModel 内部直接计算 log_probs 和 entropy"""
    # 正常走 decoder 获取 hidden_states
    hidden_states = self.decoder(...)

    # 不走 output_layer,直接用融合 kernel
    log_probs, entropy = linear_cross_entropy(
        hidden_states, self.output_layer.weight, labels, temperature, "none")

    return CausalLMOutputForPPO(log_probs=log_probs, entropy=entropy)

融合前向生成器

def fused_forward_model_gen(vision_model=False):
    def fused_forward_model(model, input_ids, position_ids, attention_mask,
                            labels, labels_mask, temperature, multi_modal_inputs):
        # 1. patch forward
        patch_fused_forward(model)
        # 2. 预处理 + 前向
        output = model(input_ids=input_ids_rmpad, labels=labels_rmpad,
                       temperature=temperature, ...)
        # 3. 后处理
        ret = postprocess_packed_seqs_for_dict_output(labels_mask, output, ...)
        # 4. unpatch forward
        unpatch_fused_forward(model)
        return ret

核心函数列表

函数名 作用
patch_fused_forward() 替换 GPTModel forward 为融合版本
unpatch_fused_forward() 恢复原始 forward
_fused_GPTModel_forward() 融合 kernel 的 GPTModel forward
fused_forward_model_gen() 融合前向函数工厂
fused_forward_no_padding_gen() 无 padding 融合前向函数工厂

与其他模块的关系

  • 使用 verl.utils.kernel.linear_cross_entropy 的 Triton kernel
  • 被 registry.py 注册为融合前向函数
  • 使用 util.py 的打包/解包函数

小结

融合前向函数是 PPO 训练的性能关键。通过在 mcore 流水线内部直接计算 log_probs 和 entropy,避免了实例化 (batch, seq_len, vocab_size) 大小的 logits 张量,显著降低显存使用。