跳转至

nixl_checkpoint_engine.py — 这是基于 NIXL(NVIDIA Inference Xfer Library) 的权重同步...

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

文件概述

这是基于 NIXL(NVIDIA Inference Xfer Library) 的权重同步后端实现。与 NCCL 的"一对多广播"不同,NIXL 使用点对点(P2P)RDMA 通信,支持多种底层传输协议(UCX、UCCL、Mooncake 等)。

核心通信模式是流水线传递(Pipeline):Trainer 把数据发给第一个 Rollout Worker,第一个 Worker 再转发给下一个,依次传递,形成链式拓扑。

关键依赖

  • nixl._api / nixl._bindings:NIXL 库的 Python 绑定
  • cupy:CUDA 数组库
  • zmq.asyncio:ZeroMQ 异步版本(注意与 NCCL 版本使用同步 ZMQ 不同)

关键代码讲解

1. NixlAgent -- NIXL Agent 封装

class NixlAgent:
    """NIXL Agent 的封装,增加了 ZeroMQ 消息通信功能"""

    def __init__(self):
        self.agent_name = str(uuid.uuid4())        # 每个 agent 的唯一标识
        self.agent = nixl_api.nixl_agent(self.agent_name)  # 底层 NIXL agent
        self.notifications = defaultdict(deque)     # 存储 NIXL 通知

        self.start_zmq_server()                     # 启动 ZMQ PULL 服务
        self.zmq_clients = {}                       # 到其他 agent 的 ZMQ PUSH 连接
        self.messages = defaultdict(deque)           # 存储接收到的消息

NixlAgent 是对底层 NIXL agent 的封装,主要增加了两个功能: 1. ZeroMQ 消息通道:用于传输 bucket 元数据(名字、形状等) 2. 通知管理:NIXL 的 RDMA 读取完成后通过 notification 机制通知对方

def add_remote_agent(self, metadata):
    """添加一个远程 agent,建立 RDMA 和 ZMQ 连接"""
    agent_name = self.agent.add_remote_agent(metadata.agent_metadata)  # RDMA 连接
    socket = context.socket(zmq.PUSH)
    socket.connect(f"tcp://{metadata.zmq_ip}:{metadata.zmq_port}")    # ZMQ 连接
    self.zmq_clients[agent_name] = socket
    return agent_name

def send_message(self, agent_name, message):
    """通过 ZMQ 发送消息到指定 agent"""
    socket = self.zmq_clients[agent_name]
    socket.send_pyobj((self.agent_name, message), zmq.DONTWAIT)

async def read_message(self, agent_name):
    """异步等待接收来自指定 agent 的消息"""
    while len(self.messages[agent_name]) == 0:
        recv_agent_name, message = await self.socket.recv_pyobj()
        self.messages[recv_agent_name].append(message)
    return self.messages[agent_name].popleft()

ZMQ 使用 PUSH/PULL 模式(点对点),与 NCCL 版本的 PUB/SUB 不同。

2. ReadableOperation -- "可读"操作(发送端)

class ReadableOperation:
    """封装一次"让对方来读我的数据"的操作"""

    def __init__(self, agent, remote_agent, local_descs, metadata):
        self.notify_key = uuid.uuid4().bytes
        message = {
            "notify_key": self.notify_key,
            "remote_descs": local_descs,   # 告诉对方"我的数据在这里"
            **metadata
        }
        self.agent.send_message(self.remote_agent, message)

    async def wait_for_complete(self):
        """等待对方读取完成(通过 NIXL notification)"""
        notification = await self.agent.get_notification(self.remote_agent)
        assert self.notify_key == notification

NIXL 的通信模式是单边读取(One-sided Read): - 发送端不主动 push 数据,而是告诉接收端"我的数据在哪里" - 接收端通过 RDMA READ 直接从发送端的内存中读取数据 - 读取完成后通过 notification 通知发送端

3. ReadOperation -- 读取操作(接收端)

class ReadOperation:
    """封装一次从远程 agent 读取数据的操作"""

    async def read_metadata(self):
        """接收远程 agent 发来的元数据"""
        metadata = await self.agent.read_message(self.remote_agent)
        self.remote_descs = metadata.pop("remote_descs")  # 远程内存描述符
        self.notify_key = metadata.pop("notify_key")       # 完成通知 key
        return metadata

    def begin_read(self):
        """发起 RDMA READ 操作"""
        self.xfer_handle = self.agent.initialize_xfer(
            "READ", self.local_descs, self.remote_descs,
            self.remote_agent, self.notify_key
        )
        state = self.agent.transfer(self.xfer_handle)  # 开始传输
        assert state != "ERR"

    async def wait_for_complete(self):
        """等待 RDMA READ 完成"""
        while True:
            state = self.agent.check_xfer_state(self.xfer_handle)
            if state == "DONE":
                break
            await asyncio.sleep(0)  # 让出控制权
        self.agent.release_xfer_handle(self.xfer_handle)

4. NIXLCheckpointEngine -- 主体类

通信拓扑:链式传递

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

每个进程只知道自己的前一个和后一个邻居:

Trainer[0] --> Rollout[0] --> Rollout[1] --> ... --> Rollout[N-1]
 (rank 0)      (rank 1)      (rank 2)              (rank N)

prepare -- 分配并注册 RDMA 内存

