跳转至

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 训练的"总指挥"。它不直接实现训练逻辑,而是:

  1. 负责配置解析(通过 Hydra)
  2. 负责环境初始化(Ray 集群、设备检测)
  3. 负责组件选择(根据 strategy 选择 Worker 实现)
  4. 负责资源分配(GPU 资源池管理)
  5. 最终将控制权交给 RayPPOTrainer.fit() 执行实际训练

理解这个文件可以帮你建立对整个训练系统的全局认知。