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 进程。