跳转至

base.py — Actor 抽象基类

文件概述

定义 PPO Actor(策略网络)的抽象基类 BasePPOActor,规定了所有 Actor 实现必须提供的接口。

核心类

BasePPOActor(ABC)

class BasePPOActor(ABC):
    @abstractmethod
    def compute_log_prob(self, data: DataProto) -> DataProto:
        """计算给定 prompt+response 下每个 token 的对数概率

        输入: prompt_ids, response_ids, attention_mask
        输出: log_prob (shape: [batch, response_len])
        """
        pass

    @abstractmethod
    def update_policy(self, data: DataProto) -> dict:
        """执行 PPO 策略梯度更新

        输入: old_log_prob, advantages, returns, ...
        输出: 训练指标 (loss, grad_norm, lr, ...)
        """
        pass

与其他模块的关系

  • 被 dp_actor.py(FSDP)和 megatron_actor.py(Megatron)继承实现
  • 被 fsdp_workers.py 和 megatron_workers.py 使用

小结

BasePPOActor 定义了 Actor 的两个核心操作:计算 log_prob 和更新策略。所有并行后端都遵循这个接口。