跳转至

checkpoint_engine 模块总览

模块路径: verl/checkpoint_engine/

一、这个模块是做什么的?

在大模型强化学习(RLHF)训练中,通常存在两类角色:

  • Trainer(训练器):负责用 PPO 等算法更新模型权重。
  • Rollout(推理/采样器):负责用当前模型生成文本样本,供 Trainer 计算奖励和梯度。

Trainer 和 Rollout 通常运行在不同的 GPU 甚至不同的机器上。每次 Trainer 更新完权重后,需要把最新的权重同步给 Rollout,这样 Rollout 才能用最新的模型继续生成样本。

checkpoint_engine 模块就是负责这个权重同步过程的。它抽象出了统一的接口,并提供了多种通信后端实现。

二、模块架构图

                         checkpoint_engine 模块架构
 ============================================================================

  __init__.py          模块入口,导出所有类,按需加载各后端
       |
       v
    base.py            核心基础模块
       |
       +-- TensorMeta              张量元数据定义
       +-- CheckpointEngineRegistry 引擎注册表(工厂模式)
       +-- CheckpointEngine         抽象基类(定义统一接口)
       +-- CheckpointEngineWithCache 带本地缓存的扩展抽象类
       +-- ColocatedCheckpointEngine 共享GPU的简单实现(注册为"naive")
       +-- CheckpointEngineWorker    Ray Worker 封装
       +-- CheckpointEngineManager   全局协调管理器
       |
       +-----+-----+-----+-----+
       |     |     |     |     |
       v     v     v     v     v
     nccl  hccl  nixl  kimi  mooncake    <-- 五种通信后端实现
  (NVIDIA) (华为) (NIXL) (Kimi) (Mooncake)

 ============================================================================

  权重同步流程总览:

  ┌─────────────────────────────────┐     ┌──────────────────────────────┐
  │         Trainer 进程             │     │       Rollout 进程            │
  │                                 │     │                              │
  │  ModelEngine (FSDP/MCore/...)   │     │   Rollout Server (vLLM/...)  │
  │       |                         │     │        ^                     │
  │       | get_per_tensor_param()  │     │        | update_weights()   │
  │       v                         │     │        |                     │
  │  CheckpointEngine.send_weights()│ ==> │  CheckpointEngine            │
  │                                 │通信  │    .receive_weights()        │
  └─────────────────────────────────┘     └──────────────────────────────┘

  通信方式取决于后端:
  - naive:    共享 GPU 内存直接传递(Trainer 和 Rollout 在同一 GPU)
  - nccl:     NVIDIA NCCL 集合通信(broadcast)
  - hccl:     华为 HCCL 集合通信(broadcast)
  - nixl:     NIXL 点对点 RDMA 通信(流水线传递)
  - kimi:     Kimi ParameterServer 方式
  - mooncake: Mooncake TransferEngine 方式

三、文件清单与阅读顺序

建议按以下顺序阅读:

序号 文件 讲解文档 说明
1 __init__.py init.py 讲解 模块入口,了解整体结构
2 base.py base.py 讲解 最重要,定义所有核心抽象类
3 nccl_checkpoint_engine.py nccl 讲解 NCCL 后端,最典型的实现
4 hccl_checkpoint_engine.py hccl 讲解 HCCL 后端,华为 NPU 适配
5 nixl_checkpoint_engine.py nixl 讲解 NIXL 后端,RDMA 点对点
6 kimi_checkpoint_engine.py kimi 讲解 Kimi ParameterServer
7 mooncake_checkpoint_engine.py mooncake 讲解 Mooncake TransferEngine

四、核心概念速查

  • Bucket(桶):把多个小的权重张量打包到一块连续内存中一起传输,减少通信次数,提高带宽利用率。
  • 双缓冲(Double Buffering):用两块 buffer 交替发送和接收,实现通信与计算的流水线重叠。
  • ZeroMQ(ZMQ):轻量级消息队列,用于传输张量的元数据(名字、形状、类型、偏移量)。
  • Registry(注册表):工厂模式,通过字符串名字(如 "nccl")创建对应的引擎实例。
  • Ray:分布式计算框架,用于管理 Trainer 和 Rollout 的多个 Worker 进程。