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 集群上并行运行,大幅提高训练吞吐量。