Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
85 changes: 85 additions & 0 deletions docs/design/checkpoint-engine.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
# checkpoint-engine

## 简介和安装

[checkpoint-engine](https://github.com/MoonshotAI/checkpoint-engine) 是 Moonshot AI 开源的权重更新中间件,用于在 RL 训练中把训练侧权重高效同步到推理引擎。它的核心组件是 `ParameterServer`,支持两种更新方式:

- **Broadcast**:默认推荐路径,适合同步更新一组推理实例。
- **P2P**:适合动态新增或重启推理实例时,只向部分 rank 传输权重,依赖 `mooncake-transfer-engine` 支持 RDMA 传输。

安装方式:

```bash
# 只使用 broadcast
pip install checkpoint-engine

# 使用 P2P,会额外安装 mooncake-transfer-engine
pip install 'checkpoint-engine[p2p]'
```

## 进度

- [ x ] checkpoint-engine colocate SGLang engine 的常规权重更新
- [ ] checkpoint-engine colocate SGLang engine 的失败引擎重启
- [ ] checkpoint-engine colocate LMDeploy engine 的权重更新

## xtuner中使用方法

当前 XTuner 中 checkpoint-engine 只用于 colocate 场景下的 SGLang rollout backend。配置 rollout 时设置:

```python
rollout_config = RolloutConfig(
weight_transport_type="checkpoint_engine",
)
```
如未设置 `weight_transport_type`,colocate-RL将默认使用 `ipc`,disaggerated-RL默认使用 `NCCL`。

权重更新流程分为两步:

```python
self.train_controller.update_weights(need_register=True, need_update=False)
self.train_controller.offload(target="model")
ray.get(self.rollout_controller.onload_weights.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT)
self.train_controller.update_weights(need_register=False, need_update=True)
```

其中:

- `need_register=True, need_update=False`:从训练引擎收集权重并注册到 checkpoint-engine。
- `offload(target="model")`:释放训练侧模型显存,为 rollout onload/update 腾空间。
- `need_register=False, need_update=True`:复用已注册 checkpoint,把权重更新到 rollout engine。

## IPC/Checkpoint-engine显存和时间对比

IPC 路径直接把 tensor 通过 IPC 传给 rollout engine,链路较短,通常延迟更低,但在大模型和复杂并行场景下更容易受显存峰值影响。

checkpoint-engine 路径会先把训练侧权重注册到 `ParameterServer`,再由 checkpoint-engine 规划 bucket 并更新 rollout engine。它的好处是更适合大模型、分片权重和失败 engine 恢复;代价是会有一份模型 shard 常驻 CPU Memory,并且 D2H copy、bucket broadcast 会引入额外时间。

如果只需要在 `ParameterServer.register_checkpoint` 后显式等待一次 accelerator 侧异步操作完成,可以设置:

```python
rollout_config = RolloutConfig(
weight_transport_type="checkpoint_engine",
checkpoint_engine_sync_after_register=True,
)
```

该选项默认关闭,profile memory 发现,显式同步会释放掉 GPU 上的 tensor, 在注册时需要等 H2D 完成,而不显式同步,能更快的完成权重更新。

权重更新的显存峰值:

- IPC:trian worker weight + rollout worker weight + bucket

- checkpoint-engine(checkpoint_engine_sync_after_register=False):max(trian worker weight + parameter server shard, rollout worker weight + buffer *2)

- 可以优化降低成:max(trian worker weight + bucket, rollout worker weight + buffer *2),但这会影响性能
- parameter server shard 指将完整权重分成若干份,每一份的大小

- 这里 bucket 和 buffer 含义不同, bucket是train engine 一次导出的一个batch的权重大小, checkpoint-engine的buffer最小为单个权重的最大大小,该值默认为 8 GiB。


## debug小tips

- OOM:先看各 rank 的 checkpoint shard 是否划分均匀。日志中已打印 `[checkpoint_engine] collect matched local keys rank=xx parameter server shard total= xxx GiB `
- register 后显存不降:可将`checkpoint_engine_sync_after_register=True` 打开,注册后即可释放注册所需GPU memory。

4 changes: 3 additions & 1 deletion examples/v1/config/rl_grpo_gsm8k_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,8 @@
accelerator="GPU",
num_workers=8 * NNODE,
num_cpus_per_worker=12,
cpu_memory_per_worker=16 * 1024**3, # 16 GB
# 32 GB. Increased from 16 GB because checkpoint-engine shards use pinned memory.
cpu_memory_per_worker=32 * 1024**3,
)

