跳转至

mooncake_checkpoint_engine.py — 这是基于 Mooncake TransferEngine 的权重同步后端实现

文件路径: verl/checkpoint_engine/mooncake_checkpoint_engine.py

文件概述

这是基于 Mooncake TransferEngine 的权重同步后端实现。Mooncake 是一个高性能数据传输引擎,支持 RDMA 和 Ascend NPU 的 direct 传输。

与 NIXL 类似,Mooncake 也采用点对点传输,但它使用了更简洁的设计: 1. 通信协调使用 vLLM 的 StatelessProcessGroup(而非 ZMQ) 2. 传输完成的通知使用"魔数写回"机制(而非 NIXL 的 notification) 3. 数据链式转发:Trainer -> Rollout[0] -> Rollout[1] -> ...

关键依赖

  • mooncake.engine.TransferEngine:Mooncake 传输引擎
  • vllm.distributed.utils.StatelessProcessGroup:进程间通信协调

关键代码讲解

1. 初始化 -- TransferEngine

@CheckpointEngineRegistry.register("mooncake")
class MooncakeCheckpointEngine(CheckpointEngine):
    def __init__(self, bucket_size, device="cuda", rollout_dtype=torch.bfloat16,
                 device_name="", is_master=False, rebuild_group=False):
        self.bucket_size = bucket_size
        self.device = device

        # 初始化 TransferEngine
        self.engine = TransferEngine()
        hostname = ray.util.get_node_ip_address().strip("[]")
        ret = self.engine.initialize(
            hostname,
            "P2PHANDSHAKE",
            "ascend_direct" if self.device == "npu" else "rdma",  # 支持 NPU 和 GPU
            device_name,
        )
        assert ret == 0

        rpc_port = self.engine.get_rpc_port()
        self.session_id = f"{hostname}:{rpc_port}"

        # 分配双缓冲 + 魔数缓冲
        self.buf = torch.empty(2 * self.bucket_size, dtype=torch.uint8, device=self.device)
        self.magic_buf = torch.empty(4 * 1024, dtype=torch.uint8, device=self.device)

        # 注册内存(用于 RDMA)
        ret = self.engine.batch_register_memory(
            [self.buf.data_ptr(), self.magic_buf.data_ptr()],
            [2 * self.bucket_size, 4 * 1024],
        )

关键点: - 双缓冲在一块连续内存中分成两半:buf[:bucket_size] 和 buf[bucket_size:] - magic_buf 用于"魔数写回"通知机制 - 支持 GPU(rdma)和 NPU(ascend_direct)两种传输模式

2. build_topology -- 星形 + 链式拓扑

@classmethod
def build_topology(cls, trainer_world_size, rollout_world_size, metadatas):
    trainer_kwargs = {
        "rank": [0] + [-1] * (trainer_world_size - 1),
        "world_size": [rollout_world_size + 1] * trainer_world_size,
        "metadata": [metadatas[0]] * trainer_world_size,
    }
    rollout_kwargs = {
        "rank": list(range(1, rollout_world_size + 1)),
        "world_size": [rollout_world_size + 1] * rollout_world_size,
        "metadata": [metadatas[0]] * rollout_world_size,
    }
    return trainer_kwargs, rollout_kwargs

3. init_process_group -- 使用 StatelessProcessGroup

def init_process_group(self, rank, world_size, metadata):
    self.rank = rank
    self.world_size = world_size
    if rank < 0:
        return

    # 创建无状态进程组
    self.store = StatelessProcessGroup.create(
        host=metadata["addr"], port=metadata["port"],
        rank=rank, world_size=world_size,
    )

    # 所有进程交换 buffer 信息
    info = {"session_id": self.session_id, "ptr": self.buf.data_ptr()}
    info_list = self.store.all_gather_obj(info)

    # 每个 Rollout Worker 记住前一个 Worker 的 buffer 信息
    self.buffer_info = None if rank == 0 else info_list[rank - 1]

通过 all_gather_obj 让所有进程交换各自的 session_id 和 buffer 地址,然后每个 Worker 记住前一个 Worker 的信息,形成链式拓扑。

4. send_weights -- 发送权重

async def send_weights(self, weights):
    if self.rank < 0:
        for name, weight in weights:
            pass
        return

    bufs = [self.buf[:self.bucket_size], self.buf[self.bucket_size:]]
    idx = 0
    current = bufs[idx]
    should_wait = False

    for name, weight in weights:
        weight = weight.to(self.rollout_dtype)

        if offset + weight.nbytes > self.bucket_size:
            get_torch_device().synchronize()
            info = {
                "bucket_meta": bucket_meta,
                "ptr": current.data_ptr(),  # 告诉接收端数据的 GPU 地址
                "len": offset,
                "is_last": False,
            }
            self.store.send_obj(info, 1)  # 发送元数据给 rank 1

            # 切换到另一个缓冲区
            idx ^= 1
            current = bufs[idx]
            bucket_meta = {}
            offset = 0

            # 等待上一个缓冲区被对方读走
            if should_wait:
                await self.wait_for_complete(current)
            should_wait = True

        current[offset:offset+weight.nbytes].copy_(
            weight.view(-1).view(torch.uint8), non_blocking=True
        )
        offset += weight.nbytes

    # 发送最后一个 bucket
    info = {"bucket_meta": bucket_meta, "ptr": current.data_ptr(),
            "len": offset, "is_last": True}
    self.store.send_obj(info, 1)
    await self.wait_for_complete(current)  # 等待对方确认

