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 对称。