跳转至

nccl_checkpoint_engine.py — 这是基于 **NCCL(NVIDIA Collective Communications Libra...

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

文件概述

这是基于 NCCL(NVIDIA Collective Communications Library) 的权重同步后端实现。NCCL 是 NVIDIA 提供的高性能 GPU 集合通信库,支持 broadcast、all-reduce 等操作。

本文件实现了通过 NCCL broadcast 操作将 Trainer 的模型权重广播给所有 Rollout Worker 的功能,是最典型、使用最广泛的后端实现。

关键依赖

  • cupy:CUDA 数组库,用于在 master 进程分配 GPU 内存(避免 PyTorch expandable_segments 的兼容问题)
  • ray.util.collective:Ray 对 NCCL 集合通信的封装
  • zmq(ZeroMQ):轻量级消息队列,用于传输张量元数据

关键代码讲解

1. MasterMetadata -- 主节点元数据

@dataclass
class MasterMetadata:
    zmq_ip: str
    zmq_port: int

Trainer 的 rank 0 进程作为 master,需要把自己的 ZeroMQ 服务地址告知所有 Rollout Worker。

2. BroadcastOperation -- 异步广播操作

class BroadcastOperation:
    def __init__(self, rank, group_name, bucket, metadata, socket, topic):
        self.rank = rank
        self.group_name = group_name
        self.bucket = bucket
        self.metadata = metadata
        self.socket = socket
        self.topic = topic

        loop = asyncio.get_running_loop()
        self._task = loop.run_in_executor(None, self._run)

    def _run(self):
        # 1. 通过 ZeroMQ 广播张量元数据(名字、形状、类型、偏移)
        if self.rank == 0:
            self.socket.send_string(self.topic, flags=zmq.SNDMORE)
            self.socket.send_pyobj(self.metadata)
        else:
            self.socket.recv_string()
            self.metadata = self.socket.recv_pyobj()

        # 2. 通过 NCCL 广播张量数据
        collective.broadcast(self.bucket, src_rank=0, group_name=self.group_name)

    async def wait_for_complete(self):
        await self._task
        return self.metadata

关键设计点: - 双通道通信:元数据走 ZeroMQ(CPU 侧,支持任意 Python 对象),张量数据走 NCCL(GPU 侧,高带宽) - 异步执行:使用 run_in_executor 把阻塞的 NCCL 操作放到线程池中,不阻塞事件循环 - ZMQ 使用 PUB/SUB 模式:master 发布,所有 worker 订阅

3. NCCLCheckpointEngine -- 主体类

初始化

@CheckpointEngineRegistry.register("nccl")
class NCCLCheckpointEngine(CheckpointEngine):
    def __init__(self, bucket_size, group_name="default", rebuild_group=False,
                 is_master=False, rollout_dtype=torch.bfloat16):
        self.bucket_size = bucket_size      # bucket 大小(字节)
        self.group_name = group_name        # NCCL 进程组名
        self.rebuild_group = rebuild_group  # 是否每次重建进程组
        self.rollout_dtype = rollout_dtype  # 推理端权重精度
        self.is_master = is_master          # 是否为 master
        if self.is_master:
            self._start_zmq_server()        # master 启动 ZMQ 服务

prepare -- 分配双缓冲

def prepare(self):
    if self.is_master:
        # master 使用 cupy 分配,避免 PyTorch expandable_segments 兼容问题
        self.send_buf = cp.zeros(self.bucket_size, dtype=cp.uint8)
        self.recv_buf = cp.zeros(self.bucket_size, dtype=cp.uint8)
    else:
        # worker 使用 PyTorch 分配
        self.send_buf = torch.zeros(self.bucket_size, dtype=torch.uint8, device="cuda")
        self.recv_buf = torch.zeros(self.bucket_size, dtype=torch.uint8, device="cuda")
    return MasterMetadata(zmq_ip=self.ip, zmq_port=self.listen_port) if self.is_master else None

每个进程分配两块大小为 bucket_size 的 GPU 内存(send_buf 和 recv_buf),用于双缓冲交替传输。

build_topology -- 构建通信拓扑

@classmethod
def build_topology(cls, trainer_world_size, rollout_world_size, metadata):
    trainer_kwargs = {
        "rank": [0] + [-1] * (trainer_world_size - 1),  # 只有 rank 0 参与
        "world_size": [rollout_world_size + 1] * trainer_world_size,
        "master_metadata": [metadata[0]] * trainer_world_size,
    }
    rollout_kwargs = {
        "rank": list(range(1, rollout_world_size + 1)),  # rollout 从 rank 1 开始
        "world_size": [rollout_world_size + 1] * rollout_world_size,
        "master_metadata": [metadata[0]] * rollout_world_size,
    }
    return trainer_kwargs, rollout_kwargs

通信拓扑是星形的: - Trainer rank 0 是 NCCL 的 rank 0(广播源) - Trainer 的其他 rank 设为 -1(不参与通信) - 所有 Rollout Worker 依次编号为 rank 1, 2, 3, ...

send_weights -- 发送权重(Trainer 端)

@torch.no_grad()
async def send_weights(self, weights):
    # Trainer 的非 rank 0 进程:消费权重但不发送
    if self.rank < 0:
        for name, weight in weights:
            pass
        return

    send_buf, recv_buf = self.send_buf, self.recv_buf
    broadcast_op = None
    bucket_meta = {}
    offset = 0

    for name, weight in weights:
        # 当 bucket 装满时,发送当前 bucket
        if offset + weight.nbytes > self.bucket_size:
            torch.cuda.synchronize()

            # 等待上一个 broadcast 完成
            if broadcast_op is not None:
                await broadcast_op.wait_for_complete()

            # 启动新的 broadcast
            broadcast_op = BroadcastOperation(
                rank=self.rank, group_name=self.group_name,
                bucket=send_buf,
                metadata={"bucket_meta": bucket_meta, "is_last": False},
                socket=self.socket, topic=self.topic,
            )

            # 交换发送和接收缓冲区(双缓冲)
            send_buf, recv_buf = recv_buf, send_buf
            bucket_meta = {}
            offset = 0

        # 记录元数据并拷贝数据到 bucket
        bucket_meta[name] = {
            "name": name, "shape": weight.shape,
            "dtype": weight.dtype, "offset": offset,
        }
        send_buf[offset : offset + weight.nbytes] = cp.asarray(
            weight.view(-1).view(torch.uint8)
        )
        offset += weight.nbytes

    # 发送最后一个 bucket(标记 is_last=True)
    broadcast_op = BroadcastOperation(
        ..., metadata={"bucket_meta": bucket_meta, "is_last": True}, ...
    )
    await broadcast_op.wait_for_complete()

发送流程的核心思想: 1. 把权重张量逐个填入 bucket(连续内存) 2. bucket 满了就触发一次 NCCL broadcast 3. 使用双缓冲:一边发送上一个 bucket,一边填充下一个 bucket 4. 最后一个 bucket 用 is_last=True 标记结束

receive_weights -- 接收权重(Rollout 端)

@torch.no_grad()
async def receive_weights(self):
    send_buf, recv_buf = self.send_buf, self.recv_buf

    # 接收第一个 bucket
    broadcast_op = BroadcastOperation(
        rank=self.rank, group_name=self.group_name,
        bucket=recv_buf, metadata=None,
        socket=self.socket, topic=self.topic,
    )
    metadata = await broadcast_op.wait_for_complete()
    send_buf, recv_buf = recv_buf, send_buf

    while not metadata["is_last"]:
        # 1. 开始接收下一个 bucket(异步)
        broadcast_op = BroadcastOperation(
            ..., bucket=recv_buf, metadata=None, ...
        )

        # 2. 从已接收的 bucket 中提取张量并 yield
        for name, meta in metadata["bucket_meta"].items():
            dtype, shape = meta["dtype"], meta["shape"]
            size = dtype.itemsize * shape.numel()
            tensor = send_buf[meta["offset"]:meta["offset"]+size] \
                .view(dtype=dtype).view(shape)
            yield name, tensor

        # 3. 等待下一个 bucket 接收完成
        metadata = await broadcast_op.wait_for_complete()
        torch.cuda.synchronize()
        send_buf, recv_buf = recv_buf, send_buf

    # 提取最后一个 bucket 的张量
    for name, meta in metadata["bucket_meta"].items():
        ...
        yield name, tensor

接收流程的核心思想: 1. 通过 NCCL broadcast 接收 bucket 数据 2. 根据元数据从 bucket 中还原出各个张量 3. 双缓冲流水线:接收下一个 bucket 的同时处理当前 bucket 的数据

4. ZeroMQ 服务

def _start_zmq_server(self):
    self.ip = ray.util.get_node_ip_address().strip("[]")
    self.listen_port, _ = get_free_port(self.ip)
    context = zmq.Context()
    self.socket = context.socket(zmq.PUB)  # 发布者模式
    self.socket.bind(f"tcp://{self.ip}:{self.listen_port}")

def _connect_zmq_client(self, metadata):
    context = zmq.Context()
    self.socket = context.socket(zmq.SUB)  # 订阅者模式
    self.socket.connect(f"tcp://{metadata.zmq_ip}:{metadata.zmq_port}")
    self.socket.setsockopt_string(zmq.SUBSCRIBE, self.topic)

ZMQ 使用 PUB/SUB 模式:master 是 Publisher,所有 worker 是 Subscriber。

核心类/函数列表

名称 类型 作用
MasterMetadata dataclass master 的 ZMQ 地址信息
BroadcastOperation 类 封装一次异步 NCCL broadcast
NCCLCheckpointEngine 类 NCCL 后端的完整实现

数据流图

send_weights 流程(Trainer rank 0):

  weights 生成器
       |
       v
  ┌─────────────────────────────────┐
  │  填入 send_buf(bucket 打包)     │
  │  记录 bucket_meta               │
  └─────────┬───────────────────────┘
            | bucket 满或最后一个
            v
  ┌─────────────────────────────────┐
  │  ZMQ PUB: 发送 bucket_meta      │  (CPU)
  │  NCCL broadcast: 发送 bucket    │  (GPU)
  │  双缓冲交换 send_buf / recv_buf  │
  └─────────────────────────────────┘


receive_weights 流程(Rollout worker):

  ┌─────────────────────────────────┐
  │  ZMQ SUB: 接收 bucket_meta      │  (CPU)
  │  NCCL broadcast: 接收 bucket    │  (GPU)
  └─────────┬───────────────────────┘
            |
            v
  ┌─────────────────────────────────┐
  │  根据 meta 从 buffer 还原张量     │
  │  yield (name, tensor)           │
  │  双缓冲交换                      │
  └─────────────────────────────────┘

与其他模块的关系

  • base.py:继承 CheckpointEngine,通过 @CheckpointEngineRegistry.register("nccl") 注册。
  • ray.util.collective:使用 Ray 封装的 NCCL 集合通信。
  • verl.utils.net_utils:获取空闲端口、检查 IPv6。

小结

NCCL 后端是最典型的实现,其核心思想是: 1. 用 bucket 打包减少通信次数 2. 用 双缓冲实现通信与计算的流水线 3. 用 ZMQ+NCCL 双通道分别传输元数据和张量数据 4. 使用 asyncio + run_in_executor 实现非阻塞异步通信