base.py — 这是整个 checkpoint_engine 模块最核心的文件¶
文件路径:
verl/checkpoint_engine/base.py
文件概述¶
这是整个 checkpoint_engine 模块最核心的文件。它定义了:
1. 张量元数据类型 TensorMeta
2. 引擎注册表 CheckpointEngineRegistry(工厂模式)
3. 抽象基类 CheckpointEngine(定义所有后端必须实现的接口)
4. 带缓存的扩展抽象类 CheckpointEngineWithCache
5. 共享GPU的简单实现 ColocatedCheckpointEngine
6. Ray Worker 封装 CheckpointEngineWorker
7. 全局管理器 CheckpointEngineManager
关键代码讲解¶
1. TensorMeta -- 张量元数据¶
这是一个 TypedDict,描述了一个张量在 bucket(桶)中的位置信息:
- name:权重名称,如 "model.layers.0.self_attn.q_proj.weight"
- shape:张量形状,如 torch.Size([4096, 4096])
- dtype:数据类型,如 torch.bfloat16
- offset:在 bucket 内存中的字节偏移量
接收端通过这些信息,从连续的 bucket 内存中还原出每个张量。
2. CheckpointEngineRegistry -- 引擎注册表¶
class CheckpointEngineRegistry:
"""Checkpoint engine registry."""
_registry: dict[str, type["CheckpointEngine"]] = {}
def register(backend: str):
"""注册一个后端,用作装饰器"""
def wrapper(cls: type["CheckpointEngine"]):
CheckpointEngineRegistry._registry[backend] = cls
return cls
return wrapper
@classmethod
def get(cls, backend: str) -> type["CheckpointEngine"]:
"""根据名字获取引擎类"""
return cls._registry[backend]
@classmethod
def new(cls, backend: str, *args, **kwargs) -> "CheckpointEngine":
"""根据名字创建引擎实例"""
if backend not in cls._registry:
raise ValueError(f"Checkpoint engine {backend} not registered")
return cls._registry[backend](*args, **kwargs)
这是工厂模式 + 注册表模式的经典实现。各后端引擎在定义时通过装饰器注册:
使用时只需传入字符串名字即可创建实例:
3. CheckpointEngine -- 抽象基类¶
class CheckpointEngine(ABC):
"""CheckpointEngine is an abstraction to transfer weights from trainer to rollout."""
@abstractmethod
def prepare(self) -> dict[str, Any]:
"""准备工作:分配 buffer,注册 RDMA 内存,返回通信元数据"""
raise NotImplementedError
@classmethod
@abstractmethod
def build_topology(cls, trainer_world_size, rollout_world_size, metadata):
"""根据所有 worker 的元数据,构建通信拓扑"""
raise NotImplementedError
@abstractmethod
def init_process_group(self, **kwargs):
"""初始化进程组"""
raise NotImplementedError
@abstractmethod
def finalize(self):
"""清理资源:释放 buffer,销毁进程组"""
raise NotImplementedError
@abstractmethod
async def send_weights(self, weights):
"""发送权重(Trainer 端调用)"""
raise NotImplementedError
@abstractmethod
async def receive_weights(self):
"""接收权重(Rollout 端调用)"""
raise NotImplementedError
这定义了权重同步的完整生命周期:
prepare() --> build_topology() --> init_process_group()
--> send_weights() / receive_weights()
--> finalize()
4. ColocatedCheckpointEngine -- 最简单的实现¶
@CheckpointEngineRegistry.register("naive")
class ColocatedCheckpointEngine(CheckpointEngine):
"""Trainer 和 Rollout 共享同一个 GPU 时使用"""
def send_weights(self, weights):
self.weights = weights # 直接保存引用
def receive_weights(self):
yield from self.weights # 直接返回引用
self.weights = None
当 Trainer 和 Rollout 在同一个 GPU 上时,不需要任何网络通信,直接通过 Python 引用传递权重即可。这是注册为 "naive" 的最简实现。
5. CheckpointEngineWorker -- Ray Worker 封装¶
class CheckpointEngineWorker(Worker):
def __init__(self, rollout_config, model_config, server_adapter=None, *args, **kwargs):
super().__init__()
# 从配置中获取后端类型和 bucket 大小
backend = self.rollout_config.checkpoint_engine.backend
bucket_size = self.rollout_config.checkpoint_engine.update_weights_bucket_megabytes << 20
engine_kwargs = self.rollout_config.checkpoint_engine.engine_kwargs.get(backend, {})
# 通过注册表创建引擎
self.checkpoint_engine = CheckpointEngineRegistry.new(
backend, bucket_size=bucket_size, **engine_kwargs
)
# 如果没有提供 server_adapter,则根据配置创建 rollout
if self.server_adapter is None:
self.server_adapter = get_rollout_class(...)( ... )
# 初始化全局进程组(用于 sglang、trt-llm 等内部通信)
initialize_global_process_group_ray(timeout_second=None, backend="cpu:gloo")
@register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False)
async def update_weights(self, global_steps=None):
"""接收权重并更新到推理引擎"""
weights = self.checkpoint_engine.receive_weights()
await self.server_adapter.update_weights(weights, global_steps=global_steps)
@register(dispatch_mode=Dispatch.DP_COMPUTE, blocking=False)
def execute_checkpoint_engine(self, method, *args, **kwargs):
"""通用方法调用转发,用于远程调用 prepare/init_process_group/finalize"""
return getattr(self.checkpoint_engine, method)(*args, **kwargs)
这个类是 Rollout 端的 Worker,它将 CheckpointEngine 与推理引擎(如 vLLM)组合在一起。关键设计:
- update_weights 先从 checkpoint engine 接收权重,再更新到推理引擎
- execute_checkpoint_engine 是一个通用的方法代理,可以远程调用 engine 的任意方法
6. CheckpointEngineManager -- 全局管理器¶
这是整个模块的指挥中心,协调 Trainer 和所有 Rollout 副本之间的权重同步。
class CheckpointEngineManager:
def __init__(self, config, trainer, replicas):
self.backend = config.backend
self.backend_cls = CheckpointEngineRegistry.get(config.backend)
self.trainer = trainer
self.replicas = replicas
核心方法 update_weights 的流程:
async def update_weights(self, global_steps=None):
# 0. 如果是 naive 模式(共享GPU),直接更新
if self.backend == "naive":
ray.get(self.trainer.update_weights(global_steps=global_steps))
return
# 1. 中断 Rollout 正在进行的请求(partial rollout 场景)
await asyncio.gather(*[r.abort_all_requests() for r in self.replicas])
# 2. 把所有 Rollout 副本的 worker 合成一个临时 WorkerGroup
workers = []
for replica in self.replicas:
workers.extend(replica.workers)
rollout = RayWorkerGroup(worker_handles=workers, ...)
# 3. 建立通信进程组
self.build_process_group(rollout)
# 4. 执行权重同步
ray.get(trainer.update_weights(...) + rollout.update_weights(...))
# 5. 清理资源
ray.get(trainer.execute_checkpoint_engine(["finalize"] * ...) + ...)
# 6. 恢复之前中断的请求
await asyncio.gather(*[r.resume_generation() for r in self.replicas])
build_process_group 方法实现了三步建立通信的流程:
def build_process_group(self, rollout):
# 1. 所有 worker 执行 prepare(),返回各自的元数据
metadata = ray.get(
trainer.execute_checkpoint_engine(["prepare"] * trainer.world_size)
+ rollout.execute_checkpoint_engine(["prepare"] * rollout.world_size)
)
# 2. 根据元数据构建通信拓扑
trainer_kwargs, rollout_kwargs = self.backend_cls.build_topology(
trainer.world_size, rollout.world_size, metadata
)
# 3. 所有 worker 执行 init_process_group()
ray.get(
trainer.execute_checkpoint_engine(**trainer_kwargs)
+ rollout.execute_checkpoint_engine(**rollout_kwargs)
)
源代码中有一幅 ASCII 架构图,展示了 Trainer 和 Rollout 之间的关系:
┌────────┬────────┬─────┬────────┐ ┌───────────────────┬───────────────────┐
│ ┌────┐ │ ┌────┐ │ │ ┌────┐ │ │ Replica 0 │ Replica 1 │
│ │ ME0│ │ │ ME1│ │ │ │ MEn│ │ ├────┬────┬────┬────┼────┬────┬────┬────┤
│ └──┬─┘ │ └────┘ │ ... │ └────┘ │ │ 0 │ 1 │ 2 │ 3 │ 0 │ 1 │ 2 │ 3 │
│ v | | | | └──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┘
| ┌──┴─┐ │ ┌────┐ │ │ ┌────┐ │ ^ ^ ^ cuda ipc ^ ^ ^
│ │ CE │ │ │ CE │ │ │ │ CE │ │ ┌──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┐
│ └──┬─┘ │ └────┘ │ │ └────┘ │ │ CE │ CE │ CE │ CE │ CE │ CE │ CE │ CE |
└────┼───┴────────┴─────┴────────┘ └──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┘
v | | | | | | | |
└─────────────(nccl/nixl/..)─────────────┴────┴────┴────┴────┴────┴────┴────┘
- ME = ModelEngine(训练引擎),CE = CheckpointEngine
- Trainer 端只有 rank 0 的 CE 实际发送数据
- Rollout 端每个 worker 的 CE 都接收数据,然后通过 cuda ipc 更新推理引擎
核心类/函数列表¶
| 名称 | 类型 | 作用 |
|---|---|---|
TensorMeta |
TypedDict | 描述张量在 bucket 中的位置 |
CheckpointEngineRegistry |
类 | 工厂模式注册表 |
CheckpointEngine |
ABC 抽象类 | 定义权重同步的统一接口 |
CheckpointEngineWithCache |
ABC 抽象类 | 增加本地缓存能力(用于 partial rollout) |
ColocatedCheckpointEngine |
具体类 | 共享 GPU 的最简实现 |
CheckpointEngineWorker |
Ray Worker | Rollout 端的 Worker 封装 |
CheckpointEngineManager |
管理器 | 协调全局权重同步流程 |
与其他模块的关系¶
verl.single_controller:CheckpointEngineWorker继承自Worker基类,使用@register装饰器暴露 Ray 远程方法。verl.workers.rollout:CheckpointEngineWorker内部持有BaseRollout实例,用于将接收到的权重更新到推理引擎。verl.workers.config:CheckpointEngineConfig、RolloutConfig等配置类。- 各后端实现:
nccl_checkpoint_engine.py等文件中的类都继承自CheckpointEngine,并通过@CheckpointEngineRegistry.register(...)注册。
小结¶
base.py 是整个模块的骨架。它通过抽象基类定义了统一的权重同步接口,通过注册表模式实现了后端的可插拔切换,通过 Manager 类封装了完整的同步流程。理解了这个文件,就理解了整个 checkpoint_engine 模块的设计思路。