跳转至

fully_async_main.py — 这是全异步 PPO 训练的主入口文件

文件路径: verl/experimental/fully_async_policy/fully_async_main.py

文件概述

这是全异步 PPO 训练的主入口文件。它定义了 FullyAsyncTaskRunner,负责初始化和协调 Rollouter(样本生成器)和 Trainer(训练器)两个独立组件,通过 MessageQueue 进行异步通信。

关键代码讲解

1. 资源池管理

def create_resource_pool_manager(config, roles):
    resource_pool_spec = {}
    mapping = {}

    # 训练池:Actor、Critic、RefPolicy 等共享
    trainer_pool = [config.trainer.n_gpus_per_node] * config.trainer.nnodes
    resource_pool_spec["trainer_pool"] = trainer_pool

    # Rollout 池:独立 GPU
    rollout_pool = [config.rollout.n_gpus_per_node] * config.rollout.nnodes
    resource_pool_spec["rollout_pool"] = rollout_pool

    return ResourcePoolManager(resource_pool_spec=resource_pool_spec, mapping=mapping)

全异步模式下,训练 GPU 和推理 GPU 是完全分离的。

2. FullyAsyncTaskRunner - 核心编排器

@ray.remote(num_cpus=1)
class FullyAsyncTaskRunner:
    def run(self, config):
        self._initialize_components(config)
        self._run_training_loop()

初始化流程:

def _initialize_components(self, config):
    # 1. 加载 tokenizer 和 processor
    tokenizer = hf_tokenizer(local_path, ...)
    processor = hf_processor(local_path, ...)

    # 2. 创建 Trainer(训练器,使用训练 GPU)
    trainer = FullyAsyncTrainer.remote(...)
    ray.get(trainer.init_workers.remote())

    # 3. 创建 Rollouter(推理器,使用推理 GPU)
    rollouter = FullyAsyncRollouter.remote(...)
    ray.get(rollouter.init_workers.remote())

    # 4. 创建 MessageQueue(连接 Rollouter 和 Trainer)
    message_queue = MessageQueue.remote(config, max_queue_size)
    message_queue_client = MessageQueueClient(message_queue)

    # 5. 设置参数同步器
    param_synchronizer = ParameterSynchronizer.remote(
        config=config, trainer=trainer, rollouter=rollouter, mq=message_queue_client
    )

    # 6. 初始同步:加载 checkpoint + 同步权重
    param_version = ray.get(trainer.load_checkpoint.remote())
    ray.get(param_synchronizer.sync_weights.remote(version=param_version, ...))

3. 训练主循环

def _run_training_loop(self):
    # 并行启动 Rollouter 和 Trainer
    rollouter_future = self.components["rollouter"].fit.remote()
    trainer_future = self.components["trainer"].fit.remote()

    # 使用 ray.wait 监控两个组件
    while futures:
        done_futures, remaining_futures = ray.wait(futures, num_returns=1, timeout=None)
        for future in done_futures:
            ray.get(future)  # 检查是否有异常
        futures = remaining_futures

4. Hydra 入口

@hydra.main(config_path="config", config_name="fully_async_ppo_trainer", version_base=None)
def main(config):
    config.actor_rollout_ref.rollout.nnodes = config.rollout.nnodes
    config.actor_rollout_ref.rollout.n_gpus_per_node = config.rollout.n_gpus_per_node
    run_ppo(config, task_runner_class=FullyAsyncTaskRunner)

全异步架构图

┌─────────────────────────────────────────────────────────────────┐
│                     FullyAsyncTaskRunner                        │
│                                                                 │
│  ┌──────────────────┐    MessageQueue    ┌──────────────────┐  │
│  │                  │                     │                  │  │
│  │  FullyAsync      │ ──── put_sample ──► │  FullyAsync      │  │
│  │  Rollouter       │                     │  Trainer         │  │
│  │  (推理 GPU)      │ ◄─── sync_weights ──│  (训练 GPU)      │  │
│  │                  │                     │                  │  │
│  └──────────────────┘                     └──────────────────┘  │
│           ▲                                        │            │
│           │          ParameterSynchronizer          │            │
│           └────────────────────────────────────────┘            │
└─────────────────────────────────────────────────────────────────┘

核心类/函数列表

名称 类型 说明
create_resource_pool_manager() 函数 创建 GPU 资源池
create_role_worker_mapping() 函数 创建角色到 Worker 类的映射
FullyAsyncTaskRunner Ray Remote 类 全异步训练编排器
main() 函数 Hydra 命令行入口

与其他模块的关系

  • 创建并管理 FullyAsyncRollouter(fully_async_rollouter.py)
  • 创建并管理 FullyAsyncTrainer(fully_async_trainer.py)
  • 创建并管理 MessageQueue(message_queue.py)
  • 创建并管理 ParameterSynchronizer(param_sync.py)

小结

fully_async_main.py 是全异步训练的总指挥。它将训练过程拆分为两个独立的异步组件(Rollouter 和 Trainer),通过 MessageQueue 通信、ParameterSynchronizer 同步参数。这种架构使得推理和训练可以在不同的 GPU 集群上并行运行,大幅提高训练吞吐量。