From 9a82cf0ba962982aeb77703dbf19644bdbabf06f Mon Sep 17 00:00:00 2001 From: Pengcheng Shi Date: Thu, 24 Sep 2026 10:44:04 +0800 Subject: [PATCH 1/2] fix IPC barrier --- tests/rl/test_update_weight_colocate.py | 153 +++++++++++++++++-- tests/rl/test_update_weight_disaggregated.py | 100 ++++++++++-- xtuner/v1/rl/weight_update/transport.py | 3 +- 3 files changed, 230 insertions(+), 26 deletions(-) diff --git a/tests/rl/test_update_weight_colocate.py b/tests/rl/test_update_weight_colocate.py index f2bc0a7ff5..04adc25500 100644 --- a/tests/rl/test_update_weight_colocate.py +++ b/tests/rl/test_update_weight_colocate.py @@ -1,17 +1,33 @@ -# 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. +"""共卡部署下的模型权重更新正确性测试。 + +覆盖 IPC 与 checkpoint-engine:训练侧权重写入推理引擎后,检查引擎状态是否与 +更新前一致。两类用例互补,不是重复。 + +Generate 检查(默认启用) + greedy generate 两次,确认引擎自身可复现;再走 IPC 更新后第三次 generate。 + 比较 response / response_ids / sampled-token logprobs。 + SGLang greedy 用 temperature=0;LMDeploy /generate 用 top_k=1 且 temperature=1.0。 + - test_sglang_colocate_ipc_update_weight_and_generate + - test_lmdeploy_colocate_ipc_update_weight_and_generate + +参数逐点检查(当前 skip,依赖 SGLang WeightChecker patch) + snapshot_parameters -> reset_parameters -> update_weights -> compare_parameters。 + 不 generate,直接比对引擎参数。WeightChecker: + https://github.com/PengchengShi00/sglang/commit/05e89d63b5a1a80671b267ff4494ad950b2aba75 + - test_sglang_colocate_ipc_update_weight + - test_sglang_colocate_checkpoint_engine_update_weight_train_register +""" import os import tempfile import unittest +import numpy as np import ray import requests from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status from xtuner.v1.model import Qwen3_5_VLMoE35BA3Config from xtuner.v1.module.mtp import MTPConfig @@ -31,6 +47,28 @@ ) from xtuner.v1.rl.rollout.worker_registry import WorkerLifecycleState +RL_TRAINER_RAY_GET_TIMEOUT = 3600 +TEST_TEXT_MESSAGES = [{"role": "user", "content": "Hello!"}] +# Tokens can still match when a few shards are wrong; sampled-token logprobs +# catch that. Keep this tight: identical weights should stay near exact. +GENERATE_LOGPROB_RTOL = 1e-5 +GENERATE_LOGPROB_ATOL = 1e-5 +# SGLang: temperature=0 is greedy. LMDeploy /generate always sets do_sample=True +# and divides logits by temperature, so greedy is top_k=1 with temperature=1.0. +SGLANG_GREEDY_SAMPLE_PARAMS = SampleParams( + temperature=0.0, + max_tokens=128, + top_k=1, + return_logprob=True, + return_token_ids=True, +) +LMDEPLOY_GREEDY_SAMPLE_PARAMS = SampleParams( + temperature=1.0, + max_tokens=128, + top_k=1, + return_logprob=True, + return_token_ids=True, +) MODEL_PATH = os.environ["QWEN3_5_MOE_PATH"] @@ -43,8 +81,12 @@ def setUpClass(cls) -> None: 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" - + if os.environ.get("XTUNER_USE_SGLANG", "0") == "1": + os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:False" + elif os.environ.get("XTUNER_USE_LMDEPLOY", "0") == "1": + os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" + else: + raise unittest.SkipTest("XTUNER_USE_SGLANG or XTUNER_USE_LMDEPLOY is not set") @classmethod def tearDownClass(cls) -> None: del os.environ["XTUNER_USE_FA3"] @@ -72,10 +114,11 @@ def tearDown(self): self.temp_dir.cleanup() self.temp_dir = None - def init_config(self, *, weight_transport_type: str): + def init_config(self, *, weight_transport_type: str, extra_rollout_config: dict | None = None): 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")) + rollout_tp_size = 4 + rollout_ep_size = 1 self.resources_cfg = AcceleratorResourcesConfig( accelerator="GPU", @@ -91,7 +134,7 @@ def init_config(self, *, weight_transport_type: str): tokenizer_path=MODEL_PATH, rollout_cross_node_comm=False, tensor_parallel_size=rollout_tp_size, - expert_parallel_size=2, + expert_parallel_size=rollout_ep_size, gpus_per_node=int(os.environ.get("GPUS_PER_NODE", "8")), dtype="bfloat16", skip_load_weights=False, @@ -100,6 +143,7 @@ def init_config(self, *, weight_transport_type: str): 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")), + extra_rollout_config=extra_rollout_config or {}, ) model_cfg = Qwen3_5_VLMoE35BA3Config(freeze_vision=True, freeze_projector=True) @@ -136,11 +180,14 @@ def init_config(self, *, weight_transport_type: str): pack_max_length=int(os.environ.get("PACK_MAX_LENGTH", str(10 * 1024))), ) - def _setup_engines(self, *, weight_transport_type: str): + def _setup_engines(self, *, weight_transport_type: str, extra_rollout_config: dict | None = None): 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.init_config( + weight_transport_type=weight_transport_type, + extra_rollout_config=extra_rollout_config, + ) self.pg = AutoAcceleratorWorkers.build_placement_group( self.resources_cfg, name=f"test_update_weight_colocate_{id(self)}", @@ -178,6 +225,88 @@ def _check_sglang_weights(self, rollout_controller, action): results.append(response.json()) return results + def _assert_generate_outputs_match( + self, + actual: RolloutState, + expected: RolloutState, + err_msg: str, + ) -> None: + self.assertEqual(actual.status, Status.COMPLETED, actual.error_msg) + self.assertEqual(actual.response, expected.response) + self.assertEqual(actual.response_ids, expected.response_ids) + self.assertIsNotNone(expected.logprobs) + self.assertIsNotNone(actual.logprobs) + self.assertGreater(len(expected.logprobs), 0) + self.assertEqual(len(actual.logprobs), len(expected.logprobs)) + self.assertEqual(len(actual.logprobs), len(actual.response_ids or [])) + np.testing.assert_allclose( + actual.logprobs, + expected.logprobs, + rtol=GENERATE_LOGPROB_RTOL, + atol=GENERATE_LOGPROB_ATOL, + err_msg=err_msg, + ) + + def _run_colocate_ipc_update_weight_and_generate( + self, + extra_rollout_config: dict | None = None, + sample_params: SampleParams | None = None, + ): + train_controller, rollout_controller = self._setup_engines( + weight_transport_type="ipc", + extra_rollout_config=extra_rollout_config, + ) + + sample_params = sample_params or SGLANG_GREEDY_SAMPLE_PARAMS + + def _generate() -> RolloutState: + return ray.get( + rollout_controller.generate.remote( + rollout_state=RolloutState(message=TEST_TEXT_MESSAGES, sample_params=sample_params), + ), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + + res_baseline = _generate() + self.assertEqual(res_baseline.status, Status.COMPLETED, res_baseline.error_msg) + res_repeat = _generate() + self._assert_generate_outputs_match( + res_repeat, + res_baseline, + err_msg="rollout logprobs changed between repeated generates before weight update", + ) + + # Colocate IPC: free rollout, bring train weights online, then update. + ray.get(rollout_controller.offload.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + train_controller.onload(target="model") + ray.get(rollout_controller.onload_weights.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + targets = ray.get(rollout_controller.get_weight_update_targets.remote()) + train_controller.bind_rollout_weight_update( + targets=targets, + rollout_config=self.rollout_cfg, + ) + train_controller.weight_update() + train_controller.offload(target="model") + ray.get(rollout_controller.onload_kvcache.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + + res_update_weight = _generate() + self._assert_generate_outputs_match( + res_update_weight, + res_baseline, + err_msg="rollout logprobs changed after weight update", + ) + + @unittest.skipIf(os.environ.get("XTUNER_USE_SGLANG", "0") == "0", "sglang backend is not enabled") + def test_sglang_colocate_ipc_update_weight_and_generate(self): + self._run_colocate_ipc_update_weight_and_generate() + + @unittest.skipIf(os.environ.get("XTUNER_USE_LMDEPLOY", "0") == "0", "lmdeploy backend is not enabled") + def test_lmdeploy_colocate_ipc_update_weight_and_generate(self): + self._run_colocate_ipc_update_weight_and_generate( + extra_rollout_config={"lmdeploy_backend": "pytorch"}, + sample_params=LMDEPLOY_GREEDY_SAMPLE_PARAMS, + ) + @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') diff --git a/tests/rl/test_update_weight_disaggregated.py b/tests/rl/test_update_weight_disaggregated.py index 8db8dfa947..12593c0004 100644 --- a/tests/rl/test_update_weight_disaggregated.py +++ b/tests/rl/test_update_weight_disaggregated.py @@ -2,6 +2,7 @@ import tempfile import unittest +import numpy as np import ray import requests @@ -25,8 +26,27 @@ from xtuner.v1.rl.rollout.worker_registry import WorkerLifecycleState RL_TRAINER_RAY_GET_TIMEOUT = 3600 - TEST_TEXT_MESSAGES = [{"role": "user", "content": "Hello!"}] +# Tokens can still match when a few shards are wrong; sampled-token logprobs +# catch that. Keep this tight: identical weights should stay near exact. +GENERATE_LOGPROB_RTOL = 1e-5 +GENERATE_LOGPROB_ATOL = 1e-5 +# SGLang: temperature=0 is greedy. LMDeploy /generate always sets do_sample=True +# and divides logits by temperature, so greedy is top_k=1 with temperature=1.0. +SGLANG_GREEDY_SAMPLE_PARAMS = SampleParams( + temperature=0.0, + max_tokens=128, + top_k=1, + return_logprob=True, + return_token_ids=True, +) +LMDEPLOY_GREEDY_SAMPLE_PARAMS = SampleParams( + temperature=1.0, + max_tokens=128, + top_k=1, + return_logprob=True, + return_token_ids=True, +) MODEL_PATH = os.environ["QWEN3_VL_DENSE_PATH"] class TestUpdateWeightDisaggregated(unittest.TestCase): @@ -138,6 +158,27 @@ def _check_sglang_weights(self, rollout_controller, action): results.append(response.json()) return results + def _assert_generate_outputs_match( + self, + actual: RolloutState, + expected: RolloutState, + err_msg: str, + ) -> None: + self.assertEqual(actual.response, expected.response) + self.assertEqual(actual.response_ids, expected.response_ids) + self.assertIsNotNone(expected.logprobs) + self.assertIsNotNone(actual.logprobs) + self.assertGreater(len(expected.logprobs), 0) + self.assertEqual(len(actual.logprobs), len(expected.logprobs)) + self.assertEqual(len(actual.logprobs), len(actual.response_ids or [])) + np.testing.assert_allclose( + actual.logprobs, + expected.logprobs, + rtol=GENERATE_LOGPROB_RTOL, + atol=GENERATE_LOGPROB_ATOL, + err_msg=err_msg, + ) + @unittest.skipIf(os.environ.get("XTUNER_USE_SGLANG", "0") == "0", "sglang backend is not enabled") def test_sglang_disaggregated_update_weight_and_generate(self): TrainingWorker = ray.remote( @@ -157,9 +198,23 @@ def test_sglang_disaggregated_update_weight_and_generate(self): self.rollout_cfg.skip_load_weights = False rollout_controller = self.rollout_cfg.build(self.rollout_pg) - sample_params = SampleParams(temperature=0.0, max_tokens=128, top_k=1) - input_state = RolloutState(message=TEST_TEXT_MESSAGES, sample_params=sample_params) - res_baseline = ray.get(rollout_controller.generate.remote(rollout_state=input_state)) + sample_params = SGLANG_GREEDY_SAMPLE_PARAMS + + def _generate() -> RolloutState: + return ray.get( + rollout_controller.generate.remote( + rollout_state=RolloutState(message=TEST_TEXT_MESSAGES, sample_params=sample_params), + ), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + + res_baseline = _generate() + res_repeat = _generate() + self._assert_generate_outputs_match( + res_repeat, + res_baseline, + err_msg="rollout logprobs changed between repeated generates before weight update", + ) # 1) 清 KV + 释放旧权重(sleep level=2 -> meta) ray.get( @@ -184,8 +239,12 @@ def test_sglang_disaggregated_update_weight_and_generate(self): timeout=RL_TRAINER_RAY_GET_TIMEOUT, ) - res_update_weight = ray.get(rollout_controller.generate.remote(rollout_state=input_state)) - self.assertEqual(res_update_weight.response, res_baseline.response) + res_update_weight = _generate() + self._assert_generate_outputs_match( + res_update_weight, + res_baseline, + err_msg="rollout logprobs changed after weight update", + ) ray.get(rollout_controller.shutdown.remote(), timeout=60) @unittest.skip("skip sglang parameter-only weight check test until the parameter-check-only patch is applied") @@ -226,7 +285,6 @@ def test_sglang_disaggregated_update_weight_equal_after_reset(self): finally: ray.get(rollout_controller.shutdown.remote(), timeout=60) - @unittest.skip("skip lmdeploy disaggregated update-weight generation test until PR4638 is merged") def test_lmdeploy_disaggregated_update_weight_and_generate(self): # TODO(shipengcheng): Remove skip when CI update lmdeploy. TrainingWorker = ray.remote( @@ -250,9 +308,23 @@ def test_lmdeploy_disaggregated_update_weight_and_generate(self): } rollout_controller = self.rollout_cfg.build(self.rollout_pg) - sample_params = SampleParams(temperature=0.0, max_tokens=128, top_k=1) - input_state = RolloutState(message=TEST_TEXT_MESSAGES, sample_params=sample_params) - res_baseline = ray.get(rollout_controller.generate.remote(rollout_state=input_state)) + sample_params = LMDEPLOY_GREEDY_SAMPLE_PARAMS + + def _generate() -> RolloutState: + return ray.get( + rollout_controller.generate.remote( + rollout_state=RolloutState(message=TEST_TEXT_MESSAGES, sample_params=sample_params), + ), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + + res_baseline = _generate() + res_repeat = _generate() + self._assert_generate_outputs_match( + res_repeat, + res_baseline, + err_msg="rollout logprobs changed between repeated generates before weight update", + ) # 1) 清 KV + 释放旧权重(sleep level=2 -> meta) ray.get( @@ -277,8 +349,12 @@ def test_lmdeploy_disaggregated_update_weight_and_generate(self): timeout=RL_TRAINER_RAY_GET_TIMEOUT, ) - res_update_weight = ray.get(rollout_controller.generate.remote(rollout_state=input_state)) - self.assertEqual(res_update_weight.response, res_baseline.response) + res_update_weight = _generate() + self._assert_generate_outputs_match( + res_update_weight, + res_baseline, + err_msg="rollout logprobs changed after weight update", + ) ray.get(rollout_controller.shutdown.remote(), timeout=60) if __name__ == "__main__": diff --git a/xtuner/v1/rl/weight_update/transport.py b/xtuner/v1/rl/weight_update/transport.py index eb4ded8686..95a89411a0 100644 --- a/xtuner/v1/rl/weight_update/transport.py +++ b/xtuner/v1/rl/weight_update/transport.py @@ -110,7 +110,6 @@ def update(self, weight_iterator: Any, **_: Any) -> None: for batch in batches: self._send(batch) self.after_update_per_group() - DEVICE_MODULE.empty_cache() finally: self.after_update_all_groups() DEVICE_MODULE.empty_cache() @@ -355,7 +354,7 @@ def after_update_per_batch( self, finished: bool, cpu_group: dist.ProcessGroup, train_enable_ep: bool = False ) -> None: # TODO(chenchiyu): narrow this condition. - if finished or (train_enable_ep and self.rollout_tp > 1): + if finished or self.rollout_tp > 1: # Make each TP head rank sync with other ranks in engine_parallel group. # FSDP all-gather of the next state_dict cannot cover this case, so without # this barrier some ranks could overwrite the IPC tensor before LMDeploy loads it. From 8078f6375e0fed28d96ac7f507823fe04d0a5506 Mon Sep 17 00:00:00 2001 From: Pengcheng Shi Date: Mon, 28 Sep 2026 10:27:22 +0800 Subject: [PATCH 2/2] set deterministic reasoning for sglang --- tests/rl/test_update_weight_colocate.py | 27 ++++++++++++-------- tests/rl/test_update_weight_disaggregated.py | 5 ++++ 2 files changed, 21 insertions(+), 11 deletions(-) diff --git a/tests/rl/test_update_weight_colocate.py b/tests/rl/test_update_weight_colocate.py index 04adc25500..f386fff40b 100644 --- a/tests/rl/test_update_weight_colocate.py +++ b/tests/rl/test_update_weight_colocate.py @@ -1,16 +1,17 @@ """共卡部署下的模型权重更新正确性测试。 -覆盖 IPC 与 checkpoint-engine:训练侧权重写入推理引擎后,检查引擎状态是否与 -更新前一致。两类用例互补,不是重复。 +覆盖 IPC 与 checkpoint-engine:训练侧权重写入推理引擎后,检查推理引擎状态是否与 +更新前一致。 Generate 检查(默认启用) greedy generate 两次,确认引擎自身可复现;再走 IPC 更新后第三次 generate。 比较 response / response_ids / sampled-token logprobs。 - SGLang greedy 用 temperature=0;LMDeploy /generate 用 top_k=1 且 temperature=1.0。 + SGLang greedy 用 temperature=0,并打开 enable_deterministic_inference。 + LMDeploy /generate 用 top_k=1 且 temperature=1.0。 - test_sglang_colocate_ipc_update_weight_and_generate - test_lmdeploy_colocate_ipc_update_weight_and_generate -参数逐点检查(当前 skip,依赖 SGLang WeightChecker patch) +checkpoint 检查(当前 skip,依赖 SGLang WeightChecker patch) snapshot_parameters -> reset_parameters -> update_weights -> compare_parameters。 不 generate,直接比对引擎参数。WeightChecker: https://github.com/PengchengShi00/sglang/commit/05e89d63b5a1a80671b267ff4494ad950b2aba75 @@ -28,8 +29,7 @@ from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status -from xtuner.v1.model import Qwen3_5_VLMoE35BA3Config -from xtuner.v1.module.mtp import MTPConfig +from xtuner.v1.model import Qwen3VLDense4BConfig from xtuner.v1.rl.loss import GRPOLossConfig as LossConfig from xtuner.v1.rl.rollout.worker import RolloutConfig @@ -61,7 +61,11 @@ top_k=1, return_logprob=True, return_token_ids=True, + sampling_seed=1024, ) +SGLANG_DETERMINISTIC_EXTRA_CONFIG = { + "sglang_enable_deterministic_inference": True, +} LMDEPLOY_GREEDY_SAMPLE_PARAMS = SampleParams( temperature=1.0, max_tokens=128, @@ -69,7 +73,7 @@ return_logprob=True, return_token_ids=True, ) -MODEL_PATH = os.environ["QWEN3_5_MOE_PATH"] +MODEL_PATH = os.environ["QWEN3_VL_DENSE_PATH"] class TestUpdateWeightColocate(unittest.TestCase): @@ -146,9 +150,7 @@ def init_config(self, *, weight_transport_type: str, extra_rollout_config: dict extra_rollout_config=extra_rollout_config or {}, ) - 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 + model_cfg = Qwen3VLDense4BConfig() optim_cfg = AdamWConfig(lr=1e-6, foreach=False, weight_decay=0.1) fsdp_cfg = FSDPConfig(torch_compile=False, cpu_offload=False, ep_size=1) @@ -298,7 +300,10 @@ def _generate() -> RolloutState: @unittest.skipIf(os.environ.get("XTUNER_USE_SGLANG", "0") == "0", "sglang backend is not enabled") def test_sglang_colocate_ipc_update_weight_and_generate(self): - self._run_colocate_ipc_update_weight_and_generate() + self._run_colocate_ipc_update_weight_and_generate( + extra_rollout_config=SGLANG_DETERMINISTIC_EXTRA_CONFIG, + sample_params=SGLANG_GREEDY_SAMPLE_PARAMS, + ) @unittest.skipIf(os.environ.get("XTUNER_USE_LMDEPLOY", "0") == "0", "lmdeploy backend is not enabled") def test_lmdeploy_colocate_ipc_update_weight_and_generate(self): diff --git a/tests/rl/test_update_weight_disaggregated.py b/tests/rl/test_update_weight_disaggregated.py index 12593c0004..dbe4d5a0ac 100644 --- a/tests/rl/test_update_weight_disaggregated.py +++ b/tests/rl/test_update_weight_disaggregated.py @@ -39,7 +39,11 @@ top_k=1, return_logprob=True, return_token_ids=True, + sampling_seed=1024, ) +SGLANG_DETERMINISTIC_EXTRA_CONFIG = { + "sglang_enable_deterministic_inference": True, +} LMDEPLOY_GREEDY_SAMPLE_PARAMS = SampleParams( temperature=1.0, max_tokens=128, @@ -196,6 +200,7 @@ def test_sglang_disaggregated_update_weight_and_generate(self): train_controller = TrainingController(workers=train_workers) self.rollout_cfg.skip_load_weights = False + self.rollout_cfg.extra_rollout_config = dict(SGLANG_DETERMINISTIC_EXTRA_CONFIG) rollout_controller = self.rollout_cfg.build(self.rollout_pg) sample_params = SGLANG_GREEDY_SAMPLE_PARAMS