Skip to content

[Feature] Add Checkpoint Engine as a transport between train and rollout - #1993

Open
PengchengShi00 wants to merge 3 commits into
InternLM:mainfrom
PengchengShi00:checkpoint-engine
Open

[Feature] Add Checkpoint Engine as a transport between train and rollout#1993
PengchengShi00 wants to merge 3 commits into
InternLM:mainfrom
PengchengShi00:checkpoint-engine

Conversation

@PengchengShi00

Copy link
Copy Markdown
Collaborator

Checkpoint Engine 权重同步

1. 背景

XTuner RL 训练需要周期性把 train engine 中的最新权重同步到 rollout engines。现有权重同步路径包括:

路径 适用场景 说明
IPC 共卡模式 训练侧直接把权重通过 IPC 更新给 rollout backend。
Checkpoint Engine 共卡模式 训练侧把权重注册到 in-process ParameterServer,再由 ParameterServer 推送给 rollout engines。
NCCL 分离模式 train workers 和 rollout workers 通过 NCCL broadcast 同步权重。

Checkpoint Engine 路径的收益:

  • 降低共卡权重更新时的显存峰值。训练侧可以先把权重注册到 ParameterServer,再 offload train engine、onload rollout engine,最后由 ParameterServer 推送权重。
  • 支持恢复失败的 rollout engine。rollout engine 重启后,可以复用上一次注册成功的 checkpoint 重新推送权重。

当前实现只在共卡训练路径中启用 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 初始化

  1. 创建ParameterServer

训练 worker 内部创建 ParameterServer,读取 train world size,作为 ParameterServer world size,不额外启动独立 PS 进程。

ParameterServer(
    auto_pg=False,
    rank=self.rank,
    world_size=self.ps_world_size,
)

auto_pg=False 表示 Checkpoint Engine 复用 XTuner 已经初始化好的 torch.distributed 默认进程组,不在 update 后销毁该进程组。

  1. 划分参数shard

根据 HF checkpoint 的 model.safetensors.index.json 计算当前 PS-rank 负责的参数 key 集合。每个 PS rank 只负责注册自己分到的参数 shard。这些 key 后续用于从 WeightIterator 中过滤当前 PS rank 需要注册的 tensor,避免每个 rank 都注册全量参数。

2.2 更新

CheckpointEngineWeightTransport.update(...) 被拆成两个可选阶段:

update(weight_iterator, need_register=True, need_update=True)
  1. Register 阶段

need_register=True 时,transport 会:

  • weight_iterator.iter_batch_groups() 收集 train engine 当前权重。
  • 根据本 rank 的 local checkpoint keys 过滤 tensor。
  • 注销上一轮 checkpoint 名称,限制 pinned host memory 占用。
  • 调用 ParameterServer.register_checkpoint(...) 注册新 checkpoint:

如果 HF index 中的某些 key 没有从 train engine 收集到,transport 会记录错误日志。MTP-only key 和非 MTP key 会分开打印,便于排查模型结构差异。

  1. Update 阶段

need_update=True 时,transport 会把当前 checkpoint 推送到 rollout engines:

  • rollout_info.active_update_targets 获取当前活跃 rollout targets。
  • 校验每个 target 声明的 update_ranks 非空、不越界、不重复。
  • 判断本次更新是否覆盖全部 PS ranks:
    • 覆盖全部 ranks 时,调用 ParameterServer.update(..., ranks=None),使用 Checkpoint Engine broadcast 路径。
    • 只覆盖部分 ranks 时,调用 ParameterServer.update(..., ranks=active_ranks),使用 p2p update 路径。
  1. need_registerneed_update

need_registerneed_update 用于控制 Checkpoint Engine 更新的两个阶段。

参数 True False
need_register 从 train engine 注册一个新的 checkpoint。 复用上一次注册成功的 checkpoint 名称。
need_update 把当前 checkpoint 更新到 rollout engines。 只完成注册,不推送到 rollout engines。

典型使用方式:

# 常规更新:注册 train 权重并同步到 rollout engines
train_controller.update_weights(need_register=True, need_update=True)

# rollout 恢复:复用上一次 checkpoint,只重新更新 rollout engines
train_controller.update_weights(need_register=False, need_update=True)

# 显存紧张:先注册 checkpoint,稍后再更新 rollout engines
train_controller.update_weights(need_register=True, need_update=False)
offload(train)
onload(rollout)
train_controller.update_weights(need_register=False, need_update=True)

need_update=False 的核心用途是把注册和 rollout 更新拆成两个时刻执行。这样训练侧可以先把权重注册到 ParameterServer,然后释放或切换部分资源,再让 rollout engines 加载 checkpoint,从而降低峰值显存压力。

3. 共卡模式 IPC / Checkpoint Engine 速度对比

使用两种更新路径的端到端权重更新时间,包括了rollout onload_weightstrain offload时间。

模型 并行配置 Backend IPC更新时间 Checkpoint Engine更新时间
Qwen3-4B TP=2 SGLang 0.86s 2.3s
Qwen3-8B TP=2 SGLang 1.1s 2.8s
Qwen3-30B-A3B TP=2 SGLang 3.7s 5s
Qwen3.5-35B-A3B EP=4 SGLang 4.3s 6.5s

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

注明下修改原因

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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

我有点疑问,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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

然后这种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):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

同样道理

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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

训练侧 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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

ce 这个依赖,好安装吗?还有如果 p2p 还需要 mooncake。如果都好安装,可以写到pyproject.toml 的 rl item 中,可以考虑锁定版本

weight_update_port=weight_update_port,
)

new_transport_signature = self.rollout_info.transport_signature

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

根据 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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

参数数量均衡不代表字节数均衡。如果要考虑更均衡可能会过于复杂,不过就目前而言,应该要加一些打印性能来表征这种不平衡,防止后续出现性能瓶颈或者无法解释情况下不知道为啥?

"""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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个地方有比较大的显存峰值吧?要想想这个地方会不会成为 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

感觉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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

对了,感觉这个权重更新的名字可以改为 weight_updater,感觉更合理一点

def build_update_url(self, server_url: str) -> str:
raise NotImplementedError

def before_update(self) -> None:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这两个接口目前有用到嘛 ?

return f"{server_url.rstrip('/')}/update_weights_from_ipc"


class CheckpointEngineWeightTransport(WeightTransport[CheckpointEngineBackendAdapter]):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 的单元测试。

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants