diff --git a/docs/design/checkpoint-engine.md b/docs/design/checkpoint-engine.md new file mode 100644 index 0000000000..68c775c440 --- /dev/null +++ b/docs/design/checkpoint-engine.md @@ -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。 + diff --git a/examples/v1/config/rl_grpo_gsm8k_async.py b/examples/v1/config/rl_grpo_gsm8k_async.py index 37f76ad9bd..ded9a4a427 100644 --- a/examples/v1/config/rl_grpo_gsm8k_async.py +++ b/examples/v1/config/rl_grpo_gsm8k_async.py @@ -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 @@ -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 diff --git a/pyproject.toml b/pyproject.toml index bb62f70004..a941493dfe 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -74,7 +74,8 @@ rl = [ "fastapi", "uvicorn", "mathruler", - "pylatexenc" + "pylatexenc", + "checkpoint-engine[p2p]" ] video = [ "decord", diff --git a/tests/rl/test_qwen35_vl_moe_async_train_2step.py b/tests/rl/test_qwen35_vl_moe_async_train_2step.py index a70b36a98a..34a28ac9ac 100644 --- a/tests/rl/test_qwen35_vl_moe_async_train_2step.py +++ b/tests/rl/test_qwen35_vl_moe_async_train_2step.py @@ -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: @@ -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]] = [] diff --git a/tests/rl/test_rl_colocate_trainer.py b/tests/rl/test_rl_colocate_trainer.py index 69f902d6d0..e407be3836 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -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=[ { diff --git a/tests/rl/test_rl_disaggregated_trainer.py b/tests/rl/test_rl_disaggregated_trainer.py index 80b96c6dae..b278cbc718 100644 --- a/tests/rl/test_rl_disaggregated_trainer.py +++ b/tests/rl/test_rl_disaggregated_trainer.py @@ -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 测试中。 @@ -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 仍能在事件循环中继续推进。 """ @@ -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( @@ -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): @@ -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() @@ -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}') diff --git a/tests/rl/test_rl_trainer_checkpoint.py b/tests/rl/test_rl_trainer_checkpoint.py index cb2977b6c8..444ea83167 100644 --- a/tests/rl/test_rl_trainer_checkpoint.py +++ b/tests/rl/test_rl_trainer_checkpoint.py @@ -132,17 +132,11 @@ 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}" @@ -150,7 +144,7 @@ def onload(self, target="all"): def offload(self, target="all"): return f"offload:{target}" - def update_weights(self): + def weight_update(self): self.update_weights_count += 1 return "updated" diff --git a/tests/rl/test_rollout_logic.py b/tests/rl/test_rollout_logic.py index 5b0aa5ad8d..7850de2fa2 100644 --- a/tests/rl/test_rollout_logic.py +++ b/tests/rl/test_rollout_logic.py @@ -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"}, ) @@ -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): diff --git a/tests/rl/test_update_weight_colocate.py b/tests/rl/test_update_weight_colocate.py new file mode 100644 index 0000000000..4a058c3c48 --- /dev/null +++ b/tests/rl/test_update_weight_colocate.py @@ -0,0 +1,223 @@ +# Scope: colocate model weight update correctness for IPC and checkpoint-engine. +# This test currently covers only the SGLang backend with a parameter-only check. +# The SGLang parameter-only WeightChecker actions are implemented in +# https://github.com/PengchengShi00/sglang/commit/05e89d63b5a1a80671b267ff4494ad950b2aba75. +# Flow: snapshot_parameters -> reset_parameters -> update_weights -> compare_parameters. + +import os +import tempfile +import unittest + +import ray +import requests + +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.model import Qwen3_5_VLMoE35BA3Config +from xtuner.v1.module.mtp import MTPConfig + +from xtuner.v1.rl.loss import GRPOLossConfig as LossConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.trainer import ( + TrainingController, + TrainingWorker as BaseTrainingWorker, + WorkerConfig, +) +from xtuner.v1.rl.utils import ( + AcceleratorResourcesConfig, + AutoAcceleratorWorkers, + CPUResourceManager, + clear_cpu_resource_manager, + set_cpu_resource_manager, +) + +MODEL_PATH = os.environ["QWEN3_5_MOE_PATH"] + + +class TestUpdateWeightColocate(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + if MODEL_PATH is None: + raise unittest.SkipTest("MODEL_PATH is not set") + os.environ["XTUNER_USE_FA3"] = "1" + os.environ["NCCL_CUMEM_ENABLE"] = "0" + os.environ["NCCL_IB_HCA"] = "mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_4,mlx5_5,mlx5_6,mlx5_7" + os.environ["PS_P2P_STORE_RDMA_DEVICES"] = "mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_4,mlx5_5,mlx5_6,mlx5_7" + os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:False" + + @classmethod + def tearDownClass(cls) -> None: + del os.environ["XTUNER_USE_FA3"] + del os.environ["NCCL_CUMEM_ENABLE"] + del os.environ["NCCL_IB_HCA"] + del os.environ["PS_P2P_STORE_RDMA_DEVICES"] + del os.environ["PYTORCH_CUDA_ALLOC_CONF"] + + def setUp(self): + self.model_path = MODEL_PATH + self.temp_dir = None + self.train_controller = None + self.rollout_controller = None + + def tearDown(self): + if self.train_controller is not None: + self.train_controller = None + if self.rollout_controller is not None: + ray.get(self.rollout_controller.shutdown.remote(), timeout=60) + self.rollout_controller = None + clear_cpu_resource_manager() + if ray.is_initialized(): + ray.shutdown() + if self.temp_dir is not None: + self.temp_dir.cleanup() + self.temp_dir = None + + def init_config(self, *, weight_transport_type: str): + nnodes = int(os.environ.get("WORLD_SIZE", "1")) + num_workers = int(os.environ.get("COLOCATE_NUM_WORKERS", str(8 * nnodes))) + rollout_tp_size = int(os.environ.get("ROLLOUT_TP_SIZE", "1")) + + self.resources_cfg = AcceleratorResourcesConfig( + accelerator="GPU", + num_workers=num_workers, + num_cpus_per_worker=12, + cpu_memory_per_worker=32 * 1024**3, + ) + self.rollout_cfg = RolloutConfig( + env="test_rollout", + device=self.resources_cfg.accelerator, + model_path=MODEL_PATH, + model_name=os.path.basename(MODEL_PATH).lower(), + tokenizer_path=MODEL_PATH, + rollout_cross_node_comm=False, + tensor_parallel_size=rollout_tp_size, + expert_parallel_size=2, + gpus_per_node=int(os.environ.get("GPUS_PER_NODE", "8")), + dtype="bfloat16", + skip_load_weights=False, + weight_transport_type=weight_transport_type, + checkpoint_name_prefix=f"test-update-weight-colocate-{id(self)}", + context_length=int(os.environ.get("ROLLOUT_CONTEXT_LENGTH", "10240")), + worker_log_dir=self.worker_log_dir, + gpu_memory_utilization=float(os.environ.get("ROLLOUT_GPU_MEMORY_UTILIZATION", "0.8")), + ) + + model_cfg = Qwen3_5_VLMoE35BA3Config(freeze_vision=True, freeze_projector=True) + model_cfg.text_config.mtp_config = MTPConfig(num_layers=1) + model_cfg.text_config.ep_size = 1 + + optim_cfg = AdamWConfig(lr=1e-6, foreach=False, weight_decay=0.1) + fsdp_cfg = FSDPConfig(torch_compile=False, cpu_offload=False, ep_size=1) + lr_cfg = LRConfig(lr_type="constant", warmup_ratio=0, lr_min=1e-6) + self.worker_cfg = WorkerConfig( + model_cfg=model_cfg, + load_from=MODEL_PATH, + optim_cfg=optim_cfg, + loss_cfg=LossConfig( + policy_loss_cfg=dict( + cliprange_high=0.28, + cliprange_low=0.2, + loss_type="vanilla", + clip_ratio_c=10.0, + log_prob_diff_min=-20.0, + log_prob_diff_max=20.0, + ), + ignore_idx=-100, + use_kl_loss=False, + kl_loss_coef=0.0, + kl_loss_type="low_var_kl", + mode="chunk", + chunk_size=512, + ), + lr_cfg=lr_cfg, + fsdp_cfg=fsdp_cfg, + sp_size=1, + optimizer_steps=1, + pack_max_length=int(os.environ.get("PACK_MAX_LENGTH", str(10 * 1024))), + ) + + def _setup_engines(self, *, weight_transport_type: str): + ray.init(num_cpus=128, ignore_reinit_error=True) + self.temp_dir = tempfile.TemporaryDirectory() + self.worker_log_dir = os.path.join(self.temp_dir.name, "work_dirs") + self.init_config(weight_transport_type=weight_transport_type) + self.pg = AutoAcceleratorWorkers.build_placement_group( + self.resources_cfg, + name=f"test_update_weight_colocate_{id(self)}", + ) + set_cpu_resource_manager(CPUResourceManager(accelerator_placement_groups=[self.pg])) + + TrainingWorker = ray.remote( + runtime_env={ + "env_vars": { + "RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES": "1", + "RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES": "1", + } + }, + )(BaseTrainingWorker) + train_workers, _ = AutoAcceleratorWorkers.from_placement_group(TrainingWorker, self.worker_cfg, self.pg) + ray.get([worker.test_all_reduce.remote() for worker in train_workers]) + self.train_controller = TrainingController(workers=train_workers) + self.train_controller.offload(target="all") + + self.rollout_controller = self.rollout_cfg.build(self.pg) + return self.train_controller, self.rollout_controller + + def _check_sglang_weights(self, rollout_controller, action): + targets = ray.get(rollout_controller.get_weight_update_targets.remote()) + active_urls = [target.server_url for target in targets if target.is_active] + self.assertGreater(len(active_urls), 0) + results = [] + for url in active_urls: + response = requests.post( + f"{url}/weights_checker", + json={"action": action}, + timeout=300, + ) + response.raise_for_status() + results.append(response.json()) + return results + + @unittest.skip("skip sglang parameter-only weight check test until the parameter-check-only patch is applied") + def test_sglang_colocate_ipc_update_weight(self): + train_controller, rollout_controller = self._setup_engines(weight_transport_type='ipc') + + self._check_sglang_weights(rollout_controller, action="snapshot_parameters") + self._check_sglang_weights(rollout_controller, action="reset_parameters") + + targets = ray.get(rollout_controller.get_weight_update_targets.remote()) + train_controller.bind_rollout_weight_update( + targets=targets, + rollout_config=self.rollout_cfg, + ) + + ray.get(rollout_controller.offload.remote(), timeout=300) + ray.get(self.rollout_controller.onload_weights.remote(), timeout=300) + train_controller.onload(target="model") + train_controller.weight_update() + + self._check_sglang_weights(rollout_controller, action="compare_parameters") + + + @unittest.skip("skip sglang parameter-only weight check test until the parameter-check-only patch is applied") + def test_sglang_colocate_checkpoint_engine_update_weight_train_register(self): + train_controller, rollout_controller = self._setup_engines(weight_transport_type="checkpoint_engine") + + self._check_sglang_weights(rollout_controller, action="snapshot_parameters") + self._check_sglang_weights(rollout_controller, action="reset_parameters") + + targets = ray.get(rollout_controller.get_weight_update_targets.remote()) + train_controller.bind_rollout_weight_update( + targets=targets, + rollout_config=self.rollout_cfg, + ) + ray.get(rollout_controller.offload.remote(), timeout=300) + train_controller.onload(target="model") + train_controller.weight_update(need_register=True, need_update=False) + train_controller.offload(target="model") + ray.get(self.rollout_controller.onload_weights.remote(), timeout=300) + train_controller.weight_update(need_register=False, need_update=True) + + self._check_sglang_weights(rollout_controller, action="compare_parameters") + +if __name__ == "__main__": + unittest.main() diff --git a/tests/rl/test_update_weight_disaggregated.py b/tests/rl/test_update_weight_disaggregated.py index 7eea959779..850ad0610a 100644 --- a/tests/rl/test_update_weight_disaggregated.py +++ b/tests/rl/test_update_weight_disaggregated.py @@ -85,6 +85,7 @@ def init_config(self): expert_parallel_size=1, gpus_per_node=int(os.environ.get("GPUS_PER_NODE", "8")), dtype="bfloat16", + weight_transport_type="nccl", skip_load_weights=True, context_length=256, worker_log_dir=self.worker_log_dir, @@ -159,9 +160,8 @@ def test_sglang_disaggregated_update_weight_and_generate(self): train_controller.bind_rollout_weight_update( targets=targets, rollout_config=self.rollout_cfg, - weight_transport_type="nccl", ) - train_controller.update_weights() + train_controller.weight_update() res_update_weight = ray.get(rollout_controller.generate.remote(rollout_state=input_state)) self.assertEqual(res_update_weight.response, res_baseline.response) @@ -198,9 +198,8 @@ def test_sglang_disaggregated_update_weight_equal_after_reset(self): train_controller.bind_rollout_weight_update( targets=targets, rollout_config=self.rollout_cfg, - weight_transport_type="nccl", ) - train_controller.update_weights() + train_controller.weight_update() self._check_sglang_weights(rollout_controller, action="compare_parameters") finally: @@ -237,9 +236,8 @@ def test_lmdeploy_disaggregated_update_weight_and_generate(self): train_controller.bind_rollout_weight_update( targets=targets, rollout_config=self.rollout_cfg, - weight_transport_type="nccl", ) - train_controller.update_weights() + train_controller.weight_update() res_update_weight = ray.get(rollout_controller.generate.remote(rollout_state=input_state)) self.assertEqual(res_update_weight.response, res_baseline.response) diff --git a/xtuner/v1/profiler/cuda_profile.py b/xtuner/v1/profiler/cuda_profile.py index 9c1aab4754..f2363419e1 100644 --- a/xtuner/v1/profiler/cuda_profile.py +++ b/xtuner/v1/profiler/cuda_profile.py @@ -91,6 +91,12 @@ def __init__(self, profile_dir: Path): torch.cuda.memory._record_memory_history(max_entries=MEMORY_SNAPSHOT_MAX_ENTRIES, stacks="python") self.profile_dir = profile_dir + def close(self): + try: + torch.cuda.memory._record_memory_history(enabled=None) + except TypeError: + torch.cuda.memory._record_memory_history(False) + def step(self, exit_ctx: bool = False): if dist.is_initialized(): rank = torch.distributed.get_rank() @@ -125,8 +131,15 @@ def profiling_memory(profile_dir: Path): yield return profiler = MemoryProfiler(profile_dir) - yield try: - profiler.step(exit_ctx=False) - except torch.OutOfMemoryError: - profiler.step(exit_ctx=True) + try: + yield + except torch.OutOfMemoryError: + profiler.step(exit_ctx=True) + raise + else: + profiler.step(exit_ctx=False) + finally: + # close() flushes and releases profiler resources on all exit paths. + # Without it, consecutive memory profiling runs may overwrite each other's output. + profiler.close() diff --git a/xtuner/v1/rl/rollout/worker.py b/xtuner/v1/rl/rollout/worker.py index faa6abbc35..6e11ba2b81 100644 --- a/xtuner/v1/rl/rollout/worker.py +++ b/xtuner/v1/rl/rollout/worker.py @@ -108,6 +108,14 @@ class RolloutConfig(BaseModel): 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. + weight_transport_type (Optional[str]): Transport path used to update rollout weights from training workers. + Supported values are "ipc" and "checkpoint_engine" for colocated mode, "nccl" for disaggregated mode. If not set, "ipc" will use in colocate and "nccl" will use in disaggregated. "checkpoint_engine" is currently supported only by the SGLang rollout backend in colocated mode. Defaults to None. + 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. + Defaults to 300.0. + checkpoint_engine_sync_after_register (bool): Whether to explicitly synchronize the accelerator after + registering a checkpoint into Checkpoint Engine. Defaults to False. rollout_max_batch_size_per_instance (int): Maximum batch size for the rollout worker. If not set, it will be determined automatically based on `context_length`. Defaults to 512. allow_over_concurrency_ratio (float): Deprecated compatibility option. Rollout runtime concurrency is @@ -187,6 +195,18 @@ class RolloutConfig(BaseModel): help="Base port number for distributed communication among rollout workers.", ), ] = 25000 + weight_transport_type: Annotated[ + Optional[str], + Parameter( + group=infer_group, + help=( + "Transport path used to update rollout weights from training workers. " + "Supported values: 'ipc' for colocated in-process transfer, 'nccl' for disaggregated GPU-to-GPU " + "transfer, and 'checkpoint_engine' for SGLang rollout workers that fetch sharded checkpoints " + "from Checkpoint Engine ParameterServer instances. Defaults to None." + ), + ), + ] = None weight_update_host: Annotated[ Optional[str], Parameter( @@ -207,6 +227,30 @@ class RolloutConfig(BaseModel): ), ), ] = 30000 + checkpoint_name_prefix: Annotated[ + str, + Parameter( + group=infer_group, + help="Prefix used for Checkpoint Engine checkpoint names.", + ), + ] = "xtuner-rl" + checkpoint_engine_timeout: Annotated[ + float, + Parameter( + group=infer_group, + help="Timeout in seconds for Checkpoint Engine rollout weight update requests.", + ), + ] = 300.0 + checkpoint_engine_sync_after_register: Annotated[ + bool, + Parameter( + group=infer_group, + help=( + "Whether to explicitly synchronize the accelerator after registering a checkpoint into " + "Checkpoint Engine." + ), + ), + ] = False rollout_max_batch_size_per_instance: Annotated[ Optional[int], Parameter( diff --git a/xtuner/v1/rl/trainer/controller.py b/xtuner/v1/rl/trainer/controller.py index 0d15f61ff0..3e005b1d22 100644 --- a/xtuner/v1/rl/trainer/controller.py +++ b/xtuner/v1/rl/trainer/controller.py @@ -321,26 +321,20 @@ def bind_rollout_weight_update( *, targets, rollout_config, - weight_transport_type, - weight_update_host=None, - weight_update_port=None, ): ray.get( [ worker.bind_rollout_weight_update.remote( targets=targets, rollout_config=rollout_config, - weight_transport_type=weight_transport_type, - weight_update_host=weight_update_host, - weight_update_port=weight_update_port, ) for worker in self.workers ] ) - def update_weights(self): - """Update the weights of the training workers.""" - handles = [worker.update_weights.remote() for worker in self.workers] + def weight_update(self, **kwargs): + """Update the weights from the training workers.""" + handles = [worker.weight_update.remote(**kwargs) for worker in self.workers] ray.get(handles, timeout=TRAIN_RAY_GET_TIMEOUT) return diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index 7d1dbc7ce0..05475174c4 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -46,7 +46,7 @@ from xtuner.v1.profiler import profiling_memory, profiling_time from xtuner.v1.rl.loss import BaseRLLossConfig, BaseRLLossContext, finalize_train_policy_metrics, kl_penalty from xtuner.v1.rl.utils import SingleAcceleratorWorker -from xtuner.v1.rl.weight_update import UpdateWeighter +from xtuner.v1.rl.weight_update import WeightUpdater from xtuner.v1.train.trainer import LoadCheckpointConfig from xtuner.v1.utils import ( XTUNER_DETERMINISTIC, @@ -274,7 +274,7 @@ def __init__( if hasattr(worker_cfg.model_cfg.text_config, "mtp_config"): self.mtp_config = worker_cfg.model_cfg.text_config.mtp_config - self.update_weighter = UpdateWeighter( + self.update_weighter = WeightUpdater( rank=self.rank, logger=self.logger, config=self.config, @@ -325,8 +325,8 @@ def bind_rollout_weight_update(self, *args, **kwargs): return self.update_weighter.bind_rollout_weight_update(*args, **kwargs) @ray_method - def update_weights(self): - return self.update_weighter.update_weights() + def weight_update(self, **kwargs): + return self.update_weighter.weight_update(**kwargs) def _init_sft(self, worker_cfg: WorkerConfig): self._sft_dataloader_config = worker_cfg.sft_dataloader_cfg diff --git a/xtuner/v1/rl/utils/ray_utils.py b/xtuner/v1/rl/utils/ray_utils.py index 2e02739806..87120ed642 100644 --- a/xtuner/v1/rl/utils/ray_utils.py +++ b/xtuner/v1/rl/utils/ray_utils.py @@ -160,9 +160,6 @@ def bind_train_rollout( train_workers, rollout_controller, rollout_config, - weight_transport_type, - weight_update_host=None, - weight_update_port=None, ) -> None: """Bind the training and rollout workers for updating weights. @@ -180,9 +177,6 @@ def bind_train_rollout( worker.bind_rollout_weight_update.remote( targets=targets, rollout_config=rollout_config, - weight_transport_type=weight_transport_type, - weight_update_host=weight_update_host, - weight_update_port=weight_update_port, ) for worker in train_workers ] diff --git a/xtuner/v1/rl/weight_update/__init__.py b/xtuner/v1/rl/weight_update/__init__.py index e2268c11fc..7e509a4066 100644 --- a/xtuner/v1/rl/weight_update/__init__.py +++ b/xtuner/v1/rl/weight_update/__init__.py @@ -6,6 +6,7 @@ WeightUpdateBatch, ) from .transport import ( + CheckpointEngineWeightTransport, IPCBackendAdapter, IPCWeightTransport, LMDeployIPCBackendAdapter, @@ -16,11 +17,12 @@ WeightTransport, WeightUpdateRequest, ) -from .update_weighter import UpdateWeighter +from .update_weighter import WeightUpdater from .weight_iterator import WeightIterator __all__ = [ + "CheckpointEngineWeightTransport", "IPCBackendAdapter", "IPCWeightTransport", "LMDeployIPCBackendAdapter", @@ -31,7 +33,7 @@ "RolloutWeightUpdateInfo", "SGLangIPCBackendAdapter", "SGLangNCCLBackendAdapter", - "UpdateWeighter", + "WeightUpdater", "WeightIterator", "WeightTransportType", "WeightUpdateBatch", diff --git a/xtuner/v1/rl/weight_update/data.py b/xtuner/v1/rl/weight_update/data.py index 08007b5a4e..232489dea1 100644 --- a/xtuner/v1/rl/weight_update/data.py +++ b/xtuner/v1/rl/weight_update/data.py @@ -12,7 +12,7 @@ RolloutBackend: TypeAlias = Literal["sglang", "vllm", "pytorch", "turbomind"] # Rollout inference backend. -WeightTransportType: TypeAlias = Literal["ipc", "nccl"] # Supported weight transport types. +WeightTransportType: TypeAlias = Literal["ipc", "nccl", "checkpoint_engine"] # Supported weight transport types. def _resolve_rollout_backend(rollout_config: RolloutConfig) -> RolloutBackend: @@ -40,11 +40,18 @@ def _validate_transport_type( assert weight_transport_type is not None, "bind_rollout_weight_update() must set weight_transport_type." transport_type = weight_transport_type.lower() - if transport_type not in ("ipc", "nccl"): - raise ValueError(f"Unsupported weight_transport_type: {weight_transport_type!r}. Expected 'ipc' or 'nccl'.") + if transport_type not in ("ipc", "nccl", "checkpoint_engine"): + raise ValueError( + f"Unsupported weight_transport_type: {weight_transport_type!r}. " + "Expected 'ipc', 'nccl' or 'checkpoint_engine'." + ) transport_type = cast(WeightTransportType, transport_type) if transport_type == "nccl" and backend in ("vllm", "turbomind"): raise NotImplementedError(f"NCCL weight transport is not supported for {backend} backend.") + if transport_type == "checkpoint_engine" and backend != "sglang": + raise NotImplementedError( + f"Checkpoint Engine weight transport currently only supports sglang, got backend={backend!r}." + ) return transport_type @@ -86,6 +93,12 @@ class RolloutWeightUpdateInfo: weight_update_host: str | None = None # Optional port used by NCCL external weight update groups. weight_update_port: int | None = None + # Optional prefix used by checkpoint-engine + checkpoint_name_prefix: str | None = None + # Optional timeout used by checkpoint-engine + checkpoint_engine_timeout: float | None = None + # Whether to explicitly synchronize after registering checkpoint-engine tensors. + checkpoint_engine_sync_after_register: bool = False @classmethod def from_targets( @@ -94,16 +107,16 @@ def from_targets( rollout_config: RolloutConfig, weight_update_targets: tuple[RolloutWeightUpdateTarget, ...], train_rank: int, - weight_transport_type: WeightTransportType | str, - weight_update_host: str | None = None, - weight_update_port: int | None = None, ) -> RolloutWeightUpdateInfo: backend = _resolve_rollout_backend(rollout_config) tp = rollout_config.tensor_parallel_size ep = rollout_config.expert_parallel_size assert tp == 1 or ep == 1, "Either tensor parallel size or engine parallel size must be 1." + transport_type = rollout_config.weight_transport_type + if transport_type is None: + raise ValueError("rollout_config.weight_transport_type should be set in RL training") transport_type = _validate_transport_type( - weight_transport_type=weight_transport_type, + weight_transport_type=transport_type, backend=backend, ) return cls( @@ -112,8 +125,13 @@ def from_targets( train_rank=train_rank, transport_type=transport_type, backend=backend, - weight_update_host=weight_update_host, - weight_update_port=weight_update_port if weight_update_port is not None else 30000, + weight_update_host=rollout_config.weight_update_host, + weight_update_port=rollout_config.weight_update_port + if rollout_config.weight_update_port is not None + else 30000, + checkpoint_name_prefix=rollout_config.checkpoint_name_prefix, + checkpoint_engine_timeout=rollout_config.checkpoint_engine_timeout, + checkpoint_engine_sync_after_register=rollout_config.checkpoint_engine_sync_after_register, ) @property diff --git a/xtuner/v1/rl/weight_update/transport.py b/xtuner/v1/rl/weight_update/transport.py index 637da86486..d961455ba2 100644 --- a/xtuner/v1/rl/weight_update/transport.py +++ b/xtuner/v1/rl/weight_update/transport.py @@ -8,7 +8,8 @@ from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from datetime import timedelta -from typing import Any, Callable, Protocol, cast +from pathlib import Path +from typing import Any, Callable, Generic, Protocol, Sequence, TypeVar, cast import requests import torch @@ -31,9 +32,14 @@ monkey_unpatch_torch_reductions, ) -from .data import RolloutWeightUpdateInfo, WeightUpdateBatch +from .data import RolloutWeightUpdateInfo, RolloutWeightUpdateTarget, WeightUpdateBatch +try: + from checkpoint_engine.ps import ParameterServer +except ImportError: + ParameterServer = None + DEVICE = get_device() DEVICE_MODULE = get_torch_device_module() @@ -52,7 +58,10 @@ def before_update(self) -> None: ... def after_update_all_groups(self) -> None: ... -class WeightTransport(ABC): +AdapterT = TypeVar("AdapterT", bound=WeightTransportAdapter) + + +class WeightTransport(ABC, Generic[AdapterT]): def __init__(self, *, rollout_info: RolloutWeightUpdateInfo, logger: Any, rank: int): self.rollout_info = rollout_info self.logger = logger @@ -60,7 +69,7 @@ def __init__(self, *, rollout_info: RolloutWeightUpdateInfo, logger: Any, rank: self.backend = self.rollout_info.backend self.rollout_ep = self.rollout_info.ep self.rollout_tp = self.rollout_info.tp - self._adapter: WeightTransportAdapter | None = None + self._adapter: AdapterT | None = None self.rollout_url = self.rollout_info.rollout_url @@ -75,7 +84,7 @@ def post_json(url: str, endpoint: str, payload: dict, *, api_key=None) -> dict: response.raise_for_status() return response.json() - def update(self, weight_iterator: Any) -> None: + def update(self, weight_iterator: Any, **_: Any) -> None: assert self._adapter is not None self._adapter.before_update() DEVICE_MODULE.empty_cache() @@ -417,7 +426,7 @@ def after_update_per_batch( dist.barrier(group=cpu_group) -class IPCWeightTransport(WeightTransport): +class IPCWeightTransport(WeightTransport[IPCBackendAdapter]): _adapter: IPCBackendAdapter def __init__( @@ -445,7 +454,8 @@ def _build_adapter(self) -> IPCBackendAdapter: if self.backend == "vllm": return VLLMIPCBackendAdapter(rollout_tp=self.rollout_info.tp) elif self.backend == "sglang": - return SGLangIPCBackendAdapter(rollout_tp=self.rollout_info.tp) + tp_size = self.rollout_info.tp if self.rollout_info.tp > 1 else self.rollout_info.ep + return SGLangIPCBackendAdapter(rollout_tp=tp_size) elif self.backend == "pytorch" or self.backend == "turbomind": return LMDeployIPCBackendAdapter( rollout_tp=self.rollout_info.tp, @@ -588,7 +598,7 @@ def build_weight_update_payload(self, batch: WeightUpdateBatch, group_name: str) return payload, flattened_tensor, weight_names else: # finalize-only request: no tensors to broadcast, just trigger the - # rollout side's mod.update_weights() finalization hooks. + # rollout side's mod.weight_update() finalization hooks. payload = { "names": [], "dtypes": [], @@ -606,7 +616,7 @@ def build_request( return WeightUpdateRequest(endpoint="update_weights_from_distributed", body=payload) -class NCCLWeightTransport(WeightTransport): +class NCCLWeightTransport(WeightTransport[NCCLBackendAdapter]): _adapter: NCCLBackendAdapter def __init__(self, *, rank: int, logger: Any, rollout_info: RolloutWeightUpdateInfo): @@ -838,3 +848,299 @@ def teardown(self) -> None: self.group_name = None self.engine_urls = [] self.external_group_world_size = None + + +class CheckpointEngineAdapter: + """Build adapter for CheckpointEngine.""" + + def __init__(self, *, rank: int): + self.rank = rank + + def build_update_url(self, server_url: str) -> str: + return f"{server_url.rstrip('/')}/update_weights_from_ipc" + + +class CheckpointEngineWeightTransport: + """In-process Checkpoint Engine transport via ParameterServer. + + Each train rank owns a PS and collectively register / gather_metas / update. + """ + + _adapter: CheckpointEngineAdapter + + def __init__( + self, + *, + rank: int, + logger: Any, + rollout_info: RolloutWeightUpdateInfo, + ): + """Build PS and split weight keys in HF json file.""" + + self.rollout_info = rollout_info + self.logger = logger + self.rank = rank + self.backend = self.rollout_info.backend + self.rollout_ep = self.rollout_info.ep + self.rollout_tp = self.rollout_info.tp + self.rollout_url = self.rollout_info.rollout_url + + os.environ["NCCL_IB_HCA"] = "mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_4,mlx5_5,mlx5_6,mlx5_7" + os.environ["PS_P2P_STORE_RDMA_DEVICES"] = "mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_4,mlx5_5,mlx5_6,mlx5_7" + + assert dist.is_initialized(), "Checkpoint Engine requires an initialized torch.distributed process group." + self.ps_world_size = dist.get_world_size() + + self._adapter = CheckpointEngineAdapter(rank=rank) + # record the update counter of PS + self._update_counter = 0 + self._checkpoint_path = self.rollout_info.rollout_config.model_path + self._checkpoint_name_prefix = self.rollout_info.rollout_config.checkpoint_name_prefix + self._timeout = self.rollout_info.rollout_config.checkpoint_engine_timeout + self._sync_after_register = self.rollout_info.checkpoint_engine_sync_after_register + # record the checkpoint name of PS and will use it to register the checkpoint and unregister the previous checkpoint + self._checkpoint_name: str | None = None + + self._ps = self.build_parameter_server() + # record the local checkpoint keys per PS-rank + self._local_checkpoint_keys = self.split_tensors_for_rank(self._checkpoint_path, self.ps_world_size, self.rank) + + def build_parameter_server(self): + """Build the Checkpoint Engine ParameterServer. + + The world size is the same as the train world size. The ParameterServer is built with auto_pg=False, so the + torch.distributed process group is not destroyed after update. + """ + if ParameterServer is None: + raise ImportError("Checkpoint Engine is not available. Please install the Checkpoint Engine package.") + + # auto_pg=False keeps the default PG alive. + ps = ParameterServer( + auto_pg=False, + rank=self.rank, + world_size=self.ps_world_size, + ) + self.logger.info(f"[checkpoint_engine] ParameterServer ready rank={self.rank} world_size={self.ps_world_size}") + return ps + + def split_tensors_for_rank(self, checkpoint_path: str | Path, world_size: int, rank: int) -> set[str]: + """Split an HF keys for each ParameterServer.""" + + path = Path(checkpoint_path) + index_path = path / "model.safetensors.index.json" + + if not index_path.exists(): + raise FileNotFoundError(f"model.safetensors.index.json file not found: {index_path}") + # TODO: split tensor according to tenser size of header metadata in .safetensors + with open(index_path) as f: + weight_map: dict[str, str] = json.load(f)["weight_map"] + weight_keys = [key_name for key_name, file_name in weight_map.items()] + local_keys = set(weight_keys[rank::world_size]) + + self.logger.info( + f"[checkpoint_engine] split keys from {index_path} " + f"rank={rank} tensors={len(local_keys)}/{len(weight_keys)}" + ) + return local_keys + + def _collect_named_tensors(self, weight_iterator, local_keys=None): + """Collect all train weights from the iterator onto CPU.""" + named = {} + named_total_bytes = 0 + for batches in weight_iterator.iter_batch_groups(): + for batch in batches: + sd = batch.state_dict + if not sd: + continue + # batch 级快路径:完全无交集可跳过 + if local_keys is not None and sd.keys().isdisjoint(local_keys): + sd.clear() + del sd, batch + DEVICE_MODULE.empty_cache() + continue + for key, tensor in list(sd.items()): + if local_keys is not None and key not in local_keys: + sd.pop(key) + del tensor + continue + # 占显存更少,但速度慢 + # named[key] = tensor.detach().to("cpu", non_blocking=True) + # 占显存多,速度快 + named[key] = tensor + named_total_bytes += named[key].numel() * named[key].element_size() + sd.clear() + del sd, batch + DEVICE_MODULE.empty_cache() + DEVICE_MODULE.empty_cache() + if local_keys is not None: + self.logger.info( + f"[checkpoint_engine] collect matched local keys rank={self.rank} " + f"parameter server shard total={named_total_bytes / 1024**3:.3f}GiB " + ) + return named + + def register_checkpoint_from_train_engine(self, weight_iterator: Any) -> None: + """Register current train engine weights into Checkpoint Engine PS.""" + + # 0. Drop previous train checkpoint to limit pinned host memory. + if self._checkpoint_name is not None: + self._ps.unregister_checkpoint(self._checkpoint_name) + + # 1. Collect named tensors from weight iterator + all_tensors = self._collect_named_tensors(weight_iterator, local_keys=self._local_checkpoint_keys) + + # 2. Shard tensors for the current parameter server + missing = self._local_checkpoint_keys - all_tensors.keys() + if missing: + missing_mtp_keys = {key for key in missing if key.startswith("mtp.")} + missing_non_mtp_keys = missing - missing_mtp_keys + raise RuntimeError( + f"[checkpoint_engine] ParameterServer Rank={self.rank}. Missing non-MTP keys: [{missing_non_mtp_keys}]. Missing MTP-only keys: [{missing_mtp_keys}]" + ) + + shard = {k: all_tensors[k] for k in self._local_checkpoint_keys if k in all_tensors} + + # 3. Register checkpoint + self._update_counter += 1 + name = f"{self._checkpoint_name_prefix}-train-{self._update_counter}" + self.logger.info( + f"[checkpoint_engine] register train checkpoint name={name} " + f"rank={self.rank} tensors={len(shard)}/{len(all_tensors)}" + ) + self._ps.register_checkpoint(name, files=[], named_tensors=shard, use_shared_memory_pool=True) + if self._sync_after_register: + DEVICE_MODULE.synchronize() + dist.barrier() + self._checkpoint_name = name + + def _make_req_func(self, targets: Sequence[RolloutWeightUpdateTarget]): + """Build CE ``req_func``: source rank POSTs IPC update to its SGLang + target.""" + rank = self.rank + + # Map each train rank to its SGLang engine target and group src rank. + rank_to_target: dict[int, RolloutWeightUpdateTarget] = {} + for target in targets: + for r in target.update_ranks: + if r in rank_to_target: + raise ValueError(f"Duplicate update rank {r} across active CE targets.") + rank_to_target[r] = target + adapter = self._adapter + + def req_func(socket_paths: list[tuple[str, str]]) -> None: + target = rank_to_target.get(rank) + if target is None: + return + src = min(target.update_ranks) + if rank != src: + return + update_url = adapter.build_update_url(target.server_url) + rank_to_socket = dict(enumerate(socket_paths)) + payload = { + "zmq_handles": dict(rank_to_socket[r] for r in target.update_ranks), + "flush_cache": True, + } + headers = {"Content-Type": "application/json"} + api_key = self.rollout_info.api_key + if api_key is not None: + token = api_key[0] if isinstance(api_key, list) else api_key + headers["Authorization"] = f"Bearer {token}" + response = requests.post(update_url, headers=headers, json=payload, timeout=self._timeout) + response.raise_for_status() + self.logger.info(f"[checkpoint_engine] rank{rank} updated {update_url}") + + return req_func + + def _get_target_update_ranks(self, targets: Sequence[RolloutWeightUpdateTarget], world_size: int) -> list[int]: + """Return validated update ranks from active rollout targets.""" + + covered: set[int] = set() + for target in targets: + if not target.update_ranks: + raise ValueError(f"Empty update_ranks for target {target.server_url!r}") + for r in target.update_ranks: + if r < 0 or r >= world_size: + raise ValueError( + f"PS/train rank {r} from target {target.server_url!r} out of range for world_size={world_size}." + ) + if r in covered: + raise ValueError(f"Duplicate PS rank {r} in active update targets.") + covered.add(r) + return sorted(covered) + + @staticmethod + def _can_broadcast_to_update_ranks(update_ranks: Sequence[int], world_size: int) -> bool: + """Whether the pending update ranks satisfy CE broadcast requirements. + + Checkpoint Engine uses full broadcast only when every PS/train rank + participates. Partial active ranks must be sent through p2p by passing + ``ranks`` to ``ParameterServer.update``. + """ + + return len(update_ranks) == world_size and list(update_ranks) == list(range(world_size)) + + def _update_engines(self) -> None: + """``gather_metas`` then ``update`` to push checkpoint to rollout + engines.""" + + targets = self.rollout_info.active_update_targets + if not targets: + raise RuntimeError("Checkpoint Engine found no active weight-update targets.") + update_ranks = self._get_target_update_ranks(targets, self.ps_world_size) + use_broadcast = self._can_broadcast_to_update_ranks(update_ranks, self.ps_world_size) + ranks = None if use_broadcast else update_ranks + req_func = self._make_req_func(targets) + self.logger.info( + f"[checkpoint_engine] gather_metas+update name={self._checkpoint_name} " + f"active_targets={len(targets)}/{len(self.rollout_info.weight_update_targets)} " + f"method={'broadcast' if use_broadcast else 'p2p'} ranks={ranks}" + ) + self._ps.gather_metas(self._checkpoint_name) + self._ps.update(self._checkpoint_name, req_func, ranks=ranks) + + def update(self, weight_iterator: Any, **kwargs: Any) -> None: + """Update rollout engine weights through the checkpoint parameter + server. + + The update consists of two stages: registering a checkpoint from the training + engine, and loading that checkpoint into the rollout engines. + + Parameters + ---------- + weight_iterator : Any + Iterator that provides the latest training engine weights. + need_register : bool, optional + Whether to register a new checkpoint from ``weight_iterator``. If False, + reuse the checkpoint name from the previous update. + need_update : bool, optional + Whether to load the registered checkpoint into rollout engines. Set this to + False to split registration and rollout update into separate calls, which + can reduce peak GPU memory usage under memory pressure. + """ + + need_register = kwargs.pop("need_register", True) + need_update = kwargs.pop("need_update", True) + assert need_register or need_update, ( + "At least one of need_register or need_update must be True when use checkpoint engine update." + ) + if not need_register and self._checkpoint_name is None: + raise RuntimeError("CheckpointEngineWeightTransport cannot update without a registered checkpoint.") + + # 1. Register checkpoint from train engine + if need_register: + self.register_checkpoint_from_train_engine(weight_iterator) + + # 2. Broadcast checkpoint to engines + if need_update: + self._update_engines() + + def reset_rollout_info(self, rollout_info: RolloutWeightUpdateInfo): + self.rollout_info = rollout_info + self.rollout_url = rollout_info.rollout_url + + def teardown(self) -> None: + if self._ps is None: + return + if self._checkpoint_name: + self._ps.unregister_checkpoint(self._checkpoint_name, force=True) + self._ps = None diff --git a/xtuner/v1/rl/weight_update/update_weighter.py b/xtuner/v1/rl/weight_update/update_weighter.py index db4558845e..0d619e2669 100644 --- a/xtuner/v1/rl/weight_update/update_weighter.py +++ b/xtuner/v1/rl/weight_update/update_weighter.py @@ -7,13 +7,12 @@ from .data import ( RolloutWeightUpdateInfo, RolloutWeightUpdateTarget, - WeightTransportType, ) -from .transport import IPCWeightTransport, NCCLWeightTransport, WeightTransport +from .transport import CheckpointEngineWeightTransport, IPCWeightTransport, NCCLWeightTransport from .weight_iterator import WeightIterator -class UpdateWeighter: +class WeightUpdater: def __init__(self, *, rank: int, logger: Any, config: Any, engine: Any): self.rank = rank self.logger = logger @@ -25,7 +24,7 @@ def __init__(self, *, rank: int, logger: Any, config: Any, engine: Any): self.weight_iterator: WeightIterator | None = None self._global_hf_keys_mapping_cache: dict[str, list[str]] = {} # Transport is initialized after bind_rollout_weight_update() is called. - self._transport: WeightTransport | None = None + self._transport: Any | None = None # Used to detect changes in rollout metadata that require resetting the transport. self._transport_signature: tuple[Any, ...] | None = None @@ -34,9 +33,6 @@ def bind_rollout_weight_update( *, targets: tuple[RolloutWeightUpdateTarget, ...], rollout_config: RolloutConfig, - weight_transport_type: WeightTransportType, - weight_update_host: str | None = None, - weight_update_port: int | None = None, ): """Bind this train worker to rollout weight-update targets.""" @@ -44,19 +40,21 @@ def bind_rollout_weight_update( rollout_config=rollout_config, weight_update_targets=targets, train_rank=self.rank, - weight_transport_type=weight_transport_type, - weight_update_host=weight_update_host, - weight_update_port=weight_update_port, ) - new_transport_signature = self.rollout_info.transport_signature - # Weight transports may cache resources derived from rollout metadata. - # Since rollout workers can fail and recover with new URL/status/mesh metadata, - # reset the cached transport whenever that metadata changes. - if self._transport_signature is not None and new_transport_signature != self._transport_signature: - self.logger.info("Rollout metadata changed, reset weight transport.") - self._reset_transport() - self._transport_signature = new_transport_signature + if self.rollout_info.transport_type == "checkpoint_engine" and self._transport is not None: + # When using Checkpoint Engine, the ParameterServer must not be reinitialized + # when rollout info changes. And only update rollout_info. + self._transport.reset_rollout_info(self.rollout_info) + else: + new_transport_signature = self.rollout_info.transport_signature + # Weight transports may cache resources derived from rollout metadata. + # Since rollout workers can fail and recover with new URL/status/mesh metadata, + # reset the cached transport whenever that metadata changes. + if self._transport_signature is not None and new_transport_signature != self._transport_signature: + self.logger.info("Rollout metadata changed, reset weight transport.") + self._reset_transport() + self._transport_signature = new_transport_signature self.weight_iterator = WeightIterator( config=self.config, @@ -67,16 +65,16 @@ def bind_rollout_weight_update( if self._transport is None: self._set_transport() - def update_weights(self): + def weight_update(self, **kwargs: Any) -> None: """Update the model weights.""" - assert self.rollout_info is not None, "bind_rollout_weight_update() must be called before update_weights()." + assert self.rollout_info is not None, "bind_rollout_weight_update() must be called before weight_update()." assert self._transport is not None, ( f"Weight transport is not initialized. transport_type={self.rollout_info.transport_type!r}, " f"backend={self.rollout_info.backend!r}." ) assert self.weight_iterator is not None, "Weight iterator is not initialized." - self._transport.update(self.weight_iterator) + self._transport.update(self.weight_iterator, **kwargs) def _set_transport(self) -> None: rollout_info = self.rollout_info @@ -90,6 +88,12 @@ def _set_transport(self) -> None: ) elif rollout_info.transport_type == "nccl": self._transport = NCCLWeightTransport(rank=self.rank, logger=self.logger, rollout_info=rollout_info) + elif rollout_info.transport_type == "checkpoint_engine": + self._transport = CheckpointEngineWeightTransport( + rank=self.rank, + logger=self.logger, + rollout_info=rollout_info, + ) else: raise NotImplementedError diff --git a/xtuner/v1/rl/weight_update/weight_iterator.py b/xtuner/v1/rl/weight_update/weight_iterator.py index 55f236ba11..36806f3f6e 100644 --- a/xtuner/v1/rl/weight_update/weight_iterator.py +++ b/xtuner/v1/rl/weight_update/weight_iterator.py @@ -190,7 +190,9 @@ def iter_hf_batches(self, submodule=None, final_update=False): ) train_enable_ep = model.fsdp_config is not None and model.fsdp_config.ep_size > 1 - should_gather_train_ep_shards = self.rollout_info.transport_type == "nccl" and train_enable_ep + should_gather_train_ep_shards = ( + self.rollout_info.transport_type == "nccl" and train_enable_ep + ) or self.rollout_info.transport_type == "checkpoint_engine" if train_enable_ep: if self.rollout_info.transport_type == "ipc" and self.rollout_info.ep > 1: diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 9e17bf0de7..13a322ab20 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -53,7 +53,6 @@ set_cpu_resource_manager, sort_rollout_state_for_deterministic, ) -from xtuner.v1.rl.weight_update.data import WeightTransportType from xtuner.v1.train.trainer import LoadCheckpointConfig, XTunerMeta from xtuner.v1.utils import XTUNER_DETERMINISTIC, get_logger, is_hf_model_path, set_deterministic, timer from xtuner.v1.utils.device import get_device, get_torch_device_module @@ -121,9 +120,6 @@ def bind_train_rollout( train_controller: TrainingController, rollout_controller: RolloutControllerProxy, rollout_config: RolloutConfig, - weight_transport_type: WeightTransportType | str, - weight_update_host: str | None = None, - weight_update_port: int | None = None, ) -> None: """Bind the training and rollout workers for update weights.""" targets = ray.get( @@ -133,9 +129,6 @@ def bind_train_rollout( train_controller.bind_rollout_weight_update( targets=targets, rollout_config=rollout_config, - weight_transport_type=weight_transport_type, - weight_update_host=weight_update_host, - weight_update_port=weight_update_port, ) return @@ -1611,11 +1604,13 @@ def __init__(self, cfg: RLColocateTrainerConfig): self.train_controller.offload(target="all") self.rollout_controller = self._rollout_config.build(self._pg) + if self._rollout_config.weight_transport_type is None: + self._rollout_config.weight_transport_type = "ipc" + bind_train_rollout( train_controller=self.train_controller, rollout_controller=self.rollout_controller, rollout_config=self._rollout_config, - weight_transport_type="ipc", ) replay_buffer = cfg.replay_buffer_config.build() @@ -1629,14 +1624,25 @@ def __init__(self, cfg: RLColocateTrainerConfig): self._sync_weights_from_train_workers() def _sync_weights_from_train_workers(self) -> None: - self.logger.info("Rollout workers skip load weights, update weights from train workers.") - ray.get(self.rollout_controller.offload.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) - self.train_controller.onload(target="model") - ray.get(self.rollout_controller.onload_weights.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) - self.train_controller.update_weights() - self.train_controller.offload(target="model") - ray.get(self.rollout_controller.onload_kvcache.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) - self.logger.info("Rollout workers updated weights from train workers.") + if self._rollout_config.weight_transport_type == "checkpoint_engine": + ray.get(self.rollout_controller.offload.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + self.train_controller.onload(target="model") + self.train_controller.weight_update(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.weight_update(need_register=False, need_update=True) + ray.get(self.rollout_controller.onload_kvcache.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + self.logger.info("Rollout workers updated weights from Checkpoint Engine.") + return + else: + self.logger.info("Rollout workers skip load weights, update weights from train workers.") + ray.get(self.rollout_controller.offload.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + self.train_controller.onload(target="model") + ray.get(self.rollout_controller.onload_weights.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + self.train_controller.weight_update() + self.train_controller.offload(target="model") + ray.get(self.rollout_controller.onload_kvcache.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + self.logger.info("Rollout workers updated weights from train workers.") def fit(self): try: @@ -1777,15 +1783,25 @@ def _sync_weights_and_save(self, train_step: int, step_timer_dict: dict) -> bool train_controller=self.train_controller, rollout_controller=self.rollout_controller, rollout_config=self._rollout_config, - weight_transport_type="ipc", - ) - ray.get( - self.rollout_controller.onload_weights.remote(), - timeout=RL_TRAINER_RAY_GET_TIMEOUT, ) - self.train_controller.update_weights() + + if self._rollout_config.weight_transport_type == "checkpoint_engine": + self.train_controller.weight_update(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.weight_update(need_register=False, need_update=True) + + else: + ray.get( + self.rollout_controller.onload_weights.remote(), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + self.train_controller.weight_update() + self.train_controller.offload(target="model") self.logger.info("Rollout workers update weights successfully in colocate mode") - self.train_controller.offload(target="model") suspend_train_nccl = ( os.getenv( "XTUNER_SUSPEND_TRAIN_NCCL_AFTER_SYNC", @@ -1824,6 +1840,11 @@ def __init__(self, cfg: RLDisaggregatedTrainerConfig): self.train_controller = self._train_worker_cfg.build(self._train_pg) self.rollout_controller = self._rollout_config.build(self._rollout_pg) + if self._rollout_config.weight_transport_type != "nccl": + self.logger.warning( + "Currently, disaggregated mode only support nccl as weight update transport. It will use NCCL for weight transport." + ) + self._rollout_config.weight_transport_type = "nccl" replay_buffer = cfg.replay_buffer_config.build() self._build_agent_loop_components(cfg, replay_buffer) # 非共卡 producer 不允许早停,否则 consumer 可能永久等不到 batch。 @@ -1833,13 +1854,11 @@ def __init__(self, cfg: RLDisaggregatedTrainerConfig): "In disaggregated mode, should_continue_fn must be default, " "because it does not allow early stopping in production." ) + bind_train_rollout( train_controller=self.train_controller, rollout_controller=self.rollout_controller, rollout_config=self._rollout_config, - weight_transport_type="nccl", - weight_update_host=self._rollout_config.weight_update_host, - weight_update_port=self._rollout_config.weight_update_port, ) if self._load_checkpoint_cfg.checkpoint_path is not None: @@ -1877,7 +1896,7 @@ def _resume_from_checkpoint(self, checkpoint_path: Path | str) -> None: saved_model_step = asyncio_run(self._resume_agent_loop_manager(checkpoint_path)) assert self._cur_step == saved_model_step - self.update_weights() + self.weight_update() asyncio_run(self.agent_loop_manager.continue_produce(model_step=saved_model_step)) def fit(self): @@ -2033,13 +2052,10 @@ async def _sync_weights_and_save(self, model_step: int, step_timer_dict: dict): train_controller=self.train_controller, rollout_controller=self.rollout_controller, rollout_config=self._rollout_config, - weight_transport_type="nccl", - weight_update_host=self._rollout_config.weight_update_host, - weight_update_port=self._rollout_config.weight_update_port, ) - self.update_weights() + self.weight_update() - def update_weights(self): + def weight_update(self): # rollout 恢复由 AgentLoopManager 控制。 - self.train_controller.update_weights() + self.train_controller.weight_update() self.logger.info("Rollout workers update weights successfully in disaggregated mode")