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 的计算。