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 两种硬件
"魔数写回"是一个值得注意的设计:通过向远程内存写入特定字节序列来通知对方操作完成,避免了额外的通知通道。