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 读取)的配合体现了单边读取的设计范式。