跳转至

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 提供了异步和同步两套接口,适配不同调用场景。