跳转至

base.py — Critic 抽象基类

文件概述

定义 PPO Critic(价值网络)的抽象基类 BasePPOCritic。

核心类

BasePPOCritic(ABC)

class BasePPOCritic(ABC):
    @abstractmethod
    def compute_values(self, data: DataProto) -> DataProto:
        """计算每个 token 位置的价值估计 \(V(s)\)

        输入: prompt+response 的 token 序列
        输出: values (shape: [batch, response_len])
        """
        pass

    @abstractmethod
    def update_critic(self, data: DataProto) -> dict:
        """用价值损失更新 Critic 网络

        输入: values, returns, response_mask
        输出: 训练指标 (vf_loss, vf_clipfrac, ...)
        """
        pass

与其他模块的关系

  • 被 dp_critic.py 和 megatron_critic.py 继承实现
  • 被 Workers 使用来执行价值估计和 Critic 训练

小结

BasePPOCritic 定义了 Critic 的两个核心操作:计算价值和更新价值网络。结构与 BasePPOActor 对称。