verl Trainer 模块总览¶
模块简介¶
verl/trainer/ 是 verl 框架的核心训练入口模块。它实现了基于 Ray 分布式框架的 PPO (Proximal Policy Optimization) 强化学习训练流程,同时也支持 SFT (Supervised Fine-Tuning) 监督微调训练。
本模块是整个 RLHF(基于人类反馈的强化学习)训练流程的"大脑",负责协调 Actor(策略模型)、Critic(价值模型)、Reward Model(奖励模型)和 Reference Policy(参考策略)等多个组件。
目录结构¶
verl/trainer/
├── __init__.py # 模块初始化(仅包含许可声明)
├── constants_ppo.py # PPO 运行时环境变量常量定义
├── main_ppo.py # PPO 训练主入口(最核心)
├── main_eval.py # 离线评估入口
├── main_generation_server.py # 独立生成服务器入口
├── sft_trainer.py # SFT 单机训练器(SPMD 模式)
├── sft_trainer_ray.py # SFT 分布式训练器(Ray 模式)
├── config/ # 配置数据类
│ ├── __init__.py # 导出所有配置类
│ ├── algorithm.py # 算法配置(KL、优势估计、Rollout Correction)
│ └── config.py # 基础配置(Checkpoint、Profile、Model)
└── ppo/ # PPO 核心算法实现
├── __init__.py # 模块初始化
├── core_algos.py # 核心算法:优势估计、策略损失函数(最重要)
├── ray_trainer.py # Ray 分布式 PPO 训练器(最核心)
├── metric_utils.py # 训练指标计算工具
├── reward.py # 奖励函数加载与管理
├── utils.py # 工具函数:Role 枚举、辅助判断
├── prefix_grouper_utils.py # 前缀共享优化工具
└── rollout_corr_helper.py # Rollout 校正助手(处理 off-policy 问题)
PPO 训练流程图¶
用户启动训练
|
v
+-------------------+
| main_ppo.py |
| main() 入口 |
+-------------------+
|
1. 初始化 Ray 集群
2. 创建 TaskRunner
|
v
+-------------------+
| TaskRunner.run |
| 配置 Worker |
+-------------------+
|
1. 加载模型/tokenizer
2. 创建数据集
3. 初始化 RayPPOTrainer
|
v
+----------------------------------------------+
| RayPPOTrainer (ray_trainer.py) |
| |
| init_workers() |
| | |
| +-- 创建 Actor/Rollout Worker Group |
| +-- 创建 Critic Worker Group |
| +-- 创建 Reference Policy Worker Group |
| +-- 创建 Reward Loop Manager |
| +-- 创建 Checkpoint Manager |
| |
| fit() -- 主训练循环 |
| | |
| v |
| +-----------------------------------------+ |
| | 每个 training step: | |
| | | |
| | 1. 从 DataLoader 获取 batch | |
| | | | |
| | v | |
| | 2. generate_sequences() | |
| | (Actor 生成回复) | |
| | | | |
| | v | |
| | 3. compute_reward() | |
| | (Reward Model 计算奖励) | |
| | | | |
| | v | |
| | 4. compute_old_log_prob() | |
| | (Actor 计算旧策略概率) | |
| | | | |
| | v | |
| | 5. compute_ref_log_prob() | |
| | (Reference 计算参考概率) | |
| | | | |
| | v | |
| | 6. compute_values() | |
| | (Critic 计算价值函数) | |
| | | | |
| | v | |
| | 7. compute_advantage() | |
| | (计算优势函数 - 在驱动进程上) | |
| | (core_algos.py) | |
| | | | |
| | v | |
| | 8. update_critic() | |
| | (更新 Critic 网络) | |
| | | | |
| | v | |
| | 9. update_actor() | |
| | (更新 Actor 策略网络) | |
| | | | |
| | v | |
| | 10. update_weights() | |
| | (同步权重到 Rollout) | |
| | | | |
| | v | |
| | 11. 保存 checkpoint / 验证 | |
| +-----------------------------------------+ |
+----------------------------------------------+
数据流图¶
Prompt Dataset
|
v
+----------+ +-----------+ +-------------+
| Actor |----->| Rollout |----->| Reward |
| (策略网络)| | (生成回复) | | (计算奖励) |
+----------+ +-----------+ +-------------+
| |
v v
+----------+ +-----------+ +-------------+
| Ref Policy| | Critic |<-----| token_level |
| (参考策略) | | (价值网络) | | _scores |
+----------+ +-----------+ +-------------+
| |
v v
+----------------------------------+
| compute_advantage() |
| (GAE / GRPO / REINFORCE++ 等) |
+----------------------------------+
|
v
+----------------------------------+
| update_actor() + update_critic()|
| (PPO 策略梯度更新) |
+----------------------------------+
推荐阅读顺序¶
- config/config.py - 了解基础配置结构
- config/algorithm.py - 了解算法配置(KL惩罚、优势估计等)
- ppo/utils.py - 了解 Role 枚举和辅助函数
- constants_ppo.py - 了解运行时环境变量
- main_ppo.py - PPO 训练主入口(重点)
- ppo/ray_trainer.py - Ray 分布式训练器(最核心,重点)
- ppo/core_algos.py - PPO 核心算法实现(重点)
- ppo/reward.py - 奖励函数管理
- ppo/metric_utils.py - 训练指标工具
- ppo/rollout_corr_helper.py - Rollout 校正
- ppo/prefix_grouper_utils.py - 前缀共享优化
- sft_trainer.py - SFT 训练器
- sft_trainer_ray.py - SFT Ray 分布式训练器
- main_eval.py - 离线评估
- main_generation_server.py - 生成服务器
关键概念速查¶
| 概念 | 说明 |
|---|---|
| Actor | 策略模型,负责生成文本回复 |
| Critic | 价值网络,估计状态的价值(用于 GAE) |
| Ref Policy | 参考策略,用于计算 KL 散度惩罚 |
| Reward Model | 奖励模型,对生成的回复打分 |
| GAE | Generalized Advantage Estimation,广义优势估计 |
| GRPO | Group Relative Policy Optimization,分组相对策略优化 |
| PPO | Proximal Policy Optimization,近端策略优化 |
| KL Penalty | KL 散度惩罚,防止策略偏离参考策略太远 |
| Rollout | 策略生成回复的过程 |
| Hybrid Engine | 混合引擎,Actor 和 Rollout 共享同一进程 |