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 参数用于控制生成策略的"温度"。