跳转至

fully_async_rollouter.py — 实现了全异步样本生成器(Rollouter)

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

文件概述

实现了全异步样本生成器(Rollouter),负责持续不断地从数据集中取样、调用 LLM 生成响应、并将生成的样本放入 MessageQueue 供 Trainer 消费。这是全异步训练架构中的"生产者"角色。

核心特性: - 流式处理:样本一生成完毕就立即放入队列,无需等待整个 batch - 并发控制:限制同时进行的 rollout 数量 - 暂停/恢复:支持参数同步时暂停生成、同步后恢复 - 部分回滚(Partial Rollout):支持在参数更新时中断未完成的 rollout

关键代码讲解

1. 类定义

@ray.remote(num_cpus=10, max_concurrency=100)
class FullyAsyncRollouter(SeparateRayPPOTrainer):
    def __init__(self, config, tokenizer, ...):
        # 关键约束
        assert not self.hybrid_engine  # 不能用混合引擎
        assert self.config.data.train_batch_size == 0  # batch_size 必须为 0
        assert self.config.data.gen_batch_size == 1    # 逐样本生成

        # 异步队列
        self.pending_queue = asyncio.Queue(maxsize=128)   # 待处理队列
        self.active_tasks = set()                          # 活跃的异步任务
        self.cancel_queue = asyncio.Queue()                # 被取消的任务(需要重新提交)

2. 核心异步架构

fit() 方法启动两个并行的协程:

async def fit(self):
    generation_task = asyncio.create_task(self._streaming_generation_main())
    monitor_task = asyncio.create_task(self._async_monitor_loop())
    await asyncio.gather(generation_task, monitor_task, return_exceptions=True)

_streaming_generation_main() 进一步启动两个子协程:

async def _streaming_generation_main(self):
    # 数据喂入协程:从 dataloader 读数据放入 pending_queue
    self.feed_task = asyncio.create_task(self._feed_samples())
    # 处理协程:从 pending_queue 取样本,提交生成任务
    self.processor_task = asyncio.create_task(self._processor_worker())

3. 数据喂入 - _feed_samples

async def _feed_samples(self):
    continuous_iterator = self._create_continuous_iterator()
    for epoch, batch_dict in continuous_iterator:
        full_batch = prepare_single_generation_data(batch_dict, self.config)
        rollout_sample = RolloutSample(full_batch=full_batch, ...)
        await self.pending_queue.put(rollout_sample)

        if self.global_steps >= self.total_rollout_steps:
            break
        self.global_steps += 1

    await self.pending_queue.put("DONE")  # 结束信号

4. 流式处理 - _processor_worker

async def _processor_worker(self):
    while True:
        # 检查是否需要暂停
        if self.paused or await self._should_pause_generation():
            # 等待所有活跃任务完成
            while self.active_tasks:
                done_tasks, self.active_tasks = await asyncio.wait(
                    self.active_tasks, return_when=asyncio.FIRST_COMPLETED
                )
            # 暂停等待恢复信号
            async with self.lock:
                while self.paused:
                    await self.condition.wait()
            continue

        # 优先从 cancel_queue 取(被中断的任务重新提交)
        if not self.cancel_queue.empty():
            rollout_sample = await self.cancel_queue.get()
        else:
            rollout_sample = await self.pending_queue.get()

        # 控制并发数
        while len(self.active_tasks) >= self.max_concurrent_samples:
            done_tasks, self.active_tasks = await asyncio.wait(
                self.active_tasks, return_when=asyncio.FIRST_COMPLETED
            )

        # 提交单样本处理
        task = asyncio.create_task(self._process_single_sample_streaming(rollout_sample))
        self.active_tasks.add(task)

5. 单样本处理

async def _process_single_sample_streaming(self, rollout_sample):
    ret, is_cancel = await self.async_rollout_manager.generate_single_sample_async(
        rollout_sample.full_batch, rollout_sample.agent_loop_output_list
    )

    if not is_cancel:
        # 生成成功:放入 MessageQueue
        rollout_sample.full_batch = ret
        await self.message_queue_client.put_sample(
            sample=ray.cloudpickle.dumps(rollout_sample),
            param_version=rollout_sample.param_version,
        )
    else:
        # 被取消:放入 cancel_queue 等待重新提交
        rollout_sample.agent_loop_output_list = ret
        await self.cancel_queue.put(rollout_sample)

6. 暂停与恢复

参数同步时需要暂停 rollout:

async def pause(self):
    async with self.lock:
        self.paused = True
        if self.config.async_training.partial_rollout:
            await self.async_rollout_manager.cancel()  # 取消进行中的 rollout
        if self.active_tasks:
            await asyncio.gather(*self.active_tasks, return_exceptions=True)
            self.active_tasks.clear()
        await self.async_rollout_manager.clear_kv_cache()  # 释放显存

async def resume(self, dependency_ref=None):
    async with self.lock:
        if self.config.async_training.partial_rollout:
            await self.async_rollout_manager.resume()
        self.paused = False
        self.condition.notify_all()

流程图

                    ┌─────────────┐
                    │ DataLoader   │
                    └──────┬──────┘
                           │ _feed_samples()
                           ▼
                    ┌─────────────┐
                    │pending_queue│
                    └──────┬──────┘
                           │ _processor_worker()
                    ┌──────┴──────┐
                    │             │
               ┌────▼────┐  ┌────▼────┐
               │ Task 1  │  │ Task 2  │ ... (并发执行)
               └────┬────┘  └────┬────┘
                    │             │
              ┌─────▼─────┐ ┌────▼──────┐
              │ 成功完成   │ │ 被取消    │
              └─────┬─────┘ └────┬──────┘
                    │             │
           ┌────────▼────┐  ┌────▼────────┐
           │MessageQueue │  │cancel_queue  │
           └─────────────┘  └─────────────┘
                  │                  │
                  ▼                  └──► 重新提交
           Trainer 消费

核心类/函数列表

名称 类型 说明
FullyAsyncRollouter Ray Remote 类 全异步样本生成器
fit() 方法 主入口,启动生成循环
_feed_samples() 协程 数据喂入协程
_processor_worker() 协程 流式处理协程
_process_single_sample_streaming() 协程 单样本生成
pause() / resume() 方法 暂停/恢复生成
_should_pause_generation() 方法 判断是否需要暂停

与其他模块的关系

  • 继承自 SeparateRayPPOTrainer(separation/ray_trainer.py)
  • 使用 MessageQueueClient(message_queue.py)发送样本
  • 使用 FullyAsyncAgentLoopManager(agent_loop/)管理推理
  • 被 ParameterSynchronizer(param_sync.py)调用 pause/resume
  • 使用 detach_utils.py 中的 prepare_single_generation_data 和 RolloutSample

小结

FullyAsyncRollouter 是全异步训练的"生产者"。它通过 asyncio 协程实现了高效的流式样本生成:数据喂入、并发推理、结果入队三个环节解耦。暂停/恢复机制确保参数同步时的正确性,cancel_queue 机制支持部分回滚的断点续传。