[Feature] Add Checkpoint Engine as a transport between train and rollout - #1993
[Feature] Add Checkpoint Engine as a transport between train and rollout#1993PengchengShi00 wants to merge 3 commits into
Conversation
| num_workers=8 * NNODE, | ||
| num_cpus_per_worker=12, | ||
| cpu_memory_per_worker=16 * 1024**3, # 16 GB | ||
| cpu_memory_per_worker=32 * 1024**3, # 32 GB |
| group. Defaults to None. | ||
| weight_update_port (Optional[int]): Port used by train rank 0 to initialize the external NCCL weight update | ||
| group. Defaults to 30000. | ||
| enable_checkpoint_engine (bool): Whether to use Checkpoint Engine to synchronize training weights to rollout workers. When enabled, train workers create in-process ParameterServer instances and broadcast weights through Checkpoint Engine. Defaults to False. |
There was a problem hiding this comment.
我有点疑问,checkpoint_engine是一种权重同步方式,为啥需要单独的 flag 来启动。合理的是不是应该暴露权重同步 type,比如 ipc nccl ce 和 disk。然后只要注明每种方式用于啥场合。而不是为这一种权重同步方式单独新增 flag
| def update_weights(self, need_register: bool = True, need_update: bool = True): | ||
| """Update the weights from the training workers.""" | ||
| handles = [ | ||
| worker.update_weights.remote(need_register=need_register, need_update=need_update) |
There was a problem hiding this comment.
然后这种need_register参数不是通用的,那么就不应该暴露在外面,而是通过 kw 传进去就行。在 controller 层是无感的
| @ray_method | ||
| def update_weights(self): | ||
| return self.update_weighter.update_weights() | ||
| def update_weights(self, need_register: bool = True, need_update: bool = True): |
| ray.get(self.rollout_controller.onload_weights.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) | ||
|
|
||
| start_time = time.perf_counter() | ||
| self.train_controller.update_weights(need_register=False) |
There was a problem hiding this comment.
skip_load_weights=True 时 Checkpoint Engine 初始化同步必然失败。xtuner/v1/train/rl_trainer.py:1639 第一次同步直接传入 need_register=False;但 transport 构造时只读取 key 列表,_checkpoint_name 仍为
None,随后在 xtuner/v1/rl/weight_update/transport.py:1134 明确抛出 RuntimeError
| self._set_transport() | ||
|
|
||
| def update_weights(self): | ||
| def update_weights(self, need_register: bool = True, need_update: bool = True) -> None: |
There was a problem hiding this comment.
训练侧 ep_size > 1 时,专家权重可能被错误切分并绑定到错误的 HF key。xtuner/v1/rl/weight_update/weight_iterator.py:193 只有 NCCL 会 gather train-EP tensor;新加入的 checkpoint_engine 不会走这个逻辑,应该是不对的
| gpu_memory_utilization=0.8, | ||
| context_length=max_response_length + max_prompt_length, | ||
| enable_return_routed_experts=(enable_return_routed_experts == "1"), | ||
| enable_checkpoint_engine=True |
There was a problem hiding this comment.
ce 这个依赖,好安装吗?还有如果 p2p 还需要 mooncake。如果都好安装,可以写到pyproject.toml 的 rl item 中,可以考虑锁定版本
| weight_update_port=weight_update_port, | ||
| ) | ||
|
|
||
| new_transport_signature = self.rollout_info.transport_signature |
There was a problem hiding this comment.
根据 transport signature 判断是否 teardown,而 signature 包含:
- worker URL;
- lifecycle state;
- active/inactive 状态;
- rank topology。
只要 rollout worker 重启、URL 改变,就会销毁整个 CE transport,从而销毁已经注册的 checkpoint。
但 CE/P2P 的主要价值之一,恰恰是保留现有 checkpoint,然后只把它推给恢复的 rollout worker。所以这个地方有点不合理
| with open(index_path) as f: | ||
| weight_map: dict[str, str] = json.load(f)["weight_map"] | ||
| weight_keys = list(key_name for key_name, file_name in weight_map.items()) | ||
| per_rank = (len(weight_keys) + world_size - 1) // world_size |
There was a problem hiding this comment.
参数数量均衡不代表字节数均衡。如果要考虑更均衡可能会过于复杂,不过就目前而言,应该要加一些打印性能来表征这种不平衡,防止后续出现性能瓶颈或者无法解释情况下不知道为啥?
| """Register current train engine weights into Checkpoint Engine PS.""" | ||
|
|
||
| # 1. Collect named tensors from weight iterator | ||
| all_tensors = self._collect_named_tensors(weight_iterator, local_keys=self._local_checkpoint_keys) |
There was a problem hiding this comment.
这个地方有比较大的显存峰值吧?要想想这个地方会不会成为 oom 瓶颈点。ce 默认这个地方就是这么做的吗?
| enable_checkpoint_engine (bool): Whether to use Checkpoint Engine to synchronize training weights to rollout workers. When enabled, train workers create in-process ParameterServer instances and broadcast weights through Checkpoint Engine. Defaults to False. | ||
| checkpoint_name_prefix (str): Prefix used for Checkpoint Engine checkpoint names registered in the | ||
| ParameterServer. Defaults to "xtuner-rl". | ||
| checkpoint_engine_timeout (float): Timeout in seconds for Checkpoint Engine rollout weight update requests. |
There was a problem hiding this comment.
感觉checkpoint engine的配置与rolloutconfig中指定不太合理,最终形态是想权重更新、训练、推理完全独立
| def update_weights(self): | ||
| return self.update_weighter.update_weights() | ||
| def update_weights(self, need_register: bool = True, need_update: bool = True): | ||
| return self.update_weighter.update_weights(need_register=need_register, need_update=need_update) |
There was a problem hiding this comment.
对了,感觉这个权重更新的名字可以改为 weight_updater,感觉更合理一点
| def build_update_url(self, server_url: str) -> str: | ||
| raise NotImplementedError | ||
|
|
||
| def before_update(self) -> None: |
| return f"{server_url.rstrip('/')}/update_weights_from_ipc" | ||
|
|
||
|
|
||
| class CheckpointEngineWeightTransport(WeightTransport[CheckpointEngineBackendAdapter]): |
There was a problem hiding this comment.
CheckpointEngineWeightTransport 不适合继承现有 batch-oriented WeightTransport:它的 send() 永远抛错,而父类 update(**_) 又会静默吞掉 CE 专属参数。建议改成显式 staged API:
- prepare(weight_iterator) -> CheckpointHandle
- publish(handle, targets)
- sync() 组合以上两步
这样状态就是 EMPTY -> PREPARED -> PUBLISHED,无需用 need_register/need_update 两个布尔值表达非法组合。内部再拆成 manifest planner、ParameterServer lifecycle、SGLang control client;同时注入 process group、PS factory 和 HTTP client,才能写不依赖 GPU 的单元测试。
Checkpoint Engine 权重同步
1. 背景
XTuner RL 训练需要周期性把 train engine 中的最新权重同步到 rollout engines。现有权重同步路径包括:
ParameterServer,再由ParameterServer推送给 rollout engines。Checkpoint Engine 路径的收益:
ParameterServer,再 offload train engine、onload rollout engine,最后由ParameterServer推送权重。当前实现只在共卡训练路径中启用 Checkpoint Engine。分离模式检测到
enable_checkpoint_engine=True时会打印 warning,并回退到 NCCL weight transport。2. 架构设计
共卡模式下,在 rollout config 中设置
enable_checkpoint_engine=True,可使用 Checkpoint-Engine 进行权重更新, 分离模式下即使设置enable_checkpoint_engine=True,也会回退到 NCCL transport。2.1 初始化
ParameterServer训练 worker 内部创建
ParameterServer,读取 train world size,作为ParameterServerworld size,不额外启动独立 PS 进程。auto_pg=False表示 Checkpoint Engine 复用 XTuner 已经初始化好的torch.distributed默认进程组,不在 update 后销毁该进程组。根据 HF checkpoint 的
model.safetensors.index.json计算当前 PS-rank 负责的参数 key 集合。每个 PS rank 只负责注册自己分到的参数 shard。这些 key 后续用于从WeightIterator中过滤当前 PS rank 需要注册的 tensor,避免每个 rank 都注册全量参数。2.2 更新
CheckpointEngineWeightTransport.update(...)被拆成两个可选阶段:当
need_register=True时,transport 会:weight_iterator.iter_batch_groups()收集 train engine 当前权重。ParameterServer.register_checkpoint(...)注册新 checkpoint:如果 HF index 中的某些 key 没有从 train engine 收集到,transport 会记录错误日志。MTP-only key 和非 MTP key 会分开打印,便于排查模型结构差异。
当
need_update=True时,transport 会把当前 checkpoint 推送到 rollout engines:rollout_info.active_update_targets获取当前活跃 rollout targets。update_ranks非空、不越界、不重复。ParameterServer.update(..., ranks=None),使用 Checkpoint Engine broadcast 路径。ParameterServer.update(..., ranks=active_ranks),使用 p2p update 路径。need_register和need_updateneed_register和need_update用于控制 Checkpoint Engine 更新的两个阶段。TrueFalseneed_registerneed_update典型使用方式:
need_update=False的核心用途是把注册和 rollout 更新拆成两个时刻执行。这样训练侧可以先把权重注册到ParameterServer,然后释放或切换部分资源,再让 rollout engines 加载 checkpoint,从而降低峰值显存压力。3. 共卡模式 IPC / Checkpoint Engine 速度对比
使用两种更新路径的端到端权重更新时间,包括了
rollout onload_weights、train offload时间。