# 2. rollout
Expand All @@ -64,6 +65,7 @@
gpu_memory_utilization=0.8,
context_length=max_response_length + max_prompt_length,
enable_return_routed_experts=(enable_return_routed_experts == "1"),
weight_transport_type="ipc"
)

# 3. judger
Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,8 @@ rl = [
"fastapi",
"uvicorn",
"mathruler",
"pylatexenc"
"pylatexenc",
"checkpoint-engine[p2p]"
]
video = [
"decord",
Expand Down
10 changes: 5 additions & 5 deletions tests/rl/test_qwen35_vl_moe_async_train_2step.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,7 @@ def test_qwen35_vl_moe_async_train_2step_and_metrics(self):
start_s = time.perf_counter()
trainer = self.build_config(work_dir).build()
self._record_produce_batch(trainer)
self._record_update_weights(trainer)
self._record_weight_update(trainer)
try:
trainer.fit()
finally:
Expand Down Expand Up @@ -343,14 +343,14 @@ async def produce_batch_wrapper(batch_size: int, train_step: int, *, model_step:

trainer.agent_loop_manager.produce_batch = produce_batch_wrapper

def _record_update_weights(self, trainer) -> None:
original_update_weights = trainer.train_controller.update_weights
def _record_weight_update(self, trainer) -> None:
original_weight_update = trainer.train_controller.weight_update

def update_weights_wrapper(*args, **kwargs):
self.update_weight_calls += 1
return original_update_weights(*args, **kwargs)
return original_weight_update(*args, **kwargs)

trainer.train_controller.update_weights = update_weights_wrapper
trainer.train_controller.weight_update = update_weights_wrapper

def _load_step_metrics(self, work_dir: Path) -> list[dict[str, float]]:
rows: list[dict[str, float]] = []
Expand Down
2 changes: 1 addition & 1 deletion tests/rl/test_rl_colocate_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,7 +163,7 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_
trainer.train_controller = SimpleNamespace(
onload=MagicMock(return_value="train_onloaded"),
offload=MagicMock(return_value="train_offloaded"),
update_weights=MagicMock(return_value="weights_updated"),
weight_update=MagicMock(return_value="weights_updated"),
fit=MagicMock(
return_value=[
{
Expand Down
15 changes: 6 additions & 9 deletions tests/rl/test_rl_disaggregated_trainer.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
"""RLDisaggregatedTrainer 的 public 行为测试。

Good Tests:
- 通过 fit()、update_weights()、同步周期校验和资源布局不变量验证行为。
- 通过 fit()、weight_update()、同步周期校验和资源布局不变量验证行为。
- 用 _FakeManager 和轻量 controller 替代真实 Ray worker。
- 只断言 step 递进、producer 恢复 model_step、checkpoint 文件等可观察结果。
- 对 async producer/consumer 只验证最终业务结果;内部任务编排放到 AgentLoopManager 测试中。
Expand All @@ -16,7 +16,7 @@
- disaggregated fit 遇到空 EXPIRED_BATCH 时重试同一个 train_step,不推进 _cur_step。
- 非空 EXPIRED_BATCH 仍会训练,并用当前完成的 model_step 恢复 producer。
- checkpoint 保存发生在 fit 完成的 model_step 上,且 manager.save 为 async 调用。
- eval 在 producer 恢复前运行;update_weights 本身不直接 pause/continue rollout controller。
- eval 在 producer 恢复前运行;weight_update 本身不直接 pause/continue rollout controller。
- sync/checkpoint/eval interval 必须是 sync_weights_interval 的整数倍,资源布局必须 fail fast。
- 前台训练 batch 阻塞时,后台 producer 仍能在事件循环中继续推进。
"""
Expand Down Expand Up @@ -144,7 +144,7 @@ def _make_trainer(self, agent_loop_manager):
fit=MagicMock(return_value=[{"train_metrics": [], "sft_train_metrics": {}}]),
onload=MagicMock(return_value="onload"),
offload=MagicMock(return_value="offload"),
update_weights=MagicMock(return_value="update"),
weight_update=MagicMock(return_value="update"),
)
trainer.rollout_controller = SimpleNamespace(
check_and_shutdown_inactive_workers=SimpleNamespace(
Expand Down Expand Up @@ -273,9 +273,6 @@ def test_fit_rebinds_weight_update_with_rollout_update_address(self):
train_controller=trainer.train_controller,
rollout_controller=trainer.rollout_controller,
rollout_config=trainer._rollout_config,
weight_transport_type="nccl",
weight_update_host="10.0.0.1",
weight_update_port=23456,
)

def test_fit_keeps_background_producer_running_while_training_blocks(self):
Expand Down Expand Up @@ -394,8 +391,8 @@ def test_update_weights_pauses_generation_without_onloading_rollout(self):
trainer = self._make_trainer(manager)

with patch("xtuner.v1.train.rl_trainer.ray.get", side_effect=lambda obj, timeout=None: obj):
trainer.update_weights = RLDisaggregatedTrainer.update_weights.__get__(trainer, RLDisaggregatedTrainer)
trainer.update_weights()
trainer.weight_update = RLDisaggregatedTrainer.weight_update.__get__(trainer, RLDisaggregatedTrainer)
trainer.weight_update()

trainer.rollout_controller.pause_generation.remote.assert_not_called()
trainer.rollout_controller.continue_generation.remote.assert_not_called()
Expand Down Expand Up @@ -424,7 +421,7 @@ async def manager_continue_produce(model_step: int):
resume=AsyncMock(side_effect=manager_resume),
continue_produce=AsyncMock(side_effect=manager_continue_produce),
)
trainer.update_weights = MagicMock(side_effect=lambda: events.append("update_weights"))
trainer.weight_update = MagicMock(side_effect=lambda: events.append("update_weights"))

train_state_path = Path(self.temp_dir.name) / trainer._SAVE_TRAIN_STATE_PATH
train_state_path.write_text('{"cur_step": 3}')
Expand Down
8 changes: 1 addition & 7 deletions tests/rl/test_rl_trainer_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,25 +132,19 @@ def bind_rollout_weight_update(
*,
targets,
rollout_config,
weight_transport_type,
weight_update_host=None,
weight_update_port=None,
):
self.rollout_info = {
"targets": targets,
"rollout_config": rollout_config,
}
self.weight_transport_type = weight_transport_type
self.weight_update_host = weight_update_host
self.weight_update_port = weight_update_port

def onload(self, target="all"):
return f"onload:{target}"

def offload(self, target="all"):
return f"offload:{target}"

def update_weights(self):
def weight_update(self):
self.update_weights_count += 1
return "updated"

Expand Down
7 changes: 6 additions & 1 deletion tests/rl/test_rollout_logic.py
Original file line number Diff line number Diff line change
Expand Up @@ -172,6 +172,12 @@ def _rollout_config(
expert_parallel_size=ep,
num_gpus_per_engine=num_gpus_per_engine,
gpus_per_node=gpus_per_node,
weight_transport_type="ipc",
weight_update_host=None,
weight_update_port=30000,
checkpoint_name_prefix="xtuner-rl",
checkpoint_engine_timeout=300.0,
checkpoint_engine_sync_after_register=False,
extra_rollout_config={"lmdeploy_backend": "pytorch"},
)

Expand Down Expand Up @@ -202,7 +208,6 @@ def _rollout_info(self, *, config, targets, train_rank: int):
rollout_config=config,
weight_update_targets=targets,
train_rank=train_rank,
weight_transport_type="ipc",
)

def test_rollout_topology_resolves_engine_dist_init_addr_when_created(self):
Expand Down
Loading
Loading