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 实现。