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 机制支持部分回滚的断点续传。