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 张量,显著降低显存使用。