跳转至

dp_critic.py — FSDP 数据并行 Critic

文件概述

基于 FSDP 的 DataParallelPPOCritic 实现。结构与 dp_actor.py 类似,但输出的是标量价值(value)而非概率。

核心类

DataParallelPPOCritic

class DataParallelPPOCritic(BasePPOCritic):
    def compute_values(self, data):
        """前向传播,输出每个 token 位置的价值估计"""
        output = self.critic_module(input_ids, attention_mask, position_ids)
        # Critic 模型通常是 ForTokenClassification 架构
        # 输出 shape: (batch, seq_len, 1) → squeeze 到 (batch, seq_len)
        values = output.logits.squeeze(-1)
        return values

    def update_critic(self, data):
        """计算 value loss 并更新 Critic"""
        # Clipped value loss(与 PPO 的 clip 思想一致)
        vpreds_clipped = values + (vpreds - values).clamp(-clip_range, clip_range)
        vf_loss1 = (vpreds - returns) ** 2
        vf_loss2 = (vpreds_clipped - returns) ** 2
        vf_loss = 0.5 * torch.max(vf_loss1, vf_loss2).mean()

价值损失计算

Critic 的目标是准确预测每个 token 位置的未来累积奖励(returns):

\[ L_{\text{value}} = \mathbb{E}\left[\left(V_{\text{pred}} - \text{returns}\right)^2\right] \]

其中: - \(V_{\text{pred}}\): Critic 模型预测的价值 - returns: 由 GAE 计算得到的实际回报 - 使用 clip 防止 Critic 更新幅度过大

与其他模块的关系

  • 继承自 base.py 的 BasePPOCritic
  • 被 fsdp_workers.py 的 CriticWorker 使用
  • 损失计算逻辑对应 utils/losses.py 中的 value_loss

小结

DataParallelPPOCritic 实现了 FSDP 后端的 Critic 训练,核心是价值预测和 clipped value loss 的计算。