跳转至

ray_trainer.py — 实现了一步离策略 PPO 训练器

文件路径: verl/experimental/one_step_off_policy/ray_trainer.py

文件概述

实现了一步离策略 PPO 训练器。核心创新是将推理和训练流水线化:当训练当前 batch 时,同时开始生成下一个 batch 的序列。通过 asyncio.create_task 实现非阻塞的异步生成。

关键代码讲解

1. 类定义

class OneStepOffRayTrainer(SeparateRayPPOTrainer):
    def __init__(self, config, tokenizer, role_worker_mapping, resource_pool_manager, ...):
        # 与标准 RayPPOTrainer 类似的初始化
        self.hybrid_engine = config.actor_rollout_ref.hybrid_engine
        assert not self.hybrid_engine  # 不支持混合引擎

        # 跳过 Rollout worker 映射(由 AgentLoopManager 管理)
        role_worker_mapping.pop(Role.Rollout, None)

2. 异步初始化

def _init_async_rollout_manager(self):
    from verl.experimental.one_step_off_policy.agent_loop import OneStepOffAgentLoopManager
    self.async_rollout_mode = True
    self.async_rollout_manager = OneStepOffAgentLoopManager.create(
        config=self.config, reward_loop_worker_handles=reward_loop_worker_handles
    )

3. 训练主循环 - 流水线化

async def fit(self):
    continuous_iterator = self._create_continuous_iterator()

    # 启动第一个异步生成任务
    batch_data_future = asyncio.create_task(self._async_gen_next_batch(continuous_iterator))

    while batch_data_future is not None:
        batch_data_future = await self.fit_step(batch_data_future, continuous_iterator)
        if self.is_last_step:
            return

4. 核心:异步生成下一批

async def _async_gen_next_batch(self, continuous_iterator):
    epoch, batch_dict = next(continuous_iterator)
    batch = DataProto.from_single_dict(batch_dict)
    gen_batch = self._get_gen_batch(batch)

    # 异步生成序列
    with marked_timer("generate_async", timing_raw):
        gen_batch_output = await self.async_rollout_manager.generate_sequences_async(gen_batch_output)

    batch = batch.union(gen_batch_output)
    return metrics, timing_raw, epoch, batch, future_reward

5. 单步训练 - 训练+生成并行

async def fit_step(self, batch_data_future, continuous_iterator):
    with marked_timer("step", self.timing_raw):
        # 等待当前 batch 的生成完成
        batch, batch_data_future = await self._fit_generate(batch_data_future, continuous_iterator)

        # PPO 训练步骤(与标准训练相同)
        await asyncio.sleep(0)  # 让事件循环处理其他协程
        batch = self._fit_compute_reward(batch)
        await asyncio.sleep(0)
        batch = self._fit_compute_log_prob(batch)
        # ... 更多步骤 ...
        batch = self._fit_update_actor(batch)
        await asyncio.sleep(0)
        self._fit_update_weights()  # 同步权重到推理服务器

6. 流水线关键 - _fit_generate

async def _fit_generate(self, batch_data_future, continuous_iterator):
    with marked_timer("gen", timing_raw):
        # 等待当前推理完成
        _metrics, _timing_raw, epoch, batch, future_reward = await batch_data_future

    # 同步权重
    with marked_timer("sync_rollout_weights", timing_raw):
        self._fit_update_weights()
        await self.async_rollout_manager.clear_kv_cache()

    # 立即启动下一个 batch 的异步生成
    if not self.is_last_step:
        batch_data_future = asyncio.create_task(
            self._async_gen_next_batch(continuous_iterator)
        )
        await asyncio.sleep(0)  # 让生成任务开始执行
    else:
        batch_data_future = None

    return batch, batch_data_future

流水线时序图

时间 ──────────────────────────────────────────────────────►

步骤1:  [推理 batch1]────────────►
                                 [训练 batch1]────────────►
                                 [推理 batch2]───────────► (重叠!)
                                                          [训练 batch2]────►
                                                          [推理 batch3]────►

        ├── gen ──────────────────┤
                                 ├── sync ─┤
                                           ├── train ─────┤
                                           ├── gen(async) ──────────────────┤

核心类/函数列表

名称 类型 说明
OneStepOffRayTrainer 类 一步离策略 PPO 训练器
fit() 异步方法 训练主循环
fit_step() 异步方法 单步训练
_async_gen_next_batch() 异步方法 异步生成下一批数据
_fit_generate() 异步方法 等待推理完成并启动下一轮推理

与其他模块的关系

  • 继承自 SeparateRayPPOTrainer(separation/ray_trainer.py)
  • 使用 OneStepOffAgentLoopManager(agent_loop/)管理推理
  • 被 OneStepTaskRunner(main_ppo.py)创建和运行

小结

OneStepOffRayTrainer 通过 asyncio 实现了训练和推理的流水线重叠。关键设计是在训练当前 batch 开始时就启动下一个 batch 的推理任务(asyncio.create_task),中间穿插 await asyncio.sleep(0) 让推理协程有机会运行。这种方式比标准同步训练快(因为推理和训练重叠),但比全异步训练简单(不需要 MessageQueue 和参数版本管理)。