跳转至

fsdp_workers.py — FSDP 后端 PPO Workers

文件概述

基于 PyTorch FSDP(Fully Sharded Data Parallel)后端的 PPO 训练 Worker 实现。这是 verl 最早期的 Worker 实现,将 Actor、Critic、Rollout、Reference 四个角色组装到分布式工作进程中。文件约 1800+ 行,是整个框架中最大的单文件之一。

注意:这个文件是较早的实现,新代码推荐使用 engine_workers.py 中基于统一引擎的实现。

核心类

ActorRolloutRefWorker

最重要的 Worker 类,将 Actor + Rollout + Reference 三个角色合并在同一个工作进程中。

class ActorRolloutRefWorker:
    def __init__(self, config, ...):
        # 根据 role 配置初始化不同的角色组合
        # role 可以是: actor, rollout, ref, actor_rollout, actor_rollout_ref 等
        self.role = config.role

关键方法: - init_model(): 初始化模型、优化器、FSDP 包装 - update_actor(): 执行 PPO 策略更新 - generate_sequences(): 调用推理引擎生成 response - compute_ref_log_prob(): 计算参考模型的 log_prob - compute_log_prob(): 计算当前策略的 log_prob - save_checkpoint() / load_checkpoint(): 检查点管理

CriticWorker

Critic 价值网络的 Worker 实现。

class CriticWorker:
    def __init__(self, config, ...):
        self.critic = DataParallelPPOCritic(...)

    def compute_values(self, data):
        # 计算每个 token 位置的价值估计
        ...

    def update_critic(self, data):
        # 用 value loss 更新 Critic 网络
        ...

核心流程

Hybrid Engine 权重同步流程

# 1. 推理引擎释放显存
await rollout.sleep()

# 2. Actor 训练完成后,获取最新权重
weights = actor.get_per_tensor_param()

# 3. 将权重同步到推理引擎
await rollout.update_weights(weights)

# 4. 推理引擎恢复显存占用
await rollout.wake_up()

与其他模块的关系

  • 依赖 actor/dp_actor.py 中的 DataParallelPPOActor
  • 依赖 critic/dp_critic.py 中的 DataParallelPPOCritic
  • 依赖 rollout/ 中的各推理引擎
  • 依赖 sharding_manager/ 处理数据分片
  • 被 verl/trainer/ppo/ray_trainer.py 通过 Ray 调度

小结

fsdp_workers.py 是 FSDP 后端下 PPO 训练的完整实现,将所有角色和流程组装在一起。随着框架演化,建议关注更模块化的 engine_workers.py 实现。