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 -- 主节点元数据¶
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 实现非阻塞异步通信