5. wait_for_complete -- "魔数写回"通知机制

async def wait_for_complete(self, buf):
    magic = torch.tensor([0xAB, 0xDC, 0xEF, 0x88], dtype=torch.uint8, device=self.device)
    while True:
        if torch.equal(buf[:4], magic):
            break
        await asyncio.sleep(0)

这是一个简单但巧妙的完成通知机制: 1. 接收端读取完数据后,通过 RDMA 将 4 字节魔数 [0xAB, 0xDC, 0xEF, 0x88] 写入发送端的 buffer 头部 2. 发送端轮询 buffer 头 4 字节,看到魔数就知道数据已被读走,可以复用 buffer

6. receive_weights -- 接收、转发、通知

async def receive_weights(self):
    bufs = [self.buf[:self.bucket_size], self.buf[self.bucket_size:]]
    idx = 0
    current = bufs[idx]
    # 预设魔数到 magic_buf
    self.magic_buf[:4] = torch.tensor([0xAB, 0xDC, 0xEF, 0x88], ...)

    while True:
        # 1. 接收元数据
        info = self.store.recv_obj(self.rank - 1)

        # 等待当前 buffer 被下游读走(第 3 轮开始)
        if idx >= 2 and self.rank < self.world_size - 1:
            await self.wait_for_complete(current)

        # 2. 通过 RDMA 从前一个 Worker 读取数据
        ret = self.engine.transfer_sync_read(
            self.buffer_info["session_id"],
            current.data_ptr(),      # 本地目标地址
            info["ptr"],             # 远程源地址
            info["len"],             # 数据长度
        )

        # 3. 转发元数据给下一个 Worker(更新 ptr 为自己的地址)
        info["ptr"] = current.data_ptr()
        if self.rank < self.world_size - 1:
            self.store.send_obj(info, self.rank + 1)

        # 4. 从 buffer 提取张量并 yield
        for name, meta in info["bucket_meta"].items():
            tensor = current[meta["offset"]:...].view(...)
            yield name, tensor

        # 5. 写回魔数,通知前一个 Worker "我读完了"
        ret = self.engine.transfer_sync_write(
            self.buffer_info["session_id"],
            self.magic_buf.data_ptr(),  # 本地源(魔数)
            info["ptr"],                # 远程目标(前一个 Worker 的 buffer 头部)
            4,                          # 4 字节
        )

        # 6. 切换缓冲区
        idx += 1
        current = bufs[idx % 2]
        get_torch_device().synchronize()

        if info["is_last"]:
            break

核心类/函数列表

名称 类型 作用
MooncakeCheckpointEngine 类 Mooncake 后端的完整实现
wait_for_complete 方法 魔数轮询等待完成通知

数据流图

  Trainer[0]                Rollout[0]              Rollout[1]
  (rank 0)                  (rank 1)                (rank 2)
     |                         |                       |
     | send_obj(meta)          |                       |
     | ----------------------> |                       |
     |                         | RDMA READ             |
     |                         | <--- 读取数据          |
     |                         |                       |
     |                         | send_obj(meta)        |
     |                         | --------------------> |
     |                         |                       | RDMA READ
     |                         |                       | <--- 读取数据
     |                         |                       |
     |                         | yield 张量              | yield 张量
     |                         |                       |
     |   魔数写回 (确认)         |                       |
     | <---------------------- |   魔数写回 (确认)       |
     |                         | <-------------------- |

与 NIXL 版本的对比

方面 NIXL Mooncake
传输库 NIXL (nixl_api) Mooncake TransferEngine
元数据传输 ZMQ PUSH/PULL StatelessProcessGroup
完成通知 NIXL notification 魔数写回
异步机制 NIXL 异步 xfer 同步 transfer_sync_read
设备支持 GPU GPU + NPU
内存管理 cupy + NIXL register torch + batch_register_memory

与其他模块的关系

  • base.py:继承 CheckpointEngine,注册为 "mooncake"。
  • mooncake.engine:Mooncake TransferEngine 库。
  • vllm.distributed.utils:使用 StatelessProcessGroup 协调通信。
  • verl.utils.device:设备工具函数。

小结

Mooncake 后端的设计在概念上与 NIXL 类似(都是链式 P2P 传递),但实现更为简洁: 1. 使用 StatelessProcessGroup 代替 ZMQ 传输元数据 2. 使用"魔数写回"代替专门的 notification 机制 3. 使用同步 transfer_sync_read 简化异步控制 4. 天然支持 GPU 和 NPU 两种硬件

"魔数写回"是一个值得注意的设计:通过向远程内存写入特定字节序列来通知对方操作完成,避免了额外的通知通道。