main_ppo.py — 这是 PPO 训练的主入口文件¶
文件概述¶
模块路径: verl.trainer.main_ppo
这是 PPO 训练的主入口文件,是整个 verl 强化学习训练流程的起点。当你运行 python -m verl.trainer.main_ppo 时,就是从这里开始的。
该文件负责:
1. 使用 Hydra 解析配置文件
2. 初始化 Ray 分布式集群
3. 创建并配置各种 Worker(Actor、Critic、Reference Policy 等)
4. 构建数据集和采样器
5. 初始化 RayPPOTrainer 并启动训练
在训练流程中的位置¶
这是最顶层的入口,是用户运行训练的起始点。它像一个"导演",把各个组件组织到一起。
关键代码讲解¶
1. Hydra 主入口¶
@hydra.main(config_path="config", config_name="ppo_trainer", version_base=None)
def main(config):
"""Main entry point for PPO training with Hydra configuration management."""
auto_set_device(config) # 自动检测设备(CUDA / NPU)
config = migrate_legacy_reward_impl(config) # 兼容旧版奖励配置
run_ppo(config)
Hydra 是一个配置管理框架,@hydra.main 装饰器会自动从 config/ppo_trainer.yaml 加载配置,并支持命令行覆盖参数。
2. Ray 集群初始化¶
def run_ppo(config, task_runner_class=None) -> None:
if not ray.is_initialized():
default_runtime_env = get_ppo_ray_runtime_env()
ray_init_kwargs = config.ray_kwargs.get("ray_init", {})
runtime_env_kwargs = ray_init_kwargs.get("runtime_env", {})
if config.transfer_queue.enable:
runtime_env_vars = runtime_env_kwargs.get("env_vars", {})
runtime_env_vars["TRANSFER_QUEUE_ENABLE"] = "1"
runtime_env = OmegaConf.merge(default_runtime_env, runtime_env_kwargs)
ray.init(**OmegaConf.to_container(ray_init_kwargs))
这段代码将预定义的环境变量(来自 constants_ppo.py)与用户自定义的环境变量合并,然后初始化 Ray 集群。
3. TaskRunner 远程执行¶
if task_runner_class is None:
task_runner_class = ray.remote(num_cpus=1)(TaskRunner)
runner = task_runner_class.remote()
ray.get(runner.run.remote(config))
TaskRunner 被包装成 Ray remote actor,在非 Head 节点上运行(num_cpus=1 确保不占用主节点资源)。ray.get() 会阻塞直到训练完成。
4. TaskRunner 类 - Worker 配置核心¶
class TaskRunner:
def __init__(self):
self.role_worker_mapping = {} # 角色 -> Worker类 的映射
self.mapping = {} # 角色 -> 资源池 的映射
TaskRunner 是训练配置的核心,它负责根据配置选择合适的 Worker 实现。
添加 Actor/Rollout Worker¶
def add_actor_rollout_worker(self, config):
use_legacy_worker_impl = config.trainer.get("use_legacy_worker_impl", "auto")
if use_legacy_worker_impl == "disable":
# 新版模型引擎实现
from verl.workers.engine_workers import ActorRolloutRefWorker
actor_rollout_cls = ActorRolloutRefWorker
elif config.actor_rollout_ref.actor.strategy in {"fsdp", "fsdp2"}:
# FSDP 策略
from verl.workers.fsdp_workers import AsyncActorRolloutRefWorker
actor_rollout_cls = AsyncActorRolloutRefWorker
elif config.actor_rollout_ref.actor.strategy == "megatron":
# Megatron 策略
from verl.workers.megatron_workers import AsyncActorRolloutRefWorker
actor_rollout_cls = AsyncActorRolloutRefWorker
根据训练策略(FSDP、Megatron 等)选择不同的 Worker 实现。这体现了 verl 的模型无关性设计。
添加 Critic Worker¶
def add_critic_worker(self, config):
if config.critic.strategy in {"fsdp", "fsdp2"}:
if use_legacy_worker_impl in ["auto", "enable"]:
from verl.workers.fsdp_workers import CriticWorker
elif use_legacy_worker_impl == "disable":
from verl.workers.engine_workers import TrainingWorker
CriticWorker = TrainingWorker
资源池初始化¶
def init_resource_pool_mgr(self, config):
global_pool_id = "global_pool"
resource_pool_spec = {
global_pool_id: [config.trainer.n_gpus_per_node] * config.trainer.nnodes,
}
# 可选:为 Reward Model 创建独立的资源池
if config.reward.reward_model.enable_resource_pool:
reward_pool = [config.reward.reward_model.n_gpus_per_node] * config.reward.reward_model.nnodes
resource_pool_spec["reward_pool"] = reward_pool
资源池定义了GPU的分配方式。默认所有角色共享一个全局池,但 Reward Model 可以有独立的资源池。
5. TaskRunner.run() - 完整训练配置流程¶
def run(self, config):
# 1. 解析配置
OmegaConf.resolve(config)
# 2. 添加各种 Worker
actor_rollout_cls, ray_worker_group_cls = self.add_actor_rollout_worker(config)
self.add_critic_worker(config)
self.add_reward_model_resource_pool(config)
self.add_ref_policy_worker(config, actor_rollout_cls)
# 3. 下载模型权重
local_path = copy_to_local(config.actor_rollout_ref.model.path, ...)
# 4. 初始化 tokenizer
tokenizer = hf_tokenizer(local_path, ...)
processor = hf_processor(local_path, ...) # 多模态用
# 5. 创建数据集
train_dataset = create_rl_dataset(config.data.train_files, ...)
val_dataset = create_rl_dataset(config.data.val_files, ...)
train_sampler = create_rl_sampler(config.data, train_dataset)
# 6. 初始化并启动训练器
trainer = RayPPOTrainer(
config=config,
tokenizer=tokenizer,
role_worker_mapping=self.role_worker_mapping,
resource_pool_manager=resource_pool_manager,
...
)
trainer.init_workers()
trainer.fit()
6. 数据集和采样器创建¶
def create_rl_dataset(data_paths, data_config, tokenizer, processor, ...):
dataset_cls = get_dataset_class(data_config)
dataset = dataset_cls(
data_files=data_paths,
tokenizer=tokenizer,
processor=processor,
config=data_config,
)
return dataset
def create_rl_sampler(data_config, dataset):
if data_config.sampler is not None and data_config.sampler.get("class_path", None) is not None:
# 自定义课程学习采样器
curriculum_class = load_extern_object(...)
sampler = curriculum_class(data_source=dataset, data_config=data_config)
elif data_config.shuffle:
# 随机采样(支持断点恢复)
sampler = RandomSampler(data_source=dataset, generator=train_dataloader_generator)
else:
# 顺序采样
sampler = SequentialSampler(data_source=dataset)
return sampler
注意这里使用 torchdata.stateful_dataloader.sampler.RandomSampler 而不是 torch.utils.data.RandomSampler,这是为了支持断点恢复。
核心类/函数列表¶
| 名称 | 类型 | 作用 |
|---|---|---|
main(config) |
function | Hydra 入口函数,解析配置并启动训练 |
run_ppo(config) |
function | 初始化 Ray 集群,创建 TaskRunner 并启动训练 |
TaskRunner |
class | 核心配置类,负责选择和注册各种 Worker |
TaskRunner.run(config) |
method | 执行完整的训练配置和启动流程 |
TaskRunner.add_actor_rollout_worker() |
method | 根据策略类型添加 Actor/Rollout Worker |
TaskRunner.add_critic_worker() |
method | 添加 Critic Worker |
TaskRunner.init_resource_pool_mgr() |
method | 初始化 GPU 资源池管理器 |
create_rl_dataset() |
function | 创建 RL 训练/验证数据集 |
create_rl_sampler() |
function | 创建数据采样器(支持课程学习) |
数据流和调用关系¶
main()
|
+-- auto_set_device()
+-- migrate_legacy_reward_impl()
+-- run_ppo()
|
+-- get_ppo_ray_runtime_env() (constants_ppo.py)
+-- ray.init()
+-- TaskRunner.remote()
|
+-- TaskRunner.run()
|
+-- add_actor_rollout_worker()
+-- add_critic_worker()
+-- add_reward_model_resource_pool()
+-- add_ref_policy_worker()
+-- validate_config()
+-- copy_to_local()
+-- hf_tokenizer() / hf_processor()
+-- create_rl_dataset() x2 (train + val)
+-- create_rl_sampler()
+-- RayPPOTrainer(...) (ray_trainer.py)
+-- trainer.init_workers()
+-- trainer.fit()
小结¶
main_ppo.py 是 PPO 训练的"总指挥"。它不直接实现训练逻辑,而是:
- 负责配置解析(通过 Hydra)
- 负责环境初始化(Ray 集群、设备检测)
- 负责组件选择(根据 strategy 选择 Worker 实现)
- 负责资源分配(GPU 资源池管理)
- 最终将控制权交给
RayPPOTrainer.fit()执行实际训练
理解这个文件可以帮你建立对整个训练系统的全局认知。