message_queue.py — 实现了基于 Ray 的异步消息队列¶
文件路径: verl/experimental/fully_async_policy/message_queue.py
文件概述¶
实现了基于 Ray 的异步消息队列,是 Rollouter 和 Trainer 之间的通信桥梁。MessageQueue 作为 Ray Actor 运行,提供线程安全的生产者-消费者模式。
关键代码讲解¶
1. MessageQueue - 服务端¶
@ray.remote(num_cpus=2, max_concurrency=20)
class MessageQueue:
def __init__(self, config, max_queue_size=1000):
self.queue = deque(maxlen=self.max_queue_size) # 有限长度双端队列
self.current_param_version = 0
self.val_queue = deque() # 验证结果队列
# asyncio 同步原语
self._lock = asyncio.Lock()
self._consumer_condition = asyncio.Condition(self._lock)
# 统计信息
self.total_produced = 0
self.total_consumed = 0
self.dropped_samples = 0
2. 生产者接口 - put_sample¶
async def put_sample(self, sample, param_version):
async with self._lock:
# 队列满时丢弃最老的样本
if len(self.queue) >= self.max_queue_size:
self.queue.popleft()
self.dropped_samples += 1
self.queue.append(sample)
self.total_produced += 1
self._consumer_condition.notify_all() # 唤醒等待的消费者
return True
3. 消费者接口 - get_sample¶
async def get_sample(self):
async with self._lock:
# 队列为空时等待
while len(self.queue) == 0 and self.running:
await self._consumer_condition.wait()
if not self.running and len(self.queue) == 0:
return None # 终止信号
data = self.queue.popleft()
self.total_consumed += 1
return data, len(self.queue) # 返回数据和剩余队列长度
4. MessageQueueClient - 客户端¶
封装了与 Ray Actor 的异步/同步通信:
class MessageQueueClient:
def __init__(self, queue_actor):
self.queue_actor = queue_actor
async def put_sample(self, sample, param_version):
"""异步放入样本"""
future = self.queue_actor.put_sample.remote(sample, param_version)
return await asyncio.wrap_future(future.future())
async def get_sample(self):
"""异步获取样本"""
future = self.queue_actor.get_sample.remote()
return await asyncio.wrap_future(future.future())
def get_sample_sync(self):
"""同步获取样本(Trainer 使用)"""
return ray.get(self.queue_actor.get_sample.remote())
5. 辅助功能¶
# 更新参数版本
async def update_param_version(self, version):
# 获取统计信息
async def get_statistics(self):
return {
"queue_size": len(self.queue),
"total_produced": self.total_produced,
"total_consumed": self.total_consumed,
"dropped_samples": self.dropped_samples,
"current_param_version": self.current_param_version,
}
# 验证结果队列(put_validate / get_validate)
async def put_validate(self, data): ...
async def get_validate(self): ...
# 内存使用估算
async def get_memory_usage(self): ...
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
MessageQueue |
Ray Remote Actor | 消息队列服务端 |
MessageQueueClient |
类 | 消息队列客户端(异步+同步接口) |
与其他模块的关系¶
- 被
FullyAsyncRollouter用作生产者(放入 rollout 样本) - 被
FullyAsyncTrainer用作消费者(取出样本训练) - 被
ParameterSynchronizer用于更新参数版本 - 在
fully_async_main.py中创建并分发给各组件
小结¶
MessageQueue 是全异步训练的通信中枢。它通过 asyncio 的 Lock+Condition 实现线程安全的生产者-消费者模式,支持队列满时的自动丢弃、优雅关闭、统计信息收集等功能。MessageQueueClient 提供了异步和同步两套接口,适配不同调用场景。