跳转至

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 策略梯度更新)               |
  +----------------------------------+

推荐阅读顺序

  1. config/config.py - 了解基础配置结构
  2. config/algorithm.py - 了解算法配置(KL惩罚、优势估计等)
  3. ppo/utils.py - 了解 Role 枚举和辅助函数
  4. constants_ppo.py - 了解运行时环境变量
  5. main_ppo.py - PPO 训练主入口(重点)
  6. ppo/ray_trainer.py - Ray 分布式训练器(最核心,重点)
  7. ppo/core_algos.py - PPO 核心算法实现(重点)
  8. ppo/reward.py - 奖励函数管理
  9. ppo/metric_utils.py - 训练指标工具
  10. ppo/rollout_corr_helper.py - Rollout 校正
  11. ppo/prefix_grouper_utils.py - 前缀共享优化
  12. sft_trainer.py - SFT 训练器
  13. sft_trainer_ray.py - SFT Ray 分布式训练器
  14. main_eval.py - 离线评估
  15. 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 共享同一进程