跳转至

dense_common.py — 纯文本模型通用前向

文件路径

verl/models/transformers/dense_common.py

文件概述

为纯文本密集模型(如 LLaMA、Qwen2、Mistral 等)提供 PPO 训练专用的前向传播函数。核心贡献是将标准的 "hidden_states -> logits -> loss" 流程替换为直接输出 log_probs 和 entropy,这是 PPO 算法所需的两个关键量。

关键代码讲解

PPO 专用输出结构

@dataclass
class CausalLMOutputForPPO(CausalLMOutputWithPast):
    log_probs: Optional[torch.FloatTensor] = None
    entropy: Optional[torch.FloatTensor] = None

继承自 HuggingFace 的 CausalLMOutputWithPast,额外添加了 PPO 需要的 log_probs(对数概率)和 entropy(熵)。

基础模型前向

def forward_base_model(self, input_ids, attention_mask, position_ids, ...):
    """通用的 base model 前向,适用于所有纯文本模型"""
    outputs = self.model(
        input_ids=input_ids,
        attention_mask=attention_mask,
        position_ids=position_ids,
        # ...
    )
    return outputs

这个函数直接调用模型的 decoder(如 self.model),获取 hidden states,不做 lm_head 投影。

Triton 后端前向

def forward_with_triton_backend(self, input_ids, ..., temperature=1.0, **loss_kwargs):
    from verl.utils.kernel.linear_cross_entropy import linear_cross_entropy

    outputs = forward_base_model(self, input_ids, ...)
    hidden_states = outputs[0]

    # 将 input_ids 向左移一位作为标签(next token prediction)
    rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)

    # 融合计算:hidden_states * lm_head_weight -> softmax -> log_prob + entropy
    log_probs, entropy = linear_cross_entropy(
        hidden_states, self.lm_head.weight, rolled_labels, temperature, "none"
    )

    return CausalLMOutputForPPO(log_probs=log_probs, entropy=entropy, ...)

关键优化:linear_cross_entropy 将 "线性投影 + softmax + log_prob + entropy" 四步融合为一个 Triton kernel,避免了巨大的 logits 矩阵(vocab_size 维度通常是 32k-128k)的显存开销。

Torch 后端前向

def forward_with_torch_backend(self, input_ids, ..., temperature=1.0, **loss_kwargs):
    from verl.utils.experimental.torch_functional import FusedLinearForPPO

    outputs = forward_base_model(self, input_ids, ...)
    hidden_states = outputs[0]

    rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)

    fused_linear_for_ppo = FusedLinearForPPO()
    log_probs, entropy = fused_linear_for_ppo.forward(
        hidden_states=hidden_states, vocab_weights=self.lm_head.weight,
        input_ids=rolled_labels, temperature=temperature,
    )

与 Triton 后端类似,但使用纯 PyTorch 实现的融合计算,兼容性更好。

核心类/函数列表

名称 作用
CausalLMOutputForPPO PPO 专用输出数据结构
forward_base_model() 通用 base model 前向(只过 decoder)
forward_with_triton_backend() 使用 Triton kernel 的融合前向
forward_with_torch_backend() 使用 PyTorch 的融合前向

与其他模块的关系

  • 被 monkey_patch.py 中的 patch_forward_with_backends() 调用
  • 使用 verl.utils.kernel.linear_cross_entropy 的 Triton kernel
  • 使用 verl.utils.experimental.torch_functional 的 PyTorch 实现

小结

这个文件是 PPO 训练的核心优化之一。通过融合 lm_head 投影和交叉熵计算,避免了实例化完整的 logits 张量,大幅降低了显存占用。temperature 参数用于控制生成策略的"温度"。