跳转至

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 框架的异步生成器接口。