kimi_checkpoint_engine.py — 这是基于 Kimi ParameterServer 的权重同步后端实现¶
文件路径:
verl/checkpoint_engine/kimi_checkpoint_engine.py
文件概述¶
这是基于 Kimi ParameterServer 的权重同步后端实现。与 NCCL/NIXL 不同,它使用了参数服务器(Parameter Server)架构,并且 Trainer 端在发送前会将参数先卸载到 CPU,以减少 GPU 显存压力。
关键特性:
1. 参数服务器模式:使用 Kimi 的 ParameterServer 管理参数注册、分发和收集。
2. CPU 卸载:Trainer 端先将参数从 GPU 移到 CPU,再注册到参数服务器。
3. 多 Trainer 分担:不同的 Trainer Worker 负责不同的参数切片(按轮询分配)。
4. 所有 rank 都参与通信:与 NCCL 只有 rank 0 发送不同,所有 Trainer 都参与。
关键依赖¶
checkpoint_engine.distributed:Kimi 的分布式通信库checkpoint_engine.ps:Kimi 的 ParameterServer 库
关键代码讲解¶
1. ckpt_get_named_tensor_buckets -- 参数切片分桶¶
def ckpt_get_named_tensor_buckets(
iterable, bucket_bytes, world_size, rank_id, rollout_dtype=torch.bfloat16
):
current_bucket = {}
current_size = 0
for tensor_idx, (name, tensor) in enumerate(iterable):
tensor = tensor.to(rollout_dtype) # 转换精度
if tensor_idx % world_size == rank_id: # 轮询分配给不同 rank
tensor_size = tensor.element_size() * tensor.numel()
if current_size + tensor_size > bucket_bytes:
if current_bucket:
yield current_bucket # 一个 bucket 满了
current_bucket = {}
current_size = 0
current_bucket[name] = tensor
current_size += tensor_size
if current_bucket:
yield current_bucket
关键设计:tensor_idx % world_size == rank_id 实现了参数在多个 Trainer Worker 之间的轮询分配。比如有 4 个 Trainer Worker:
- Worker 0 负责第 0, 4, 8, ... 个参数
- Worker 1 负责第 1, 5, 9, ... 个参数
- 以此类推
2. receive_tensor -- 自定义接收方法¶
async def receive_tensor(self, checkpoint_name, ranks_group, ranks=None,
bucket_size=2<<30, disable_h2d_buffer=False):
"""自定义的张量接收方法,会被 monkey-patch 到 ParameterServer 上"""
dist.barrier(group=ranks_group)
# 生成 H2D(Host to Device)桶
buckets = _gen_h2d_buckets(
self._current_global_parameter_metas, bucket_size,
self._local_rdma_devices, self._remote_rdma_devices, ranks,
)
# 分配 H2D buffer(CPU -> GPU 拷贝缓冲区)
h2d_buffer = torch.empty(bucket_size, dtype=torch.uint8, device=self.device_manager.device_type)
# 双缓冲
buffer = torch.empty(bucket_size * 2, dtype=torch.uint8, device=self.device_manager.device_type)
for i in range(max_len):
# 1. 从参数服务器拷贝数据到 h2d_buffer
if i < len(receiver_rank_buckets) and not disable_h2d_buffer:
self._copy_to_buffer(checkpoint_name, receiver_rank_buckets[i][1],
h2d_buffer, ...)
# 2. 对每个 receiver 执行 broadcast
for receiver_rank, _buckets in buckets_by_receiver_rank.items():
...
broadcast_op = BroadcastOperation(
rank=receiver_rank, ranks_group=ranks_group,
bucket=buffer_b, metadata=bucket.items,
)
...
# 3. 从 buffer 中还原张量并 yield
for item in meta_list:
tensor = buffer[offset:offset+size].view(dtype=dtype).view(shape)
yield item["name"], tensor
dist.barrier(group=ranks_group) # 清理
这个函数会通过 monkey-patch 替换掉 ParameterServer 原有的 receive_tensor 方法。
3. BroadcastOperation -- NCCL 广播¶
class BroadcastOperation:
def __init__(self, rank, ranks_group, bucket, metadata):
loop = asyncio.get_running_loop()
self._task = loop.run_in_executor(None, self._run)
def _run(self):
dist.broadcast(self.bucket, src=self.rank, group=self.ranks_group)
async def wait_for_complete(self):
await self._task
return self.metadata
注意与 NCCL 版本的区别:
- 这里的 broadcast 源 rank 不固定为 0,而是由 receiver_rank 决定(即每个 Rollout Worker 轮流作为 broadcast 源)
- 不需要 ZMQ 传元数据(元数据通过 ParameterServer 内部机制传递)
4. KIMICheckpointEngine -- 主体类¶
通信拓扑¶
@classmethod
def build_topology(cls, trainer_world_size, rollout_world_size, metadata):
trainer_kwargs = {
"rank": list(range(0, trainer_world_size)), # 所有 Trainer 都参与!
"trainer_world_size": [trainer_world_size] * trainer_world_size,
"rollout_world_size": [rollout_world_size] * trainer_world_size,
"master_metadata": [metadata[0]] * trainer_world_size,
}
rollout_kwargs = {
"rank": list(range(trainer_world_size, trainer_world_size + rollout_world_size)),
"trainer_world_size": [trainer_world_size] * rollout_world_size,
"rollout_world_size": [rollout_world_size] * rollout_world_size,
"master_metadata": [metadata[0]] * rollout_world_size,
}
return trainer_kwargs, rollout_kwargs
与 NCCL 的关键区别:所有 Trainer Worker 都参与通信(rank 0 到 trainer_world_size-1),而 NCCL 只有 rank 0 参与。
init_process_group¶
def init_process_group(self, rank, trainer_world_size, rollout_world_size, master_metadata):
self.rank = rank
self.trainer_world_size = trainer_world_size
self.world_size = trainer_world_size + rollout_world_size
if not self.initialized:
self.parameter_server = ParameterServer(
rank=rank, world_size=self.world_size,
auto_pg=False,
master_addr=master_metadata.dist_ip,
master_port=master_metadata.dist_port,
)
# monkey-patch:替换接收方法
self.parameter_server.receive_tensor = types.MethodType(
receive_tensor, self.parameter_server
)
dist.use_backend(f"vllm_{get_nccl_backend()}")
self.parameter_server.init_process_group()
# 创建 Rollout Worker 的子进程组
self.rollout_ranks = list(range(self.trainer_world_size, self.world_size))
self.rollout_group = dist.new_group(self.rollout_ranks)
self.initialized = True
send_weights -- CPU 卸载 + 参数注册¶
async def send_weights(self, weights):
named_tensors = {}
for named_tensors_gpu in ckpt_get_named_tensor_buckets(
weights, self.bucket_size, self.trainer_world_size,
self.rank, self.rollout_dtype
):
# 多线程并行将参数从 GPU 卸载到 CPU
with concurrent.futures.ThreadPoolExecutor(max_workers=32) as executor:
futures = [
executor.submit(lambda n, t: (n, t.to("cpu", non_blocking=True)), name, tensor)
for name, tensor in named_tensors_gpu.items()
]
for future in concurrent.futures.as_completed(futures):
name, tensor_cpu = future.result()
named_tensors[name] = tensor_cpu
get_torch_device().synchronize()
# 注册到参数服务器
self.parameter_server.register_checkpoint(self.checkpoint_name, named_tensors=named_tensors)
named_tensors = {}
get_torch_device().empty_cache()
# 收集元数据 + 同步
self.parameter_server.gather_metas(self.checkpoint_name)
dist.barrier()
self.parameter_server.unregister_checkpoint(self.checkpoint_name)
receive_weights -- 从参数服务器接收¶
async def receive_weights(self):
self.parameter_server.gather_metas(self.checkpoint_name)
async for name, tensor in self.parameter_server.receive_tensor(
self.checkpoint_name, self.rollout_group,
self.rollout_ranks, self.bucket_size
):
yield name, tensor
dist.barrier()
核心类/函数列表¶
| 名称 | 类型 | 作用 |
|---|---|---|
MasterMetadata |
dataclass | master 的地址信息 |
ckpt_get_named_tensor_buckets |
函数 | 参数轮询分片 + 分桶 |
receive_tensor |
函数 | 自定义的张量接收方法(monkey-patch) |
BroadcastOperation |
类 | 异步 NCCL broadcast |
KIMICheckpointEngine |
类 | Kimi 后端的完整实现 |
与其他后端的对比¶
| 方面 | NCCL | Kimi |
|---|---|---|
| 架构模式 | 集合通信 | 参数服务器 |
| Trainer 参与度 | 仅 rank 0 | 所有 rank |
| 参数分配 | 不分片 | 轮询分片 |
| CPU 卸载 | 无 | 有(多线程 GPU->CPU) |
| 接收方式 | 直接 NCCL broadcast | ParameterServer + broadcast |
与其他模块的关系¶
- base.py:继承
CheckpointEngine,注册为"kimi_ckpt_engine"。 checkpoint_engine.ps:Kimi 的 ParameterServer 库。checkpoint_engine.distributed:Kimi 的分布式通信库。verl.utils.device:获取设备相关工具函数。
小结¶
Kimi 后端采用参数服务器 + CPU 卸载的设计,让所有 Trainer Worker 分担参数传输的工作量。相比 NCCL 只由 rank 0 负责发送,这种方式在 Trainer 端的 GPU 显存压力更小(通过 CPU 卸载),且通信带宽更均衡(通过分片分担)。代码中使用了 monkey-patch 技术替换 ParameterServer 的接收方法,以适配 verl 框架的异步生成器接口。