跳转至

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 -- 张量元数据

class TensorMeta(TypedDict):
    name: str
    shape: torch.Size
    dtype: torch.dtype
    offset: int

这是一个 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)

这是工厂模式 + 注册表模式的经典实现。各后端引擎在定义时通过装饰器注册:

@CheckpointEngineRegistry.register("nccl")
class NCCLCheckpointEngine(CheckpointEngine):
    ...

使用时只需传入字符串名字即可创建实例:

engine = CheckpointEngineRegistry.new("nccl", bucket_size=1024*1024*256)

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 模块的设计思路。