def prepare(self):
    if self.device == "cuda":
        send_buf = cp.zeros(self.bucket_size, dtype=cp.uint8)
        recv_buf = cp.zeros(self.bucket_size, dtype=cp.uint8)
        self.send_buf = torch.as_tensor(send_buf, dtype=torch.uint8)
        self.recv_buf = torch.as_tensor(recv_buf, dtype=torch.uint8)
    else:
        self.send_buf = torch.zeros(..., pin_memory=True)  # CPU 用 pin memory
        self.recv_buf = torch.zeros(..., pin_memory=True)

    # 注册内存用于 RDMA
    self.send_reg_descs = self.agent.register_memory(self.send_buf)
    self.recv_reg_descs = self.agent.register_memory(self.recv_buf)

    # 获取传输描述符
    self.send_descs = self.agent.get_xfer_descs(self.send_buf)
    self.recv_descs = self.agent.get_xfer_descs(self.recv_buf)

    return self.agent.get_agent_metadata()

RDMA 通信的关键步骤:先分配内存,然后注册到 NIXL(让远程节点可以直接读取),最后获取传输描述符。

send_weights -- 发送权重

async def send_weights(self, weights):
    ...
    for name, weight in weights:
        if offset + weight.nbytes > self.bucket_size:
            # bucket 满了,等待上一个操作完成
            if readable_op is not None:
                await readable_op.wait_for_complete()

            # 通知下一个 agent "数据准备好了,来读取吧"
            readable_op = ReadableOperation(
                self.agent, self.next_agent, send_descs,
                {"bucket_meta": bucket_meta, "is_last": False},
            )
            send_buf, recv_buf = recv_buf, send_buf  # 双缓冲
            ...

        # 填充 bucket
        send_buf[offset:offset+weight.nbytes].copy_(
            weight.view(-1).view(torch.uint8), non_blocking=True
        )
        offset += weight.nbytes

receive_weights -- 接收并转发权重

async def receive_weights(self):
    # 从前一个 agent 读取第一个 bucket
    read_op = ReadOperation(self.agent, self.prev_agent, recv_descs, self.bucket_size)
    metadata = await read_op.read_metadata()
    read_op.begin_read()
    await read_op.wait_for_complete()

    while not metadata["is_last"]:
        # 1. 转发给下一个 agent(如果有的话)
        if self.next_agent is not None:
            readable_op = ReadableOperation(
                self.agent, self.next_agent, send_descs, metadata,
            )

        # 2. 同时从前一个 agent 读取下一个 bucket
        read_op = ReadOperation(self.agent, self.prev_agent, recv_descs, ...)
        next_metadata = await read_op.read_metadata()
        read_op.begin_read()

        # 3. 提取张量并 yield
        for name, meta in metadata["bucket_meta"].items():
            tensor = send_buf[meta["offset"]:...].view(...)
            yield name, tensor

        # 4. 等待读取和转发完成
        if readable_op is not None:
            await readable_op.wait_for_complete()
        await read_op.wait_for_complete()

        # 5. 双缓冲交换
        metadata = next_metadata
        send_buf, recv_buf = recv_buf, send_buf

流水线传递的核心:每个 Worker 同时做三件事(流水线重叠): 1. 从前一个 agent 读取下一个 bucket 2. 将当前 bucket 转发给下一个 agent 3. 从当前 bucket 提取张量供推理引擎更新

finalize -- 清理所有资源

def finalize(self):
    if self.prev_agent:
        self.agent.remove_remote_agent(self.prev_agent)
    if self.next_agent:
        self.agent.remove_remote_agent(self.next_agent)

    self.agent.deregister_memory(self.send_reg_descs)  # 注销 RDMA 内存
    self.agent.deregister_memory(self.recv_reg_descs)
    ...

核心类/函数列表

名称 类型 作用
NixlAgentMetadata dataclass agent 的 RDMA + ZMQ 元数据
NixlAgent 类 NIXL agent 封装,增加 ZMQ 消息通道
ReadableOperation 类 发送端:"我的数据准备好了"
ReadOperation 类 接收端:RDMA READ 读取远程数据
NIXLCheckpointEngine 类 NIXL 后端的完整实现

通信拓扑与数据流

  ┌──────────────┐    ┌──────────────┐    ┌──────────────┐    ┌──────────────┐
  │  Trainer[0]  │ -> │  Rollout[0]  │ -> │  Rollout[1]  │ -> │  Rollout[2]  │
  │   (rank 0)   │    │   (rank 1)   │    │   (rank 2)   │    │   (rank 3)   │
  └──────────────┘    └──────────────┘    └──────────────┘    └──────────────┘
        |                   |                   |                   |
        |  ReadableOp       |  ReadOp           |  ReadOp           |
        |  "数据在这里"  --> |  RDMA READ    --> |  RDMA READ    --> |
        |                   |  + 转发            |  + 转发            |
        v                   v                   v                   v
      发送方             接收+转发            接收+转发            接收方
                        (流水线)             (流水线)

与 NCCL 版本的对比

方面 NCCL NIXL
通信模式 集合通信(broadcast) 点对点 RDMA(链式传递)
拓扑结构 星形(1 对多) 链式(1 对 1 传递)
元数据传输 ZMQ PUB/SUB ZMQ PUSH/PULL
数据传输 NCCL broadcast NIXL RDMA READ
传输方向 发送端 push 接收端 pull(单边读取)
通知机制 无(同步阻塞) NIXL notification

与其他模块的关系

  • base.py:继承 CheckpointEngine,通过 @CheckpointEngineRegistry.register("nixl") 注册。
  • nixl._api / nixl._bindings:NIXL 库,提供 RDMA 传输能力。

小结

NIXL 后端的设计与 NCCL 有本质区别:它使用链式 P2P RDMA 而非集合通信。这种设计的优势是不需要所有节点同时参与通信,更适合异构网络环境和弹性扩缩容场景。代码中 ReadableOperation("数据可读"通知)和 ReadOperation(RDMA 读取)的配合体现了单边读取的设计范式。