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 和更新策略。所有并行后端都遵循这个接口。