diff --git a/tests/rl/test_qwen35_vl_moe_recover_e2e.py b/tests/rl/test_qwen35_vl_moe_recover_e2e.py new file mode 100644 index 0000000000..290ec25ac9 --- /dev/null +++ b/tests/rl/test_qwen35_vl_moe_recover_e2e.py @@ -0,0 +1,521 @@ +"""Real Qwen3.5 VLM MoE checkpoint-engine recovery E2E test. + +This test focuses only on the recovery protocol: + +1. train step 1 registers and broadcasts a checkpoint-engine weight update; +2. while train step 2 rollout is running, rank 0's backend is crashed; +3. RolloutHealthManager restarts the worker into pending_weight_update; +4. the train step 2 checkpoint-engine sync updates the pending worker; +5. train step 2 and the post-recovery train step 3 both complete. + +Run in the same 8-GPU environment used by the Qwen3.5 VLM MoE +async-training E2E test. +""" + +from __future__ import annotations + +import asyncio +import os +import threading +import time +import unittest +from pathlib import Path +from typing import Any, Callable + +import ray + +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.data_proto.rl_data import SampleParams +from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig +from xtuner.v1.datasets.rl_tokenize_fn import RLQwen3VLTokenizeFnConfig +from xtuner.v1.model import Qwen3_5_VLMoE35BA3Config +from xtuner.v1.module.mtp import MTPConfig +from xtuner.v1.rl.advantage import GRPOAdvantageConfig +from xtuner.v1.rl.agent_loop import SingleTurnAgentLoopConfig +from xtuner.v1.rl.agent_loop_manager import ( + AgentLoopManagerConfig, + AsyncProduceStrategyConfig, + SamplerConfig, + TaskSpecConfig, +) +from xtuner.v1.rl.judger import GEO3KJudgerConfig +from xtuner.v1.rl.loss import GRPOLossConfig +from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.rollout.worker_registry import WorkerLifecycleState +from xtuner.v1.rl.trainer import RolloutImportanceSampling, WorkerConfig +from xtuner.v1.rl.utils import AcceleratorResourcesConfig, CPUResourcesConfig +from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig +from xtuner.v1.utils import get_logger + + +EXPERIMENT_NAME = "qwen35_vl_moe_checkpoint_engine_recovery_e2e" +TOTAL_TRAIN_STEPS = 3 +TRAIN_BATCH_SIZE_BY_STEP = {1: 8, 2: 256, 3: 8} +PROMPT_REPEAT_K = 2 +MAX_PROMPT_LENGTH = 4096 +MAX_RESPONSE_LENGTH = 2048 +PACK_MAX_LENGTH = 8192 +RECOVERY_TIMEOUT_S = 600.0 +RAY_GET_TIMEOUT_S = 600.0 +POLL_INTERVAL_S = 0.5 +logger = get_logger() + + +class TestQwen35VLMoECheckpointEngineRecoveryE2E(unittest.TestCase): + def setUp(self) -> None: + self.model_path = self._required_path("QWEN3_5_MOE_PATH") + self.media_root = self._required_path("GEO3K_MEDIA_ROOT") + self.data_path = self._required_path("GEO3K_LONGTAIL_DATA_PATH") + + default_work_dir = ( + Path.cwd() / "work_dirs" / f"{EXPERIMENT_NAME}_{time.strftime('%Y%m%d%H%M%S')}_{os.getpid()}" + ) + self.work_dir = Path(os.environ.get("WORK_DIR", str(default_work_dir))) + self.work_dir.mkdir(parents=True, exist_ok=True) + + self._events: list[str] = [] + self._events_lock = threading.Lock() + self._step_1_weight_update_finished = threading.Event() + self._rollout_step_2_started = threading.Event() + self._rollout_step_2_finished = threading.Event() + self._rank_0_pending_weight_update = threading.Event() + self._recovery_finished = threading.Event() + self._fault_injection_error: Exception | None = None + self._rank_0_lifecycle_states: list[str] = [] + self._produce_calls: list[dict[str, int]] = [] + self._weight_update_calls: list[dict[str, int | bool]] = [] + + self._patch_env( + { + "XTUNER_USE_LMDEPLOY": "0", + "XTUNER_USE_SGLANG": "1", + "XTUNER_USE_VLLM": "0", + "XTUNER_USE_FA3": "1", + "XTUNER_DETERMINISTIC": "false", + "XTUNER_TEST_IMMEDIATE_RECOVERY": "1", + }, + unset=("RAY_ADDRESS","PYTORCH_CUDA_ALLOC_CONF"), + ) + ray.init(address="local", num_cpus=256, num_gpus=8, ignore_reinit_error=True) + + def tearDown(self) -> None: + if ray.is_initialized(): + ray.shutdown() + if hasattr(self, "_old_env"): + self._restore_env() + + @unittest.skipIf(os.environ.get("XTUNER_USE_SGLANG", "0") == "0", "sglang backend is not enabled") + def test_checkpoint_engine_backend_failure_recovery(self) -> None: + trainer = self._build_config().build() + self._install_rollout_probe(trainer) + self._install_checkpoint_engine_probe(trainer) + + fault_injection_thread = threading.Thread( + target=self._inject_failure_after_checkpoint_engine_ready, + args=(trainer,), + name="checkpoint-engine-recovery-fault-injector", + daemon=True, + ) + fault_injection_thread.start() + + try: + trainer.fit() + finally: + fault_injection_thread.join(timeout=10) + + self.assertFalse(fault_injection_thread.is_alive(), "Fault-injection coordinator did not exit.") + if self._fault_injection_error is not None: + raise AssertionError("Fault-injection coordinator failed.") from self._fault_injection_error + + unavailable_states = { + WorkerLifecycleState.INACTIVE.value, + WorkerLifecycleState.PENDING_WEIGHTS.value, + } + self.assertTrue(unavailable_states.intersection(self._rank_0_lifecycle_states)) + self.assertEqual(self._rank_0_lifecycle_states[-1], WorkerLifecycleState.ACTIVE.value) + self.assertEqual( + [call["train_step"] for call in self._produce_calls], + [1, 2, 3], + ) + self.assertEqual( + [call["batch_size"] for call in self._produce_calls], + [TRAIN_BATCH_SIZE_BY_STEP[step] for step in range(1, TOTAL_TRAIN_STEPS + 1)], + ) + self.assertEqual( + [call["train_step"] for call in self._weight_update_calls], + [1, 2], + ) + self.assertTrue(all(call["weights_synced"] for call in self._weight_update_calls)) + self._assert_recovery_event_order() + + def _install_rollout_probe(self, trainer: Any) -> None: + original_produce_batch = trainer.agent_loop_manager.produce_batch + + async def produce_batch_wrapper(batch_size: int, train_step: int, *, model_step: int) -> Any: + batch_size = TRAIN_BATCH_SIZE_BY_STEP.get(train_step, batch_size) + self._record_event(f"rollout_{train_step}_started") + if train_step == 2: + self._rollout_step_2_started.set() + + try: + result = await original_produce_batch(batch_size, train_step, model_step=model_step) + self._produce_calls.append( + { + "batch_size": batch_size, + "train_step": train_step, + "model_step": model_step, + } + ) + if train_step == 2: + pending_weight_update = await asyncio.to_thread( + self._rank_0_pending_weight_update.wait, + RECOVERY_TIMEOUT_S, + ) + if not pending_weight_update: + raise TimeoutError( + "Timed out waiting for rank 0 to restart into pending_weight_update during train step 2 " + "rollout." + ) + return result + finally: + if train_step == 2: + self._rollout_step_2_finished.set() + self._record_event(f"rollout_{train_step}_finished") + + trainer.agent_loop_manager.produce_batch = produce_batch_wrapper + + def _install_checkpoint_engine_probe(self, trainer: Any) -> None: + original_sync_weights_and_save = trainer._sync_weights_and_save + + def sync_weights_and_save_wrapper(train_step: int, step_timer_dict: dict) -> bool: + logger.info(f"[recovery-test] sync_weights_and_save enter train_step={train_step}") + weights_synced = original_sync_weights_and_save(train_step, step_timer_dict) + logger.info( + f"[recovery-test] sync_weights_and_save exit train_step={train_step} " + f"weights_synced={weights_synced}" + ) + if weights_synced: + has_registered_checkpoint = trainer.train_controller.has_registered_weight_checkpoint() + self._weight_update_calls.append( + { + "train_step": train_step, + "weights_synced": weights_synced, + "has_registered_checkpoint": has_registered_checkpoint, + } + ) + self._record_event(f"checkpoint_engine_{train_step}_updated") + if train_step == 1: + if not has_registered_checkpoint: + raise AssertionError("Train step 1 did not register a checkpoint-engine checkpoint.") + rank_0_state = self._get_rank_0_lifecycle_state(trainer) + logger.info( + f"[recovery-test] set step_1_weight_update_finished " + f"rank0_state={rank_0_state}" + ) + self._step_1_weight_update_finished.set() + return weights_synced + + trainer._sync_weights_and_save = sync_weights_and_save_wrapper + + def _inject_failure_after_checkpoint_engine_ready(self, trainer: Any) -> None: + try: + logger.info("[recovery-test] fault injector waiting for step 1 checkpoint-engine update") + if not self._step_1_weight_update_finished.wait(timeout=RECOVERY_TIMEOUT_S): + raise TimeoutError("Timed out waiting for the train step 1 checkpoint-engine update.") + + logger.info("[recovery-test] fault injector waiting for train step 2 rollout start") + if not self._rollout_step_2_started.wait(timeout=RECOVERY_TIMEOUT_S): + raise TimeoutError("Timed out waiting for train step 2 rollout to start.") + if self._rollout_step_2_finished.is_set(): + raise RuntimeError("Train step 2 rollout finished before backend failure injection.") + + initial_state = self._get_rank_0_lifecycle_state(trainer) + if initial_state != WorkerLifecycleState.ACTIVE.value: + raise RuntimeError(f"Rank 0 was not active before fault injection: state={initial_state}.") + self._record_rank_0_state(initial_state) + logger.info(f"[recovery-test] before backend crash injection rank0_state={initial_state}") + + ray.get( + trainer.rollout_controller.inject_backend_crash_for_test.remote(rank=0), + timeout=RAY_GET_TIMEOUT_S, + ) + logger.info("[recovery-test] backend crash injection returned") + self._record_event("backend_crash_injected") + + self._wait_for_rank_0_state( + trainer, + expected=lambda state: state != WorkerLifecycleState.ACTIVE.value, + description="become inactive", + ) + self._record_event("rank_0_unavailable") + self._wait_for_rank_0_state( + trainer, + expected=lambda state: state == WorkerLifecycleState.PENDING_WEIGHTS.value, + description="wait for checkpoint-engine weights", + ) + self._record_event("rank_0_pending_weight_update") + self._rank_0_pending_weight_update.set() + self._wait_for_rank_0_state( + trainer, + expected=lambda state: state == WorkerLifecycleState.ACTIVE.value, + description="recover to active", + ) + self._record_event("rank_0_recovered") + except Exception as error: + self._fault_injection_error = error + finally: + self._recovery_finished.set() + + def _wait_for_rank_0_state( + self, + trainer: Any, + *, + expected: Callable[[str], bool], + description: str, + ) -> str: + deadline = time.monotonic() + RECOVERY_TIMEOUT_S + while time.monotonic() < deadline: + state = self._get_rank_0_lifecycle_state(trainer) + self._record_rank_0_state(state) + if expected(state): + return state + time.sleep(POLL_INTERVAL_S) + raise TimeoutError( + f"Timed out waiting for rank 0 to {description}; observed states={self._rank_0_lifecycle_states}." + ) + + @staticmethod + def _get_rank_0_lifecycle_state(trainer: Any) -> str: + targets = ray.get( + trainer.rollout_controller.get_weight_update_targets.remote(), + timeout=RAY_GET_TIMEOUT_S, + ) + for target in targets: + if target.endpoint_rank == 0: + return target.lifecycle_state + raise RuntimeError(f"Rank 0 weight-update target was not found: targets={targets}.") + + def _record_rank_0_state(self, state: str) -> None: + if not self._rank_0_lifecycle_states or self._rank_0_lifecycle_states[-1] != state: + self._rank_0_lifecycle_states.append(state) + + def _assert_recovery_event_order(self) -> None: + required_events = ( + "checkpoint_engine_1_updated", + "rollout_2_started", + "backend_crash_injected", + "rank_0_unavailable", + "rank_0_pending_weight_update", + "checkpoint_engine_2_updated", + "rank_0_recovered", + "rollout_2_finished", + "rollout_3_started", + "rollout_3_finished", + ) + for event in required_events: + self.assertEqual(self._events.count(event), 1, f"Unexpected event count for {event}: {self._events}") + + positions = {event: self._events.index(event) for event in required_events} + ordered_pairs = ( + ("checkpoint_engine_1_updated", "backend_crash_injected"), + ("rollout_2_started", "backend_crash_injected"), + ("backend_crash_injected", "rank_0_unavailable"), + ("rank_0_unavailable", "rank_0_pending_weight_update"), + ("rank_0_pending_weight_update", "rank_0_recovered"), + ("rank_0_recovered", "rollout_2_finished"), + ("rollout_2_finished", "checkpoint_engine_2_updated"), + ("checkpoint_engine_2_updated", "rollout_3_started"), + ("rollout_3_started", "rollout_3_finished"), + ) + for first, second in ordered_pairs: + self.assertLess(positions[first], positions[second], f"Expected {first} before {second}: {self._events}") + + def _record_event(self, event: str) -> None: + with self._events_lock: + self._events.append(event) + + def _build_config(self) -> RLColocateTrainerConfig: + resources = AcceleratorResourcesConfig( + accelerator="GPU", + num_workers=8, + num_cpus_per_worker=12, + cpu_memory_per_worker=24 * 1024**3, + ) + rollout_config = RolloutConfig( + env=EXPERIMENT_NAME, + device=resources.accelerator, + model_path=str(self.model_path), + tokenizer_path=str(self.model_path), + dtype="bfloat16", + tensor_parallel_size=1, + expert_parallel_size=4, + gpu_memory_utilization=0.8, + context_length=MAX_PROMPT_LENGTH + MAX_RESPONSE_LENGTH, + rollout_max_batch_size_per_instance=128, + allow_over_concurrency_ratio=1.0, + enable_return_routed_experts=False, + weight_transport_type="checkpoint_engine", + skip_load_weights=True, + checkpoint_name_prefix=EXPERIMENT_NAME, + checkpoint_engine_timeout=RECOVERY_TIMEOUT_S, + # 更快的发现错误并且重启 + health_check_interval_seconds=5.0, + health_check_failure_threshold=1, + extra_rollout_config={ + "sglang_log_level": "error", + }, + ) + model_cfg = Qwen3_5_VLMoE35BA3Config(freeze_vision=True, freeze_projector=True) + model_cfg.text_config.mtp_config = MTPConfig(num_layers=1) + train_worker_cfg = WorkerConfig( + model_cfg=model_cfg, + load_from=str(self.model_path), + optim_cfg=AdamWConfig( + lr=1e-6, + betas=(0.9, 0.999), + max_grad_norm=1.0, + weight_decay=0.1, + foreach=False, + swap_optimizer=True, + ), + loss_cfg=GRPOLossConfig( + policy_loss_cfg={ + "cliprange_high": 0.28, + "cliprange_low": 0.2, + "loss_type": "vanilla", + "clip_ratio_c": 10.0, + "log_prob_diff_min": -20, + "log_prob_diff_max": 20, + }, + ignore_idx=-100, + use_kl_loss=False, + kl_loss_coef=0.0, + kl_loss_type="low_var_kl", + mode="chunk", + chunk_size=512, + rollout_is=RolloutImportanceSampling( + rollout_is_level="token", + rollout_is_mode="both", + rollout_is_threshold=(5, 0.5), + rollout_is_mask_threshold=(5, 0.5), + rollout_is_veto_threshold=(20, 0), + ), + ), + lr_cfg=LRConfig(lr_type="constant", warmup_ratio=0, lr_min=1e-6), + fsdp_cfg=FSDPConfig(torch_compile=False, cpu_offload=False, ep_size=1, fp32_lm_head=False), + sp_size=1, + optimizer_steps=8, + pack_max_length=PACK_MAX_LENGTH, + ) + + dataloader_cfg = DataloaderConfig( + dataset_config_list=[ + { + "dataset": DatasetConfig( + name=EXPERIMENT_NAME, + anno_path=self.data_path, + class_name="VLMJsonlDataset", + media_root=str(self.media_root), + ), + "tokenize_fn": RLQwen3VLTokenizeFnConfig( + processor_path=str(self.model_path), + max_length=MAX_PROMPT_LENGTH, + chat_template="qwen3.5-vl", + add_generation_prompt=True, + enable_thinking=True, + ), + } + ], + pack_max_length=PACK_MAX_LENGTH, + collator="fake_collator", + pack_level="none", + ) + agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=[ + TaskSpecConfig( + task_name="geo3k_longtail", + agent_loop_config=SingleTurnAgentLoopConfig( + hf_checkpoint=str(self.model_path), + sample_params=SampleParams( + max_tokens=MAX_RESPONSE_LENGTH, + top_k=0, + top_p=1.0, + temperature=0.0, + min_tokens=0, + return_logprob=True, + return_token_ids=True, + return_routed_experts=False, + ), + ), + judger_config=GEO3KJudgerConfig( + judger_name="hiyouga/geometry3k", + cpu_resources=CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1), + ), + produce_strategy_config=AsyncProduceStrategyConfig( + over_sample_threshold=1.0, + enable_partial_rollout=False, + max_staleness=1, + max_pending_tasks=16, + ), + sampler_config=SamplerConfig( + dataloader_cfg=dataloader_cfg, + prompt_repeat_k=PROMPT_REPEAT_K, + ), + ) + ], + ) + + return RLColocateTrainerConfig( + resources=resources, + train_worker_cfg=train_worker_cfg, + rollout_config=rollout_config, + tokenizer_path=str(self.model_path), + replay_buffer_config=AsyncReplayBufferConfig(), + agent_loop_manager_cfg=agent_loop_manager_cfg, + load_from=str(self.model_path), + total_train_steps=TOTAL_TRAIN_STEPS, + train_batch_size=TRAIN_BATCH_SIZE_BY_STEP[1], + advantage_estimator_config=GRPOAdvantageConfig(eps=1e-8), + sync_weights_interval=1, + enable_evaluate=False, + enable_initial_evaluate=False, + evaluate_step=1, + work_dir=str(self.work_dir), + checkpoint_interval=-1, + checkpoint_maxkeep=-1, + hf_interval=-1, + hf_max_keep=-1, + seed=123, + debug_rollout=False, + exp_tracker="jsonl", + ) + + @staticmethod + def _required_path(env_name: str) -> Path: + value = os.environ.get(env_name) + if not value: + raise RuntimeError(f"{env_name} must be set for the checkpoint-engine recovery E2E test.") + path = Path(value) + if not path.exists(): + raise FileNotFoundError(f"{env_name} does not exist: {path}") + return path + + def _patch_env(self, updates: dict[str, str], *, unset: tuple[str, ...] = ()) -> None: + keys = set(updates) | set(unset) + self._old_env = {key: os.environ.get(key) for key in keys} + for key, value in updates.items(): + os.environ[key] = value + for key in unset: + os.environ.pop(key, None) + + def _restore_env(self) -> None: + for key, value in self._old_env.items(): + if value is None: + os.environ.pop(key, None) + else: + os.environ[key] = value + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/rl/test_rl_colocate_trainer.py b/tests/rl/test_rl_colocate_trainer.py index 95b46b6ac0..b5ef8e96e8 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -32,6 +32,7 @@ from xtuner.v1.rl.agent_loop_manager import AsyncProduceStrategyConfig, ProduceBatchResult from xtuner.v1.rl.agent_loop_manager.agent_loop_manager import AgentLoopManager from xtuner.v1.rl.agent_loop_manager.produce_utils import _TaskRunner +from xtuner.v1.rl.health_manager import RLHealthManager from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig, SerializedRayObjectRef from xtuner.v1.train.rl_trainer import RLColocateTrainer, RLThroughputBenchmark @@ -112,6 +113,7 @@ def tearDown(self): def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_weights_interval: int = 1): trainer = RLColocateTrainer.__new__(RLColocateTrainer) + trainer._rollout_config = SimpleNamespace(weight_transport_type='ipc') trainer.logger = MagicMock() trainer._total_train_steps = total_train_steps trainer._cur_step = 0 @@ -153,13 +155,13 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_ ) trainer.rollout_controller = SimpleNamespace( - check_and_shutdown_inactive_workers=SimpleNamespace( - remote=MagicMock(return_value="rollout_inactive_workers_shutdown") + shutdown_inactive_workers=SimpleNamespace( + remote=MagicMock(side_effect=_ray_get_none_ref) ), - offload=SimpleNamespace(remote=MagicMock(return_value="rollout_offloaded")), - restart_inactive_workers=SimpleNamespace(remote=MagicMock(return_value="rollout_restarted")), - onload_weights=SimpleNamespace(remote=MagicMock(return_value="weights_loaded")), - onload_kvcache=SimpleNamespace(remote=MagicMock(return_value="kvcache_loaded")), + offload=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)), + restart_inactive_workers=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)), + onload_weights=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)), + onload_kvcache=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)), validate_registered_workers_to_proxy=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)), ) trainer.train_controller = SimpleNamespace( @@ -179,6 +181,11 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_ ] ), ) + trainer.rl_health_manager = RLHealthManager( + train_controller=trainer.train_controller, + rollout_controller=trainer.rollout_controller, + rollout_config=trainer._rollout_config, + ) return trainer def test_fit_accepts_async_strategy_manager_on_colocate_path(self): @@ -240,7 +247,7 @@ async def _produce_batch(batch_size, train_step, *, model_step): return ProduceBatchResult(rollout_states=[[_FakeRolloutState(train_step)]]) trainer = self._make_trainer(SimpleNamespace(produce_batch=_produce_batch)) - trainer.rollout_controller.check_and_shutdown_inactive_workers.remote.side_effect = RuntimeError( + trainer.rollout_controller.shutdown_inactive_workers.remote.side_effect = RuntimeError( "inactive rollout workers after recovery" ) @@ -251,7 +258,7 @@ async def _produce_batch(batch_size, train_step, *, model_step): with self.assertRaisesRegex(RuntimeError, "inactive rollout workers"): trainer.fit() - trainer.rollout_controller.check_and_shutdown_inactive_workers.remote.assert_called_once_with() + trainer.rollout_controller.shutdown_inactive_workers.remote.assert_called_once_with() trainer.rollout_controller.offload.remote.assert_not_called() trainer.train_controller.onload.assert_not_called() trainer.train_controller.fit.assert_not_called() diff --git a/tests/rl/test_rl_disaggregated_trainer.py b/tests/rl/test_rl_disaggregated_trainer.py index bad8565f74..933d9e76fd 100644 --- a/tests/rl/test_rl_disaggregated_trainer.py +++ b/tests/rl/test_rl_disaggregated_trainer.py @@ -148,7 +148,7 @@ def _make_trainer(self, agent_loop_manager): weight_update=MagicMock(return_value="update"), ) trainer.rollout_controller = SimpleNamespace( - check_and_shutdown_inactive_workers=SimpleNamespace( + shutdown_inactive_workers=SimpleNamespace( remote=MagicMock(return_value="rollout_inactive_workers_shutdown") ), restart_inactive_workers=SimpleNamespace(remote=MagicMock(return_value="rollout_restarted")), diff --git a/tests/rl/test_rl_trainer_checkpoint.py b/tests/rl/test_rl_trainer_checkpoint.py index 0c4d53fa3e..a59b11ded4 100644 --- a/tests/rl/test_rl_trainer_checkpoint.py +++ b/tests/rl/test_rl_trainer_checkpoint.py @@ -102,11 +102,12 @@ def __init__(self): self.continue_generation = _RemoteMethod(async_result=True) self.flush_cache = _RemoteMethod(return_value="cache_flushed") self.offload = _RemoteMethod(return_value="rollout_offloaded") - self.check_and_shutdown_inactive_workers = _RemoteMethod(return_value="rollout_inactive_workers_shutdown") + self.shutdown_inactive_workers = _RemoteMethod(return_value="rollout_inactive_workers_shutdown") self.restart_inactive_workers = _RemoteMethod(return_value="rollout_restarted") self.onload_weights = _RemoteMethod(return_value="weights_loaded") self.onload_kvcache = _RemoteMethod(return_value="kvcache_loaded") self.get_weight_update_targets = _RemoteMethod(return_value=()) + self.mark_worker_groups_lifecycle_state = _RemoteMethod(return_value=None) self.set_enable_partial_rollout = _RemoteMethod(return_value=None) self.validate_registered_workers_to_proxy = _RemoteMethod(return_value=_AwaitableValue(None)) diff --git a/tests/rl/test_rollout_logic.py b/tests/rl/test_rollout_logic.py index e5a1a879b3..95b3898f1c 100644 --- a/tests/rl/test_rollout_logic.py +++ b/tests/rl/test_rollout_logic.py @@ -21,6 +21,7 @@ from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status from xtuner.v1.rl.agent_loop import AgentLoopConfig +from xtuner.v1.rl.health_manager import RLHealthManager from xtuner.v1.rl.rollout.controller import RolloutController from xtuner.v1.rl.rollout.health_manager import RolloutHealthManager from xtuner.v1.rl.rollout.lmdeploy import LMDeployWorker @@ -201,7 +202,8 @@ def _weight_update_targets(self, topology: RolloutTopology): for spec in topology.server_launch_specs() ), ) - return registry.weight_update_targets() + targets, _ = registry.weight_update_targets() + return targets def _rollout_info(self, *, config, targets, train_rank: int): return RolloutWeightUpdateInfo.from_targets( @@ -430,6 +432,51 @@ def test_validate_registered_workers_to_proxy_delegates_proxy_validation(self): controller.proxy_manager.validate_registered_session_urls.assert_called_once_with() + def test_pause_generation_waits_for_worker_recovery(self): + controller = RolloutController.__new__(RolloutController) + controller.health_manager = MagicMock() + controller.health_manager.wait_recovery_done.return_value = True + controller.registry = MagicMock() + controller.registry.active_workers.return_value = () + controller.logger = MagicMock() + + with patch("xtuner.v1.rl.rollout.controller.ray.get", return_value=[]): + controller.pause_generation() + + controller.health_manager.pause.assert_called_once_with() + controller.health_manager.wait_recovery_done.assert_called_once_with(timeout=600.0) + controller.health_manager.shutdown_inactive_workers.assert_not_called() + + def test_pause_generation_fails_when_worker_recovery_times_out(self): + controller = RolloutController.__new__(RolloutController) + controller.health_manager = MagicMock() + controller.health_manager.wait_recovery_done.return_value = False + + with self.assertRaisesRegex(TimeoutError, "rollout worker recovery"): + controller.pause_generation() + + def test_mark_worker_groups_lifecycle_state_defaults_to_all_source_groups(self): + controller = RolloutController.__new__(RolloutController) + controller.registry = self._build_registry((0, 1)) + _register_started_servers( + controller.registry, + ( + (0, object(), "http://worker-0", "http://session-0"), + (1, object(), "http://worker-1", "http://session-1"), + ), + ) + controller.registry.mark_unhealthy_ranks({0, 1}) + controller.health_manager = MagicMock() + + controller.mark_worker_groups_lifecycle_state( + source_state=WorkerLifecycleState.INACTIVE, + target_state=WorkerLifecycleState.ACTIVE, + ) + + self.assertTrue(all(worker.is_active() for worker in controller.registry.all_workers())) + active_groups = controller.health_manager.notify_worker_group_active.call_args.args[0] + self.assertEqual([group.ranks for group in active_groups], [(0,), (1,)]) + class TestRolloutProxyManager(unittest.TestCase): _ROUTED_PROXY_URL = "http://routed-proxy" @@ -537,7 +584,7 @@ def test_inactive_lifecycle_listener_deletes_entrypoint_session_urls(self): manager._delete_session_url.assert_called_once_with("http://session-0") manager._register_session_url.assert_not_called() - def test_recovered_lifecycle_listener_registers_entrypoint_session_urls_without_validation(self): + def test_active_lifecycle_listener_registers_entrypoint_session_urls_without_validation(self): manager = self._build_manager() manager._register_session_url = MagicMock() worker_group = SimpleNamespace( @@ -552,7 +599,7 @@ def test_recovered_lifecycle_listener_registers_entrypoint_session_urls_without_ ) ) - manager.on_worker_group_recovered(worker_group) + manager.on_worker_group_active(worker_group) manager._register_session_url.assert_called_once_with("http://session-0") @@ -645,11 +692,34 @@ def test_registry_filters_entrypoints_and_tracks_lifecycle(self): inactive_groups = registry.inactive_worker_groups() self.assertEqual(inactive_groups[0].ranks, (0, 1)) self.assertEqual(self._worker_by_rank(registry, 0).lifecycle_state, WorkerLifecycleState.INACTIVE) - claimed_groups = registry.claim_inactive_groups_for_recovery() - self.assertEqual(claimed_groups[0].ranks, (0, 1)) - self.assertEqual(self._worker_by_rank(registry, 0).lifecycle_state, WorkerLifecycleState.RECOVERING) - registry.set_group_recovery_result(claimed_groups[0], recovered=False) + recovery_groups = registry.get_inactive_groups_for_recovery() + self.assertEqual(recovery_groups[0].ranks, (0, 1)) self.assertEqual(self._worker_by_rank(registry, 0).lifecycle_state, WorkerLifecycleState.INACTIVE) + registry.set_group_recovery_result(recovery_groups[0], recovered=False) + self.assertEqual(self._worker_by_rank(registry, 0).lifecycle_state, WorkerLifecycleState.INACTIVE) + + def test_registry_sets_groups_state_with_source_filter(self): + runtime_layout = self._runtime_layout(engine_ranks=(0,)) + registry = RolloutWorkerRegistry(rollout_topology=runtime_layout) + _register_started_servers( + registry, + ((0, object(), "http://worker-0", "http://session-0"),), + lifecycle_state=WorkerLifecycleState.PENDING_WEIGHTS, + ) + + pending_group = registry.get_target_state_worker_groups(WorkerLifecycleState.PENDING_WEIGHTS)[0] + updated_groups = registry.set_groups_state( + groups=[pending_group], + target_state=WorkerLifecycleState.ACTIVE, + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + ) + + self.assertEqual(updated_groups[0].ranks, (0,)) + self.assertEqual(registry.get_target_state_worker_groups(WorkerLifecycleState.PENDING_WEIGHTS), ()) + self.assertEqual( + tuple(worker.rank for worker in registry.get_target_state_workers(WorkerLifecycleState.ACTIVE)), + (0,), + ) def test_registry_projects_weight_update_targets_from_topology_and_runtime_state(self): runtime_layout = self._runtime_layout(engine_ranks=(0, 1)) @@ -659,16 +729,16 @@ def test_registry_projects_weight_update_targets_from_topology_and_runtime_state ((0, object(), "http://worker-0", "http://session-0"),), ) - targets = registry.weight_update_targets() + targets, group_ranks = registry.weight_update_targets() self.assertEqual(len(targets), 1) + self.assertEqual(group_ranks, ((0,),)) target = targets[0] self.assertEqual(target.endpoint_rank, 0) self.assertEqual(target.update_ranks, (0, 1)) self.assertEqual(target.engine_size, 2) self.assertEqual(target.server_url, "http://worker-0") self.assertEqual(target.lifecycle_state, WorkerLifecycleState.ACTIVE.value) - self.assertTrue(target.is_active) class TestSessionRouter(unittest.IsolatedAsyncioTestCase): @@ -1109,6 +1179,131 @@ def _build_manager( registry, ) + def test_pending_weight_recovery_is_disabled_for_ipc_and_nccl(self): + for transport_type in ("ipc", "nccl"): + with self.subTest(transport_type=transport_type): + manager = RLHealthManager( + train_controller=MagicMock(), + rollout_controller=MagicMock(), + rollout_config=SimpleNamespace(weight_transport_type=transport_type), + ) + + manager.start() + manager.set_rollout_resources_available(True) + manager.stop() + + self.assertFalse(manager.enable_pending_weight_recovery) + self.assertIsNone(manager._pending_rollout_weight_update_thread) + + def test_disabling_rollout_resources_waits_for_pending_weight_update(self): + manager = RLHealthManager( + train_controller=MagicMock(), + rollout_controller=MagicMock(), + rollout_config=SimpleNamespace(weight_transport_type="checkpoint_engine"), + ) + manager._rollout_resources_available.set() + + class _DrainLock: + def __init__(self): + self.released = False + + def acquire(self, *, timeout): + self.timeout = timeout + self.event_was_cleared = not manager._rollout_resources_available.is_set() + return True + + def release(self): + self.released = True + + drain_lock = _DrainLock() + manager._rollout_weight_update_lock = drain_lock + + with patch("xtuner.v1.rl.health_manager.ray.get", side_effect=lambda value, timeout=None: value): + manager.set_rollout_resources_available(False) + + self.assertTrue(drain_lock.event_was_cleared) + self.assertEqual(drain_lock.timeout, 600.0) + self.assertTrue(drain_lock.released) + + def test_disabling_rollout_resources_times_out_waiting_for_weight_update(self): + manager = RLHealthManager( + train_controller=MagicMock(), + rollout_controller=MagicMock(), + rollout_config=SimpleNamespace(weight_transport_type="checkpoint_engine"), + ) + manager._rollout_resources_available.set() + manager._rollout_weight_update_lock = SimpleNamespace( + acquire=MagicMock(return_value=False), + release=MagicMock(), + ) + + with patch("xtuner.v1.rl.health_manager.ray.get", side_effect=lambda value, timeout=None: value): + with self.assertRaisesRegex(TimeoutError, "pending rollout weight update"): + manager.set_rollout_resources_available(False) + + self.assertFalse(manager._rollout_resources_available.is_set()) + manager._rollout_weight_update_lock.acquire.assert_called_once_with(timeout=600.0) + manager._rollout_weight_update_lock.release.assert_not_called() + + def test_pending_weight_update_rechecks_rollout_phase_after_lock_acquire(self): + manager = RLHealthManager( + train_controller=MagicMock(), + rollout_controller=MagicMock(), + rollout_config=SimpleNamespace(weight_transport_type="checkpoint_engine"), + ) + manager._pending_rollout_weight_update_stop_event = SimpleNamespace( + wait=MagicMock(side_effect=(False, True)) + ) + manager._rollout_resources_available = SimpleNamespace( + is_set=MagicMock(side_effect=(True, False)) + ) + manager._rollout_weight_update_lock = SimpleNamespace( + acquire=MagicMock(return_value=True), + release=MagicMock(), + ) + manager._update_pending_rollout_weights_from_checkpoint_engine = MagicMock() + + manager._pending_rollout_worker_weight_update_loop() + + manager._rollout_weight_update_lock.acquire.assert_called_once_with(blocking=False) + manager._rollout_weight_update_lock.release.assert_called_once_with() + manager._update_pending_rollout_weights_from_checkpoint_engine.assert_not_called() + + def test_rl_health_manager_updates_pending_checkpoint_engine_workers(self): + pending_target = SimpleNamespace(endpoint_rank=0) + rollout_controller = SimpleNamespace( + get_weight_update_targets=SimpleNamespace(remote=MagicMock(return_value=((pending_target,), ((0,),)))), + onload_weights=SimpleNamespace(remote=MagicMock(return_value=None)), + onload_kvcache=SimpleNamespace(remote=MagicMock(return_value=None)), + mark_worker_groups_lifecycle_state=SimpleNamespace(remote=MagicMock(return_value=None)), + ) + train_controller = SimpleNamespace( + has_registered_weight_checkpoint=MagicMock(return_value=True), + bind_rollout_weight_update=MagicMock(), + weight_update=MagicMock(), + ) + rollout_config = SimpleNamespace(weight_transport_type="checkpoint_engine") + manager = RLHealthManager( + train_controller=train_controller, + rollout_controller=rollout_controller, + rollout_config=rollout_config, + ) + + with patch("xtuner.v1.rl.health_manager.ray.get", side_effect=lambda value, timeout=None: value): + updated_groups = manager._update_pending_rollout_weights_from_checkpoint_engine() + + self.assertEqual(updated_groups, ((0,),)) + train_controller.bind_rollout_weight_update.assert_called_once_with( + targets=(pending_target,), + rollout_config=rollout_config, + ) + train_controller.weight_update.assert_called_once_with(need_register=False, need_update=True) + rollout_controller.mark_worker_groups_lifecycle_state.remote.assert_called_once_with( + group_ranks=[(0,)], + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + target_state=WorkerLifecycleState.ACTIVE, + ) + def test_marks_worker_inactive_after_consecutive_health_failures(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(False)) worker_info = WorkerSnapshot(rank=0, actor=actor, url="http://worker-0") @@ -1116,13 +1311,14 @@ def test_marks_worker_inactive_after_consecutive_health_failures(self): inactive_groups = [] listener = SimpleNamespace( on_worker_group_inactive=inactive_groups.append, - on_worker_group_recovered=MagicMock(), + on_worker_group_active=MagicMock(), ) manager, registry = self._build_manager( workers_info, failure_threshold=2, worker_lifecycle_listeners=[listener], ) + manager.resume() manager.run_once() @@ -1152,10 +1348,11 @@ def on_worker_group_inactive(group): manager._worker_lifecycle_listeners = ( SimpleNamespace( on_worker_group_inactive=on_worker_group_inactive, - on_worker_group_recovered=MagicMock(), + on_worker_group_active=MagicMock(), ), ) + manager.resume() manager.run_once() self.assertEqual(lock_acquired_by_listener, [True]) @@ -1174,7 +1371,7 @@ def test_inactive_worker_is_not_cleaned_up_again(self): inactive_groups = [] listener = SimpleNamespace( on_worker_group_inactive=inactive_groups.append, - on_worker_group_recovered=MagicMock(), + on_worker_group_active=MagicMock(), ) manager, _ = self._build_manager(workers_info, worker_lifecycle_listeners=[listener]) @@ -1224,24 +1421,33 @@ def test_run_once_does_not_log_error_when_last_active_worker_becomes_inactive(se actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(False)) worker_info = WorkerSnapshot(rank=0, actor=actor, url="http://worker-0") manager, registry = self._build_manager({0: worker_info}, failure_threshold=1) + manager.resume() with patch("xtuner.v1.rl.rollout.health_manager.logger.error") as log_error: manager.run_once() - log_error.assert_not_called() + self.assertFalse( + any("No active rollout worker" in call.args[0] for call in log_error.call_args_list), + f"Expected no stale no-active-worker log, got: {log_error.call_args_list}", + ) self.assertFalse(self._worker_by_rank(registry, 0).is_active()) self.assertEqual(actor.check_health.calls, [()]) - def test_fail_fast_health_check_still_runs_when_periodic_health_check_is_disabled(self): + def test_shutdown_inactive_workers_shuts_already_inactive_groups_when_periodic_check_is_disabled(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(False)) - worker_info = WorkerSnapshot(rank=0, actor=actor, url="http://worker-0") + worker_info = WorkerSnapshot( + rank=0, + actor=actor, + url="http://worker-0", + lifecycle_state=WorkerLifecycleState.INACTIVE, + ) manager, registry = self._build_manager({0: worker_info}, failure_threshold=0) with patch.object(manager, "_shutdown_worker_group", return_value=True): - manager.check_and_shutdown_inactive_workers() + manager.shutdown_inactive_workers() self.assertFalse(self._worker_by_rank(registry, 0).is_active()) - self.assertEqual(actor.check_health.calls, [()]) + self.assertEqual(actor.check_health.calls, []) def test_health_check_uses_configured_timeout(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) @@ -1253,11 +1459,29 @@ async def fake_wait_for(awaitable, timeout): observed_timeouts.append(timeout) return await awaitable + manager.resume() with patch("xtuner.v1.rl.rollout.health_manager.asyncio.wait_for", side_effect=fake_wait_for): manager.run_once() self.assertEqual(observed_timeouts, [2.5]) + def test_background_health_check_respects_pause(self): + actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) + worker_info = WorkerSnapshot(rank=0, actor=actor, url="http://worker-0") + manager, _ = self._build_manager({0: worker_info}) + + manager.run_once() + + self.assertEqual(actor.check_health.calls, []) + + def test_wait_recovery_done_times_out_while_recovery_is_pending(self): + actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) + worker_info = WorkerSnapshot(rank=0, actor=actor, url="http://worker-0") + manager, _ = self._build_manager({0: worker_info}) + manager._recovery_done.clear() + + self.assertFalse(manager.wait_recovery_done(timeout=0.0)) + def test_wait_until_next_check_waits_for_resume_when_paused_during_interval(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) worker_info = WorkerSnapshot(rank=0, actor=actor, url="http://worker-0") @@ -1299,7 +1523,7 @@ def test_shutdown_barrier_keeps_failed_shutdown_group_inactive(self): manager, registry = self._build_manager({0: worker_info}) with patch.object(manager, "_shutdown_worker_group", return_value=False): - manager.check_and_shutdown_inactive_workers() + manager.shutdown_inactive_workers() self.assertEqual(self._worker_by_rank(registry, 0).lifecycle_state, WorkerLifecycleState.INACTIVE) @@ -1325,32 +1549,6 @@ def test_restart_barrier_keeps_failed_recovery_group_inactive(self): f"Expected restart failure log to explain why it is non-fatal, got: {log_error.call_args_list}", ) - def test_restart_barrier_notifies_recovered_group_after_success(self): - actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) - worker_info = WorkerSnapshot( - rank=0, - actor=actor, - url="http://worker-0", - session_url="http://session-0", - lifecycle_state=WorkerLifecycleState.INACTIVE, - ) - recovered_groups = [] - listener = SimpleNamespace( - on_worker_group_inactive=MagicMock(), - on_worker_group_recovered=recovered_groups.append, - ) - manager, registry = self._build_manager( - {0: worker_info}, - worker_lifecycle_listeners=[listener], - ) - - with patch.object(manager, "_restart_worker_group", return_value=True): - manager.restart_inactive_workers() - - self.assertTrue(self._worker_by_rank(registry, 0).is_active()) - self.assertEqual([group.ranks for group in recovered_groups], [(0,)]) - self.assertTrue(all(worker.is_active() for worker in recovered_groups[0].workers)) - def test_restart_barrier_cleans_claimed_groups_when_stopping(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) worker_info = WorkerSnapshot( @@ -1388,7 +1586,7 @@ def test_shutdown_without_waiting_server_down_does_not_probe_worker_server(self) lifecycle_state=WorkerLifecycleState.INACTIVE, ) manager, registry = self._build_manager({0: worker_info}) - group = registry.claim_inactive_groups_for_recovery()[0] + group = registry.get_inactive_groups_for_recovery()[0] def fake_ray_get(ref, timeout=None): del timeout @@ -1422,7 +1620,7 @@ def test_restart_worker_group_uses_reinit(self): lifecycle_state=WorkerLifecycleState.INACTIVE, ) manager, registry = self._build_manager({0: worker_info}) - group = registry.claim_inactive_groups_for_recovery()[0] + group = registry.get_inactive_groups_for_recovery()[0] def fake_ray_get(refs, timeout=None): del timeout @@ -1442,36 +1640,6 @@ def fake_ray_get(refs, timeout=None): self.assertEqual(actor.offload.calls, [()]) self.assertEqual(actor.restore_skip_load_weights.calls, [()]) - def test_recovered_listener_runs_outside_lifecycle_operation_lock(self): - actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) - worker_info = WorkerSnapshot( - rank=0, - actor=actor, - url="http://worker-0", - lifecycle_state=WorkerLifecycleState.INACTIVE, - ) - lock_acquired_by_listener = [] - manager, _ = self._build_manager({0: worker_info}) - - def on_worker_group_recovered(group): - acquired = manager._lifecycle_operation_lock.acquire(blocking=False) - lock_acquired_by_listener.append(acquired) - if acquired: - manager._lifecycle_operation_lock.release() - - manager._worker_lifecycle_listeners = ( - SimpleNamespace( - on_worker_group_inactive=MagicMock(), - on_worker_group_recovered=on_worker_group_recovered, - ), - ) - - with patch.object(manager, "_restart_worker_group", return_value=True): - manager.restart_inactive_workers() - - self.assertEqual(lock_acquired_by_listener, [True]) - - class TestPartialRolloutHandler(unittest.IsolatedAsyncioTestCase): async def test_preprocess_and_postprocess_preserve_response_prefix(self): # partial rollout 续写时应复用 prompt+历史 response,并把新 response token 追加到历史后面。 diff --git a/tests/rl/test_update_weight_colocate.py b/tests/rl/test_update_weight_colocate.py index 4a058c3c48..f2bc0a7ff5 100644 --- a/tests/rl/test_update_weight_colocate.py +++ b/tests/rl/test_update_weight_colocate.py @@ -29,6 +29,7 @@ clear_cpu_resource_manager, set_cpu_resource_manager, ) +from xtuner.v1.rl.rollout.worker_registry import WorkerLifecycleState MODEL_PATH = os.environ["QWEN3_5_MOE_PATH"] @@ -164,7 +165,7 @@ def _setup_engines(self, *, weight_transport_type: str): 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] + active_urls = [target.server_url for target in targets if target.lifecycle_state == WorkerLifecycleState.ACTIVE.value] self.assertGreater(len(active_urls), 0) results = [] for url in active_urls: diff --git a/tests/rl/test_update_weight_disaggregated.py b/tests/rl/test_update_weight_disaggregated.py index 850ad0610a..62bf694ec5 100644 --- a/tests/rl/test_update_weight_disaggregated.py +++ b/tests/rl/test_update_weight_disaggregated.py @@ -22,6 +22,7 @@ clear_cpu_resource_manager, set_cpu_resource_manager, ) +from xtuner.v1.rl.rollout.worker_registry import WorkerLifecycleState TEST_TEXT_MESSAGES = [{"role": "user", "content": "Hello!"}] MODEL_PATH = os.environ["QWEN3_VL_DENSE_PATH"] @@ -120,7 +121,7 @@ def init_config(self): 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] + active_urls = [target.server_url for target in targets if target.lifecycle_state == WorkerLifecycleState.ACTIVE.value] self.assertGreater(len(active_urls), 0) results = [] for url in active_urls: diff --git a/xtuner/v1/rl/agent_loop_manager/producer.py b/xtuner/v1/rl/agent_loop_manager/producer.py index 77be52c794..2f7f660e70 100644 --- a/xtuner/v1/rl/agent_loop_manager/producer.py +++ b/xtuner/v1/rl/agent_loop_manager/producer.py @@ -205,6 +205,9 @@ class AsyncProduceStrategyConfig(ProduceStrategyConfig): rerolls out immediately without entering tail-batch mode, and ``N > 0`` waits until the expired pool contains at least ``N`` groups before entering tail-batch mode. + max_pending_tasks (int | None): Maximum number of concurrently pending + rollout groups in one produce_batch call. Defaults to None, which + keeps the existing unbounded scheduling behavior. **Examples:** @@ -222,6 +225,7 @@ class AsyncProduceStrategyConfig(ProduceStrategyConfig): max_staleness: int = Field(default=0, ge=0) max_token_staleness: int | None = Field(default=None, ge=0) tail_batch_trigger_size: int = Field(default=-1, ge=-1) + max_pending_tasks: int | None = Field(default=None, gt=0) def build( self, @@ -250,6 +254,7 @@ def build( max_token_staleness=self.max_token_staleness, sync_weights_interval=sync_weights_interval, tail_batch_trigger_size=self.tail_batch_trigger_size, + max_pending_tasks=self.max_pending_tasks, should_continue_fn=self.should_continue_fn, ) @@ -324,6 +329,7 @@ def __init__( over_sample_threshold: float, enable_partial_rollout: bool, tail_batch_trigger_size: int, + max_pending_tasks: int | None, max_staleness: int, max_token_staleness: int | None, sync_weights_interval: int, @@ -353,6 +359,7 @@ def __init__( else calculate_stale_threshold(max_token_staleness, sync_weights_interval) ) self.tail_batch_trigger_size = tail_batch_trigger_size + self.max_pending_tasks = max_pending_tasks self._local_pending_tasks: set[asyncio.Task] = set() def pending_task_count(self) -> int: @@ -419,7 +426,9 @@ async def spawn_one() -> asyncio.Task: pending_count = len(self._local_pending_tasks) desired_pending = max(0, scheduled_target - available) - if available + pending_count < scheduled_target: + if self.max_pending_tasks is not None: + desired_pending = min(desired_pending, self.max_pending_tasks) + if pending_count < desired_pending: while len(self._local_pending_tasks) < desired_pending: self._local_pending_tasks.add(await spawn_one()) diff --git a/xtuner/v1/rl/health_manager.py b/xtuner/v1/rl/health_manager.py new file mode 100644 index 0000000000..4d0754d787 --- /dev/null +++ b/xtuner/v1/rl/health_manager.py @@ -0,0 +1,240 @@ +from __future__ import annotations + +import threading +from contextlib import contextmanager +from typing import TYPE_CHECKING + +import ray + +from xtuner.v1.rl.rollout.worker_registry import WorkerLifecycleState +from xtuner.v1.utils import get_logger + + +if TYPE_CHECKING: + from xtuner.v1.rl.rollout.controller import RolloutControllerProxy + from xtuner.v1.rl.rollout.worker import RolloutConfig + from xtuner.v1.rl.trainer.controller import TrainingController + + +RL_HEALTH_MANAGER_RAY_GET_TIMEOUT = 3600 +RL_HEALTH_MANAGER_STOP_JOIN_TIMEOUT = 30.0 +PENDING_ROLLOUT_WORKER_CHECK_INTERVAL = 1.0 +ROLLOUT_WEIGHT_UPDATE_DRAIN_TIMEOUT = 600.0 + + +class _NoOpRLHealthManager: + """No-op implementation for RLTrainer without rollout health management. + + Debug rollout and debug train modes do not initialize a real + ``RLHealthManager``. This implementation preserves the same interface so + callers can use the health manager without conditional checks. + """ + + def start(self): + pass + + def stop(self): + pass + + def set_rollout_resources_available(self, available: bool): + pass + + @contextmanager + def weight_update_guard(self): + yield + + +class RLHealthManager: + """Coordinate driver-side recovery of colocated rollout workers. + + RolloutHealthManager owns worker restart and moves successfully restarted workers to PENDING_WEIGHTS. This manager + polls pending workers on the driver and updates them from the latest registered checkpoint before promoting them to + ACTIVE. + """ + + def __init__( + self, + *, + train_controller: TrainingController, + rollout_controller: RolloutControllerProxy, + rollout_config: RolloutConfig, + ) -> None: + self.enable_pending_weight_recovery = ( + rollout_config is not None and rollout_config.weight_transport_type == "checkpoint_engine" + ) + self.train_controller = train_controller + self.rollout_controller = rollout_controller + self._rollout_config = rollout_config + self._rollout_resources_available = threading.Event() + self._rollout_weight_update_lock = threading.Lock() + self._pending_rollout_weight_update_stop_event = threading.Event() + self._pending_rollout_weight_update_thread: threading.Thread | None = None + self.logger = get_logger(tag="RLHealthManager") + + def _check_enabled_dependencies(self) -> None: + if not self.enable_pending_weight_recovery: + return + if self.train_controller is None or self.rollout_controller is None or self._rollout_config is None: + raise RuntimeError("RLHealthManager is missing Checkpoint Engine recovery dependencies.") + + def start(self) -> None: + if not self.enable_pending_weight_recovery: + return + self._check_enabled_dependencies() + if self._pending_rollout_weight_update_thread is not None: + return + + self._pending_rollout_weight_update_stop_event.clear() + self._pending_rollout_weight_update_thread = threading.Thread( + target=self._pending_rollout_worker_weight_update_loop, + name="pending-rollout-weight-update", + daemon=True, + ) + self._pending_rollout_weight_update_thread.start() + self.logger.info("Started pending rollout checkpoint-engine update thread.") + + def stop(self) -> None: + if not self.enable_pending_weight_recovery: + return + self._check_enabled_dependencies() + self._pending_rollout_weight_update_stop_event.set() + + thread = self._pending_rollout_weight_update_thread + if thread is not None: + thread.join(timeout=RL_HEALTH_MANAGER_STOP_JOIN_TIMEOUT) + if thread.is_alive(): + self.logger.warning( + "Pending rollout weight update thread did not stop before " + f"timeout={RL_HEALTH_MANAGER_STOP_JOIN_TIMEOUT}s." + ) + return + + self._pending_rollout_weight_update_thread = None + self.logger.info("Stopped pending rollout checkpoint-engine update thread.") + + def set_rollout_resources_available(self, available: bool) -> None: + # If rollout resources are unavailable, shutdown any inactive rollout workers to free up resources for training. + if not available: + ray.get( + self.rollout_controller.shutdown_inactive_workers.remote(), + timeout=RL_HEALTH_MANAGER_RAY_GET_TIMEOUT, + ) + + if not self.enable_pending_weight_recovery: + return + if available: + self._rollout_resources_available.set() + return + + # Stop admitting new pending-worker updates, then wait for an update + # that already owns the lock to finish before training reuses the + # colocated rollout resources. + self._rollout_resources_available.clear() + acquired = self._rollout_weight_update_lock.acquire(timeout=ROLLOUT_WEIGHT_UPDATE_DRAIN_TIMEOUT) + if not acquired: + raise TimeoutError( + "Timed out waiting for pending rollout weight update before switching to training: " + f"timeout={ROLLOUT_WEIGHT_UPDATE_DRAIN_TIMEOUT}s." + ) + self._rollout_weight_update_lock.release() + + @contextmanager + def weight_update_guard(self): + """Serialize normal and recovery Checkpoint Engine weight updates.""" + if not self.enable_pending_weight_recovery: + yield + return + self._check_enabled_dependencies() + with self._rollout_weight_update_lock: + yield + + def _update_pending_rollout_weights_from_checkpoint_engine(self) -> tuple[tuple[int, ...], ...]: + """Update every currently pending rollout group from Checkpoint + Engine.""" + self._check_enabled_dependencies() + assert self.rollout_controller is not None + assert self.train_controller is not None + assert self._rollout_config is not None + if not self.train_controller.has_registered_weight_checkpoint(): + self.logger.info( + "Defer pending rollout checkpoint-engine update because no train checkpoint has been registered yet." + ) + return () + + pending_targets, pending_group_ranks = ray.get( + self.rollout_controller.get_weight_update_targets.remote( + target_state=WorkerLifecycleState.PENDING_WEIGHTS, + return_group_ranks=True, + ), + timeout=RL_HEALTH_MANAGER_RAY_GET_TIMEOUT, + ) + if not pending_targets: + return () + + try: + self.logger.info( + "Updating pending rollout workers from Checkpoint Engine: " + f"group_ranks={pending_group_ranks}, targets={pending_targets}." + ) + self.train_controller.bind_rollout_weight_update( + targets=pending_targets, + rollout_config=self._rollout_config, + ) + ray.get( + self.rollout_controller.onload_weights.remote(target_state=WorkerLifecycleState.PENDING_WEIGHTS), + timeout=RL_HEALTH_MANAGER_RAY_GET_TIMEOUT, + ) + self.train_controller.weight_update(need_register=False, need_update=True) + ray.get( + self.rollout_controller.onload_kvcache.remote(target_state=WorkerLifecycleState.PENDING_WEIGHTS), + timeout=RL_HEALTH_MANAGER_RAY_GET_TIMEOUT, + ) + ray.get( + self.rollout_controller.mark_worker_groups_lifecycle_state.remote( + group_ranks=list(pending_group_ranks), + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + target_state=WorkerLifecycleState.ACTIVE, + ), + timeout=RL_HEALTH_MANAGER_RAY_GET_TIMEOUT, + ) + self.logger.info( + f"Recovered rollout workers weight updated from Checkpoint Engine: {pending_group_ranks}." + ) + return tuple(pending_group_ranks) + except Exception: + self.logger.exception( + f"Failed to update recovered rollout workers weight from Checkpoint Engine: {pending_group_ranks}." + ) + ray.get( + self.rollout_controller.mark_worker_groups_lifecycle_state.remote( + group_ranks=list(pending_group_ranks), + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + target_state=WorkerLifecycleState.INACTIVE, + ), + timeout=RL_HEALTH_MANAGER_RAY_GET_TIMEOUT, + ) + return () + + def _pending_rollout_worker_weight_update_loop(self) -> None: + while not self._pending_rollout_weight_update_stop_event.wait(PENDING_ROLLOUT_WORKER_CHECK_INTERVAL): + if not self._rollout_resources_available.is_set(): + continue + if not self._rollout_weight_update_lock.acquire(blocking=False): + continue + try: + # The rollout phase may have ended after the check above but + # before this thread acquired the lock. + if not self._rollout_resources_available.is_set(): + continue + updated_groups = self._update_pending_rollout_weights_from_checkpoint_engine() + if updated_groups: + self.logger.info( + f"Background pending rollout checkpoint-engine update completed: {updated_groups}." + ) + except Exception: + self.logger.exception("Background pending rollout weight update failed.") + finally: + self._rollout_weight_update_lock.release() + + +__all__ = ["RLHealthManager"] diff --git a/xtuner/v1/rl/rollout/controller.py b/xtuner/v1/rl/rollout/controller.py index 3db2874d03..cfa52ee13f 100644 --- a/xtuner/v1/rl/rollout/controller.py +++ b/xtuner/v1/rl/rollout/controller.py @@ -22,7 +22,7 @@ RolloutConfig, get_rollout_worker_base_cls, ) -from .worker_registry import RolloutWorkerRegistry +from .worker_registry import RolloutWorkerRegistry, WorkerLifecycleState # Keep this as a Ray actor because Ray AgentLoop actors need a shared, cross-process handle to the same controller @@ -66,9 +66,52 @@ def __init__( ) self.health_manager.start() - def get_weight_update_targets(self) -> tuple[RolloutWeightUpdateTarget, ...]: - """Return rollout endpoints that can receive weight update requests.""" - return self.registry.weight_update_targets() + def get_weight_update_targets( + self, target_state: WorkerLifecycleState | None = None, return_group_ranks: bool = False + ) -> ( + list[RolloutWeightUpdateTarget] + | tuple[ + list[RolloutWeightUpdateTarget], + list[tuple[int, ...]], + ] + ): + """Return rollout weight-update targets and their lifecycle groups.""" + + target_states: tuple[WorkerLifecycleState, ...] + if target_state is None: + target_states = ( + WorkerLifecycleState.PENDING_WEIGHTS, + WorkerLifecycleState.ACTIVE, + WorkerLifecycleState.INACTIVE, + ) + else: + target_states = (target_state,) + target_state_values = {state.value for state in target_states} + targets, group_ranks = self.registry.weight_update_targets() + + filtered_targets = [target for target in targets if target.lifecycle_state in target_state_values] + + if not return_group_ranks: + return filtered_targets + + endpoint_ranks = {target.endpoint_rank for target in filtered_targets} + filtered_group_ranks = [ranks for ranks in group_ranks if endpoint_ranks.intersection(ranks)] + + return filtered_targets, filtered_group_ranks + + def inject_backend_crash_for_test(self, *, rank: int = 0) -> None: + """Crash one active rollout backend for the immediate-recovery test.""" + worker = self.registry.active_entrypoint_by_rank(rank) + if worker is None: + raise RuntimeError(f"No active rollout request entrypoint found for test fault injection: rank={rank}.") + + accepted = ray.get( + worker.actor.inject_backend_crash_for_test.remote(), # type: ignore[attr-defined] + timeout=ROLLOUT_RAY_GET_TIMEOUT, + ) + if not accepted: + raise RuntimeError(f"Rollout worker rejected test fault injection: rank={rank}, url={worker.url}.") + self.logger.warning(f"[ImmediateRecoveryExperiment] backend_crash_injected rank={rank} url={worker.url}") def register_active_workers_to_proxy(self) -> None: if self.proxy_manager is None: @@ -133,6 +176,9 @@ def set_enable_partial_rollout(self, enable: bool) -> None: def pause_generation(self): self.health_manager.pause() + # Wait for the health manager to finish recovery before pausing generation. + if not self.health_manager.wait_recovery_done(timeout=600.0): + raise TimeoutError("Timed out waiting for rollout worker recovery before training.") active_workers = self.registry.active_workers() futures = [ worker.actor.pause_generation.remote() # type: ignore[attr-defined] @@ -152,34 +198,65 @@ def pause_generation(self): if failed_worker_urls: self.logger.warning(f"Abort request failed: worker_urls={failed_worker_urls}") - async def check_and_shutdown_inactive_workers(self): - """Run a fail-fast health barrier and shut down failed groups so - training can reuse shared rollout resources.""" - await asyncio.to_thread(self.health_manager.check_and_shutdown_inactive_workers) + async def shutdown_inactive_workers(self): + """Shut down failed groups so training can reuse shared rollout + resources.""" + await asyncio.to_thread(self.health_manager.shutdown_inactive_workers) async def restart_inactive_workers(self): """Restart inactive groups before a sync-step weight update.""" - await asyncio.to_thread(self.health_manager.restart_inactive_workers) + groups = await asyncio.to_thread(self.health_manager.restart_inactive_workers) + return tuple(group.ranks for group in groups) + + def mark_worker_groups_lifecycle_state( + self, + group_ranks: list[tuple[int, ...]] | None = None, + *, + source_state: WorkerLifecycleState, + target_state: WorkerLifecycleState, + ) -> None: + """Move selected worker groups from source_state to target_state. + + When group_ranks is omitted, every complete worker group currently in source_state is moved. When it is + provided, only exact matching groups are considered. Transitions to ACTIVE or INACTIVE notify the health + manager so routing and lifecycle listeners stay in sync. + """ + source_groups = self.registry.get_target_state_worker_groups(source_state) + # 只对目标group中命中状态的worker进行状态更新,若不提供目标group,则对所有source_state状态的worker进行状态更新 + if group_ranks is None: + groups = source_groups + else: + groups_by_ranks = {group.ranks: group for group in source_groups} + groups = tuple(groups_by_ranks[ranks] for ranks in group_ranks if ranks in groups_by_ranks) + updated_groups = self.registry.set_groups_state( + groups, + target_state, + source_state=source_state, + ) + if target_state is WorkerLifecycleState.ACTIVE: + self.health_manager.notify_worker_group_active(updated_groups) + elif target_state is WorkerLifecycleState.INACTIVE: + self.health_manager.notify_worker_group_inactive(updated_groups) def continue_generation(self): - self._broadcast_to_active_workers("continue_generation") + self._broadcast_to_workers("continue_generation", WorkerLifecycleState.ACTIVE) self.health_manager.resume() def offload(self): - self._broadcast_to_active_workers("offload") + self._broadcast_to_workers("offload", WorkerLifecycleState.ACTIVE) def flush_cache(self): self._broadcast_to_active_workers("flush_cache") def onload(self): - self._broadcast_to_active_workers("onload_weights") - self._broadcast_to_active_workers("onload_kvcache") + self._broadcast_to_workers("onload_weights", WorkerLifecycleState.ACTIVE) + self._broadcast_to_workers("onload_kvcache", WorkerLifecycleState.ACTIVE) - def onload_weights(self): - self._broadcast_to_active_workers("onload_weights") + def onload_weights(self, target_state: WorkerLifecycleState = WorkerLifecycleState.ACTIVE): + self._broadcast_to_workers("onload_weights", target_state) - def onload_kvcache(self): - self._broadcast_to_active_workers("onload_kvcache") + def onload_kvcache(self, target_state: WorkerLifecycleState = WorkerLifecycleState.ACTIVE): + self._broadcast_to_workers("onload_kvcache", target_state) def shutdown(self): """Shut down all rollout workers tracked by the controller.""" @@ -190,8 +267,8 @@ def shutdown(self): timeout=ROLLOUT_RAY_GET_TIMEOUT, ) - def _broadcast_to_active_workers(self, method_name: str, **kwargs): - workers = self.registry.active_workers() + def _broadcast_to_workers(self, method_name: str, target_state: WorkerLifecycleState, **kwargs): + workers = self.registry.get_target_state_workers(target_state) futures = [getattr(worker.actor, method_name).remote(**kwargs) for worker in workers] return ray.get(futures, timeout=ROLLOUT_RAY_GET_TIMEOUT) diff --git a/xtuner/v1/rl/rollout/health_manager.py b/xtuner/v1/rl/rollout/health_manager.py index 2026cd655a..47028735bb 100644 --- a/xtuner/v1/rl/rollout/health_manager.py +++ b/xtuner/v1/rl/rollout/health_manager.py @@ -14,7 +14,7 @@ from xtuner.v1.utils import get_logger -from .worker_registry import RolloutWorkerRegistry, WorkerGroup, WorkerSnapshot +from .worker_registry import RolloutWorkerRegistry, WorkerGroup, WorkerLifecycleState, WorkerSnapshot if TYPE_CHECKING: @@ -37,7 +37,7 @@ class RolloutWorkerLifecycleListener(Protocol): def on_worker_group_inactive(self, group: WorkerGroup) -> None: ... - def on_worker_group_recovered(self, group: WorkerGroup) -> None: ... + def on_worker_group_active(self, group: WorkerGroup) -> None: ... class _HealthManagerStopping(InterruptedError): @@ -48,8 +48,8 @@ class _HealthManagerStopping(InterruptedError): class _WorkerHealthFailureTracker: """Track per-rank health-check failures and decide when a rank fails. - Periodic checks call update_failed_ranks() to apply the configured threshold. Explicit shutdown barriers call - mark_failed_ranks() to fail unhealthy ranks immediately while still keeping failure-count bookkeeping in one place. + Periodic checks call update_failed_ranks() and report workers whose consecutive failures reach the configured + threshold. """ threshold: int @@ -84,22 +84,6 @@ def update_failed_ranks(self, worker_health_results: dict[int, bool]) -> set[int return failed_ranks - def mark_failed_ranks(self, worker_health_results: dict[int, bool]) -> set[int]: - failed_ranks: set[int] = set() - for rank, is_healthy in worker_health_results.items(): - if is_healthy: - self.failure_counts.pop(rank, None) - continue - - failure_count = self._record_failure(rank) - logger.warning( - f"Worker {rank} failed explicit health check and will be marked inactive " - f"immediately: failure_count={failure_count}." - ) - failed_ranks.add(rank) - - return failed_ranks - class RolloutHealthManager: """Own worker health state and recovery after controller startup. @@ -124,6 +108,9 @@ def __init__( self._stop_event = threading.Event() self._pause_event = threading.Event() self._pause_event.set() + # 恢复是否已经完成 + self._recovery_done = threading.Event() + self._recovery_done.set() self._thread: threading.Thread | None = None self._lifecycle_operation_lock = threading.Lock() self._worker_health_failure_tracker = _WorkerHealthFailureTracker(threshold=self._check_failure_threshold) @@ -175,6 +162,15 @@ def resume(self) -> None: self._pause_event.clear() logger.info("RolloutHealthManager resumed.") + def wait_recovery_done(self, timeout: float) -> bool: + """Wait for ongoing worker health and worker recovery.""" + deadline = time.monotonic() + timeout + if not self._lifecycle_operation_lock.acquire(timeout=timeout): + return False + self._lifecycle_operation_lock.release() + # 等待 recovery 完成,最多等待剩余的 timeout 时间 + return self._recovery_done.wait(timeout=max(0.0, deadline - time.monotonic())) + # ------------------------------------------------------------------ # Public health and lifecycle workflows # ------------------------------------------------------------------ @@ -191,6 +187,8 @@ def run_once(self) -> None: failed_groups: tuple[WorkerGroup, ...] = () logger.debug("RolloutHealthManager running health checks for active workers.") try: + if self._pause_event.is_set(): + return worker_health_results = self._check_active_workers_health() failed_ranks = self._worker_health_failure_tracker.update_failed_ranks(worker_health_results) if not failed_ranks: @@ -200,6 +198,8 @@ def run_once(self) -> None: except _HealthManagerStopping: return failed_groups = self._registry.mark_unhealthy_ranks(failed_ranks) + # 标记有重启中的worker,正在等待完成 + self._recovery_done.clear() finally: self._lifecycle_operation_lock.release() @@ -210,46 +210,30 @@ def run_once(self) -> None: event_name="inactive", notify_listener=lambda listener, group: listener.on_worker_group_inactive(group), ) + # TODO: Recovery runs synchronously on the health-check thread, so the next + # periodic health check waits until this restart finishes. Move restart to another thread. + try: + self._restart_inactive_workers() + finally: + # 标记恢复已完成 + self._recovery_done.set() - def restart_inactive_workers(self) -> None: + def restart_inactive_workers(self) -> tuple[WorkerGroup, ...]: """Synchronously restart inactive groups before the next sync-step weight update.""" - recovered_groups: list[WorkerGroup] = [] - groups_to_recover: tuple[WorkerGroup, ...] = () - + self._recovery_done.clear() try: - with self._paused_lifecycle_operation(): - groups_to_recover = self._registry.claim_inactive_groups_for_recovery() - if groups_to_recover: - recovered_groups = self._restart_claimed_recovery_groups(groups_to_recover) - except _HealthManagerStopping: - return - - if not groups_to_recover: - logger.info("No failed rollout workers detected during recovery.") - return - - self._notify_worker_lifecycle_listeners( - recovered_groups, - event_name="recovered", - notify_listener=lambda listener, group: listener.on_worker_group_recovered(group), - ) - inactive_workers = [f"rank={worker.rank}, url={worker.url}" for worker in self._registry.inactive_workers()] - if inactive_workers: - logger.error("inactive rollout workers before sync-step weight update: " + ", ".join(inactive_workers)) + return self._restart_inactive_workers() + finally: + self._recovery_done.set() - def check_and_shutdown_inactive_workers(self) -> None: - """Fail-fast health-check active workers, mark failures inactive, and - shut down every non-active group so shared resources can be reused by - training.""" + def shutdown_inactive_workers(self) -> None: + """Shut down every non-active group so shared resources can be reused + by training.""" groups_to_shutdown: tuple[WorkerGroup, ...] = () try: with self._paused_lifecycle_operation(): - worker_health_results = self._check_active_workers_health() - self._checkpoint_not_stopping() - self._mark_unhealthy_worker_groups_inactive(worker_health_results) - self._checkpoint_not_stopping() groups_to_shutdown = self._registry.inactive_worker_groups() for group in groups_to_shutdown: self._shutdown_worker_group(group) @@ -374,15 +358,6 @@ async def probe_workers(): return worker_health_results - def _mark_unhealthy_worker_groups_inactive(self, worker_health_results: dict[int, bool]) -> None: - failed_ranks = self._worker_health_failure_tracker.mark_failed_ranks(worker_health_results) - if not failed_ranks: - return - - inactive_groups = self._registry.mark_unhealthy_ranks(failed_ranks) - for group in inactive_groups: - logger.warning(f"Rollout worker group ranks={group.ranks} failed health check. Marking as inactive.") - # ------------------------------------------------------------------ # Worker group recovery state # ------------------------------------------------------------------ @@ -394,15 +369,39 @@ def _restart_claimed_recovery_groups(self, groups: tuple[WorkerGroup, ...]) -> l group_recovery_results = self._restart_worker_groups(groups) self._checkpoint_not_stopping() - recovered_groups: list[WorkerGroup] = [] + pending_weights_groups: list[WorkerGroup] = [] for group in groups: recovered = group_recovery_results.get(group.ranks, False) - recorded_group = self._registry.set_group_recovery_result(group, recovered=recovered) if recovered: + # The recovered server can accept a weight update, but must + # stay out of rollout routing until the trainer pushes weights. + logger.info( + "[recovery-test] marking recovered rollout worker group pending_weight_update: " + f"ranks={group.ranks}, workers=[" + + ", ".join( + f"rank={worker.rank}, url={worker.url}, state={worker.lifecycle_state.value}" + for worker in group.workers + ) + + "]" + ) + recorded_group = self._registry.set_groups_state( + groups=(group,), + target_state=WorkerLifecycleState.PENDING_WEIGHTS, + )[0] + logger.info( + "[recovery-test] marked rollout worker group pending_weight_update: " + f"ranks={recorded_group.ranks}, workers=[" + + ", ".join( + f"rank={worker.rank}, url={worker.url}, state={worker.lifecycle_state.value}" + for worker in recorded_group.workers + ) + + "]" + ) self._worker_health_failure_tracker.clear(group.ranks) groups_needing_cleanup.pop(group.ranks, None) - recovered_groups.append(recorded_group) + pending_weights_groups.append(recorded_group) else: + recorded_group = self._registry.set_group_recovery_result(group, recovered=False) groups_needing_cleanup.pop(group.ranks, None) logger.error( "Failed to restart rollout worker group; training can continue with remaining active " @@ -411,7 +410,7 @@ def _restart_claimed_recovery_groups(self, groups: tuple[WorkerGroup, ...]) -> l + ", ".join(f"rank={worker.rank}, url={worker.url}" for worker in recorded_group.workers) + "]" ) - return recovered_groups + return pending_weights_groups except BaseException: self._cleanup_unfinalized_recovery_groups(tuple(groups_needing_cleanup.values())) raise @@ -483,7 +482,7 @@ def _restart_worker_group( self._checkpoint_not_stopping() with self._skip_load_weights_during_restart(group): self._checkpoint_not_stopping() - ray.get( + init_results = ray.get( [ # reinit() reuses the server launch spec bound during # controller startup. @@ -492,9 +491,21 @@ def _restart_worker_group( ], timeout=ROLLOUT_RAY_GET_TIMEOUT, ) + logger.info( + "[recovery-test] reinit returned for rollout worker group " + f"ranks={group.ranks}, init_results=[" + + ", ".join( + f"rank={result.rank}, server_url={result.server_url}, session_url={result.session_url}" + for result in init_results + ) + + "]" + ) self._checkpoint_not_stopping() health_results = self._check_workers_health(group.workers) + logger.info( + f"[recovery-test] post-reinit health results for group ranks={group.ranks}: {health_results}" + ) unhealthy_ranks = [ worker.rank for worker in group.workers if not health_results.get(worker.rank, False) ] @@ -596,6 +607,44 @@ def _wait_worker_server_down(self, worker: WorkerSnapshot, *, max_wait_attempts: return False + def notify_worker_group_active(self, groups: Iterable[WorkerGroup]) -> None: + self._notify_worker_lifecycle_listeners( + groups, + event_name="active", + notify_listener=lambda listener, group: listener.on_worker_group_active(group), + ) + + def notify_worker_group_inactive(self, groups: Iterable[WorkerGroup]) -> None: + self._notify_worker_lifecycle_listeners( + groups, + event_name="inactive", + notify_listener=lambda listener, group: listener.on_worker_group_inactive(group), + ) + + def _restart_inactive_workers(self) -> tuple[WorkerGroup, ...]: + pending_weights_groups: tuple[WorkerGroup, ...] = () + groups_to_recover: tuple[WorkerGroup, ...] = () + + try: + with self._paused_lifecycle_operation(): + groups_to_recover = self._registry.get_inactive_groups_for_recovery() + if groups_to_recover: + pending_weights_groups = tuple(self._restart_claimed_recovery_groups(groups_to_recover)) + except _HealthManagerStopping: + return () + + if not groups_to_recover: + return () + + inactive_workers = [ + f"rank={worker.rank}, url={worker.url}" + for worker in self._registry.inactive_workers() + if worker.lifecycle_state is WorkerLifecycleState.INACTIVE + ] + if inactive_workers: + logger.error("inactive rollout workers before sync-step weight update: " + ", ".join(inactive_workers)) + return pending_weights_groups + # ------------------------------------------------------------------ # Worker lifecycle notifications # ------------------------------------------------------------------ diff --git a/xtuner/v1/rl/rollout/proxy_manager.py b/xtuner/v1/rl/rollout/proxy_manager.py index fdd70c1742..768f4acb6e 100644 --- a/xtuner/v1/rl/rollout/proxy_manager.py +++ b/xtuner/v1/rl/rollout/proxy_manager.py @@ -57,8 +57,8 @@ def on_worker_group_inactive(self, group: "WorkerGroup") -> None: if worker.is_request_entrypoint: self._delete_session_url(worker.session_url) - def on_worker_group_recovered(self, group: "WorkerGroup") -> None: - """Register recovered request entrypoints to routed API proxy.""" + def on_worker_group_active(self, group: "WorkerGroup") -> None: + """Register active request entrypoints to routed API proxy.""" for worker in group.workers: if worker.is_request_entrypoint: self._register_session_url(worker.session_url) diff --git a/xtuner/v1/rl/rollout/worker.py b/xtuner/v1/rl/rollout/worker.py index f3d828cf6b..7357abad71 100644 --- a/xtuner/v1/rl/rollout/worker.py +++ b/xtuner/v1/rl/rollout/worker.py @@ -733,6 +733,46 @@ def shutdown(self, *, stop_session_server: bool = False): self.logger.debug(f"Worker {self.rank} server process and its children terminated.") return + def inject_backend_crash_for_test(self) -> bool: + """Force-stop the backend server for the immediate-recovery test.""" + if os.environ.get("XTUNER_TEST_IMMEDIATE_RECOVERY", "0") != "1": + raise RuntimeError("Rollout test fault injection requires XTUNER_TEST_IMMEDIATE_RECOVERY=1.") + self.logger.warning( + f"[ImmediateRecoveryExperiment] crashing_backend_server rank={self.rank} url={self.server_url}" + ) + + if self.server_task is not None: + server_task = self.server_task + ray.cancel(server_task, force=True, recursive=True) + try: + ray.get(server_task, timeout=60) + except ray.exceptions.GetTimeoutError: + self.logger.warning(f"Worker {self.rank} server task did not stop within crash timeout.") + raise + except Exception as e: + self.logger.debug(f"Worker {self.rank} server task stopped after injected crash: {e}") + self.server_task = None + return True + + if self.server_process is not None: + import psutil + + try: + parent = psutil.Process(self.server_process.pid) + except psutil.NoSuchProcess: + self.server_process = None + return True + children = parent.children(recursive=True) + for child in children: + child.kill() + parent.kill() + parent.wait(timeout=5) + self.server_process = None + self.logger.debug(f"Worker {self.rank} server process and its children killed.") + return True + + return False + def _start_session_server(self) -> None: """Start the per-worker SessionServer proxy.""" assert self.server_launch_spec is not None diff --git a/xtuner/v1/rl/rollout/worker_registry.py b/xtuner/v1/rl/rollout/worker_registry.py index 4452af4b2a..ccdffca5f5 100644 --- a/xtuner/v1/rl/rollout/worker_registry.py +++ b/xtuner/v1/rl/rollout/worker_registry.py @@ -30,8 +30,8 @@ class WorkerLifecycleState(str, Enum): ACTIVE = "active" # Not serving rollout requests; the rollout server may still hold resources. INACTIVE = "inactive" - # Temporarily owned by recovery shutdown/init/check_health. - RECOVERING = "recovering" + # Server is healthy after recovery, but waiting for trainer-side weights.. + PENDING_WEIGHTS = "pending_weights" @dataclass(frozen=True) @@ -122,8 +122,24 @@ def all_actors(self) -> tuple[RolloutWorker, ...]: def active_workers(self) -> tuple[WorkerSnapshot, ...]: """Return workers whose lifecycle state is active.""" + return self.get_target_state_workers(WorkerLifecycleState.ACTIVE) + + def get_target_state_workers(self, target_state: WorkerLifecycleState) -> tuple[WorkerSnapshot, ...]: + """Return workers matching the requested lifecycle state.""" with self._lock: - return tuple(worker for worker in self._workers.values() if worker.is_active()) + return tuple(worker for worker in self._workers.values() if worker.lifecycle_state is target_state) + + def get_target_state_worker_groups(self, target_state: WorkerLifecycleState) -> tuple[WorkerGroup, ...]: + """Return lifecycle groups containing workers in the requested + state.""" + with self._lock: + worker_groups = self._build_worker_groups() + matched_groups = [ + group + for group in worker_groups.values() + if any(worker.lifecycle_state is target_state for worker in group.workers) + ] + return tuple(sorted(matched_groups, key=lambda group: group.ranks)) def active_entrypoints(self) -> tuple[WorkerSnapshot, ...]: """Return active workers that can receive rollout generation @@ -171,8 +187,8 @@ def inactive_worker_groups(self) -> tuple[WorkerGroup, ...]: ] return tuple(sorted(inactive_groups, key=lambda group: group.ranks)) - def claim_inactive_groups_for_recovery(self) -> tuple[WorkerGroup, ...]: - """Claim inactive worker groups by moving them to RECOVERING state.""" + def get_inactive_groups_for_recovery(self) -> tuple[WorkerGroup, ...]: + """Return inactive worker groups selected for recovery.""" with self._lock: worker_groups = self._build_worker_groups() inactive_groups = [ @@ -186,13 +202,9 @@ def claim_inactive_groups_for_recovery(self) -> tuple[WorkerGroup, ...]: worker.rank for worker in group.workers if worker.lifecycle_state is WorkerLifecycleState.INACTIVE ) logger.warning( - f"Claimed inactive rollout worker ranks={inactive_ranks} " + f"Selected inactive rollout worker ranks={inactive_ranks} " f"in worker_group_ranks={group.ranks} for recovery." ) - for rank in group.ranks: - worker = self._workers.get(rank) - if worker is not None: - self._workers[rank] = replace(worker, lifecycle_state=WorkerLifecycleState.RECOVERING) return sorted_groups def mark_unhealthy_ranks(self, ranks: set[int]) -> tuple[WorkerGroup, ...]: @@ -236,8 +248,41 @@ def set_group_recovery_result( ) return recorded_group - def weight_update_targets(self) -> tuple[RolloutWeightUpdateTarget, ...]: - """Return weight-update targets resolved with current runtime state.""" + def set_groups_state( + self, + groups: Iterable[WorkerGroup], + target_state: WorkerLifecycleState, + *, + source_state: WorkerLifecycleState | None = None, + ) -> tuple[WorkerGroup, ...]: + """Move worker groups from source_state to target_state. + + If source_state is provided, only workers currently in that state are updated. + """ + with self._lock: + groups = tuple(groups) + for group in groups: + for rank in group.ranks: + worker = self._workers.get(rank) + if worker is not None and (source_state is None or worker.lifecycle_state is source_state): + self._workers[rank] = replace(worker, lifecycle_state=target_state) + worker_groups = self._build_worker_groups() + recorded_groups = [] + for group in groups: + recorded_group = worker_groups.get(group.ranks) + if recorded_group is None: + continue + recorded_groups.append(recorded_group) + return tuple(recorded_groups) + + def weight_update_targets( + self, + ) -> tuple[ + tuple[RolloutWeightUpdateTarget, ...], + tuple[tuple[int, ...], ...], + ]: + """Return weight-update targets resolved with current runtime state and + their lifecycle group ranks.""" from xtuner.v1.rl.weight_update.data import RolloutWeightUpdateTarget with self._lock: @@ -256,4 +301,17 @@ def weight_update_targets(self) -> tuple[RolloutWeightUpdateTarget, ...]: lifecycle_state=worker.lifecycle_state.value, ) ) - return tuple(sorted(targets, key=lambda target: target.endpoint_rank)) + sorted_targets = tuple( + sorted( + targets, + key=lambda target: target.endpoint_rank, + ) + ) + group_ranks = tuple( + dict.fromkeys( + self._rollout_topology.lifecycle_group_for_server_rank(target.endpoint_rank) + for target in sorted_targets + ) + ) + + return sorted_targets, group_ranks diff --git a/xtuner/v1/rl/trainer/controller.py b/xtuner/v1/rl/trainer/controller.py index 3e005b1d22..0cc6b9e322 100644 --- a/xtuner/v1/rl/trainer/controller.py +++ b/xtuner/v1/rl/trainer/controller.py @@ -338,6 +338,10 @@ def weight_update(self, **kwargs): ray.get(handles, timeout=TRAIN_RAY_GET_TIMEOUT) return + def has_registered_weight_checkpoint(self) -> bool: + handles = [worker.has_registered_weight_checkpoint.remote() for worker in self.workers] + return all(ray.get(handles, timeout=TRAIN_RAY_GET_TIMEOUT)) + def suspend_train_nccl_process_groups(self): """Suspend train-side NCCL process groups after weight sync.""" handles = [ diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index 80929e9fff..982ed40110 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -337,6 +337,10 @@ def bind_rollout_weight_update(self, *args, **kwargs): def weight_update(self, **kwargs): return self.update_weighter.weight_update(**kwargs) + @ray_method + def has_registered_weight_checkpoint(self) -> bool: + return self.update_weighter.has_registered_weight_checkpoint() + def _init_sft(self, worker_cfg: WorkerConfig): self._sft_dataloader_config = worker_cfg.sft_dataloader_cfg self._sft_dataloader: Dataloader | None = None diff --git a/xtuner/v1/rl/weight_update/data.py b/xtuner/v1/rl/weight_update/data.py index 5355f4bcfa..fbe2e46479 100644 --- a/xtuner/v1/rl/weight_update/data.py +++ b/xtuner/v1/rl/weight_update/data.py @@ -68,10 +68,6 @@ class RolloutWeightUpdateTarget: # Registry lifecycle state value for this endpoint. lifecycle_state: str - @property - def is_active(self) -> bool: - return self.lifecycle_state == "active" - @property def engine_size(self) -> int: return len(self.update_ranks) @@ -144,7 +140,7 @@ def local_update_target(self) -> RolloutWeightUpdateTarget | None: @property def rollout_url(self) -> str | None: target = self.local_update_target - if target is None or not target.is_active: + if target is None: return None return target.server_url @@ -174,14 +170,25 @@ def ipc_engine_parallel_size(self) -> int | None: return target.engine_size @property - def active_update_targets(self) -> tuple[RolloutWeightUpdateTarget, ...]: - return tuple(target for target in self.weight_update_targets if target.is_active) + def update_targets(self) -> tuple[RolloutWeightUpdateTarget, ...]: + return tuple(target for target in self.weight_update_targets) + + @property + def update_target_infos(self) -> list[dict[str, Any]]: + return [ + { + "endpoint_rank": target.endpoint_rank, + "server_url": target.server_url, + "lifecycle_state": target.lifecycle_state, + "update_ranks": target.update_ranks, + "engine_size": target.engine_size, + } + for target in self.update_targets + ] @property def nccl_engine_infos(self) -> tuple[tuple[int, str, int], ...]: - return tuple( - (target.endpoint_rank, target.server_url, target.engine_size) for target in self.active_update_targets - ) + return tuple((target.endpoint_rank, target.server_url, target.engine_size) for target in self.update_targets) @property def transport_signature(self) -> tuple[Any, ...]: diff --git a/xtuner/v1/rl/weight_update/transport.py b/xtuner/v1/rl/weight_update/transport.py index a8fd68cc27..d14ba0a96f 100644 --- a/xtuner/v1/rl/weight_update/transport.py +++ b/xtuner/v1/rl/weight_update/transport.py @@ -73,6 +73,10 @@ def __init__(self, *, rollout_info: RolloutWeightUpdateInfo, logger: Any, rank: self.rollout_url = self.rollout_info.rollout_url + def reset_rollout_info(self, rollout_info: RolloutWeightUpdateInfo): + self.rollout_info = rollout_info + self.rollout_url = rollout_info.rollout_url + @staticmethod def post_json(url: str, endpoint: str, payload: dict, *, api_key=None) -> dict: headers = {"Content-Type": "application/json"} @@ -477,8 +481,6 @@ def after_update_per_group(self) -> None: def send(self, batch: WeightUpdateBatch) -> None: ipc_update_target = self.rollout_info._ipc_update_target assert ipc_update_target is not None, "IPC rollout target for current train rank is not resolved." - if not ipc_update_target.is_active: - return rollout_url = ipc_update_target.server_url DEVICE_MODULE.empty_cache() @@ -899,7 +901,9 @@ def __init__( self._checkpoint_name: str | None = None self._ps = self.build_parameter_server() + self._p2p_available = self._check_checkpoint_engine_p2p_available() # 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): @@ -920,6 +924,20 @@ def build_parameter_server(self): self.logger.info(f"[checkpoint_engine] ParameterServer ready rank={self.rank} world_size={self.ps_world_size}") return ps + def _check_checkpoint_engine_p2p_available(self) -> bool: + try: + from mooncake.engine import TransferEngine # noqa: F401 + except ImportError as e: + self.logger.warning( + "Checkpoint Engine P2P weight update is unavailable because " + "mooncake TransferEngine is not installed or cannot be imported. " + "Full Checkpoint Engine broadcast weight update may still work, " + "but partial rollout worker recovery requires P2P. " + f"import_error={e!r}" + ) + return False + return True + def split_tensors_for_rank(self, checkpoint_path: str | Path, world_size: int, rank: int) -> set[str]: """Split an HF keys for each ParameterServer.""" @@ -1004,9 +1022,12 @@ def register_checkpoint_from_train_engine(self, weight_iterator: Any) -> None: 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() + try: + self._ps.register_checkpoint(name, files=[], named_tensors=shard, use_shared_memory_pool=True) + if self._sync_after_register: + DEVICE_MODULE.synchronize() + except Exception: + self.logger.error(f"[checkpoint_engine] register_checkpoint failed rank={self.rank} name={name}") self._checkpoint_name = name def _make_req_func(self, targets: Sequence[RolloutWeightUpdateTarget]): @@ -1079,17 +1100,31 @@ def _update_engines(self) -> None: """``gather_metas`` then ``update`` to push checkpoint to rollout engines.""" - targets = self.rollout_info.active_update_targets + targets = self.rollout_info.update_targets + + self.logger.info( + f"[checkpoint_engine] update rollout engine info rank={self.rank} selected rollout workers for weight update: {self.rollout_info.update_target_infos}" + ) + 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 + + if not use_broadcast and not self._p2p_available: + self.logger.warning( + "Checkpoint Engine partial weight update requires P2P, but mooncake " + "TransferEngine is unavailable. update_ranks=%s world_size=%s. " + "Install mooncake TransferEngine or fall back to full broadcast update.", + update_ranks, + self.ps_world_size, + ) 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}" + f"[checkpoint_engine] gather_metas+update name={self._checkpoint_name} ranks={self.rank} " + f"selected_targets={len(targets)}/{self.ps_world_size} " + f"method={'broadcast' if use_broadcast else 'p2p'} update_ranks={update_ranks} " ) self._ps.gather_metas(self._checkpoint_name) self._ps.update(self._checkpoint_name, req_func, ranks=ranks) @@ -1116,6 +1151,7 @@ def update(self, weight_iterator: Any, **kwargs: Any) -> None: 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." ) @@ -1130,6 +1166,9 @@ def update(self, weight_iterator: Any, **kwargs: Any) -> None: if need_update: self._update_engines() + def has_registered_checkpoint(self) -> bool: + return self._checkpoint_name is not None + def reset_rollout_info(self, rollout_info: RolloutWeightUpdateInfo): self.rollout_info = rollout_info self.rollout_url = rollout_info.rollout_url diff --git a/xtuner/v1/rl/weight_update/update_weighter.py b/xtuner/v1/rl/weight_update/update_weighter.py index 0d619e2669..7223ee8c44 100644 --- a/xtuner/v1/rl/weight_update/update_weighter.py +++ b/xtuner/v1/rl/weight_update/update_weighter.py @@ -55,7 +55,6 @@ def bind_rollout_weight_update( 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, engine=self._engine, @@ -76,6 +75,15 @@ def weight_update(self, **kwargs: Any) -> None: assert self.weight_iterator is not None, "Weight iterator is not initialized." self._transport.update(self.weight_iterator, **kwargs) + def has_registered_weight_checkpoint(self) -> bool: + transport = self._transport + if transport is None: + return False + has_registered = getattr(transport, "has_registered_checkpoint", None) + if has_registered is None: + return False + return bool(has_registered()) + def _set_transport(self) -> None: rollout_info = self.rollout_info assert rollout_info is not None, "bind_rollout_weight_update() must be called before setting transport." diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 2aeba247ce..49dee2b650 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -33,6 +33,7 @@ ) from xtuner.v1.rl.agent_loop_manager.produce_utils import default_should_continue_fn from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.health_manager import RLHealthManager, _NoOpRLHealthManager from xtuner.v1.rl.replay_buffer import ( AsyncReplayBufferConfig, SyncReplayBufferConfig, @@ -41,6 +42,7 @@ ) from xtuner.v1.rl.rollout.controller import RolloutControllerProxy from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.rollout.worker_registry import WorkerLifecycleState from xtuner.v1.rl.trace import TraceConfig, close_trace, configure_trace from xtuner.v1.rl.trainer.controller import TrainingController from xtuner.v1.rl.trainer.worker import WorkerConfig, WorkerLogItem @@ -188,8 +190,17 @@ def bind_train_rollout( rollout_config: RolloutConfig, ) -> None: """Bind the training and rollout workers for update weights.""" + # Promote pending workers to active before the regular weight update so + # subsequent lifecycle operations can handle them together. + ray.get( + rollout_controller.mark_worker_groups_lifecycle_state.remote( + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + target_state=WorkerLifecycleState.ACTIVE, + ), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) targets = ray.get( - rollout_controller.get_weight_update_targets.remote(), # type: ignore[attr-defined] + rollout_controller.get_weight_update_targets.remote(target_state=WorkerLifecycleState.ACTIVE), timeout=RL_TRAINER_RAY_GET_TIMEOUT, ) train_controller.bind_rollout_weight_update( @@ -1011,10 +1022,6 @@ def _train_one_batch( # 共卡训练前切换资源:检查 rollout -> offload rollout -> onload train。 if offload_rollout_before_train: - ray.get( - self.rollout_controller.check_and_shutdown_inactive_workers.remote(), - timeout=RL_TRAINER_RAY_GET_TIMEOUT, - ) ray.get(self.rollout_controller.offload.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) if onload_train_before_train: if getattr(self, "_train_nccl_suspended", False): @@ -1668,6 +1675,7 @@ def _log_mini_batch_metrics(self, workers_log_item: List[WorkerLogItem]): class RLColocateTrainer(BaseRLTrainer): _META_PATH = ".xtuner_rl_colocate_trainer" agent_loop_manager: AgentLoopManager + rl_health_manager: RLHealthManager | _NoOpRLHealthManager # 共卡保留资源切换和权重同步流程;通用保存、日志在 BaseRLTrainer。 def __init__(self, cfg: RLColocateTrainerConfig): @@ -1681,6 +1689,7 @@ def __init__(self, cfg: RLColocateTrainerConfig): set_cpu_resource_manager(self._cpu_resource_manager) if self._debug_rollout: + self.rl_health_manager = _NoOpRLHealthManager() if self._rollout_config.skip_load_weights: self.logger.info( "debug_rollout cannot be used with rollout_config.skip_load_weights=True. force set skip_load_weights to False" @@ -1702,6 +1711,7 @@ def __init__(self, cfg: RLColocateTrainerConfig): checkpoint_path = self._resume_train_controller_and_state(checkpoint_path) if self._debug_train: + self.rl_health_manager = _NoOpRLHealthManager() assert self._debug_rollout_dir is not None self.tokenizer = AutoTokenizer.from_pretrained(cfg.tokenizer_path, trust_remote_code=True) self._debug_train_files = self._list_debug_rollout_files(self._debug_rollout_dir) @@ -1719,6 +1729,12 @@ def __init__(self, cfg: RLColocateTrainerConfig): if self._rollout_config.weight_transport_type is None: self._rollout_config.weight_transport_type = "ipc" + self.rl_health_manager = RLHealthManager( + train_controller=self.train_controller, + rollout_controller=self.rollout_controller, + rollout_config=self._rollout_config, + ) + bind_train_rollout( train_controller=self.train_controller, rollout_controller=self.rollout_controller, @@ -1757,9 +1773,11 @@ def _sync_weights_from_train_workers(self) -> None: self.logger.info("Rollout workers updated weights from train workers.") def fit(self): + self.rl_health_manager.start() try: self._fit() finally: + self.rl_health_manager.stop() self._exp_tracker.close() close_trace() @@ -1792,17 +1810,22 @@ def _fit(self): step_timer_dict = {} with timer("step", step_timer_dict): # 共卡一次调用内完成生产和消费。 + self.rl_health_manager.set_rollout_resources_available(True) self.logger.info( f"[Step {train_step}] start to generate rollout experience for train step {train_step} with model step {model_step}" ) - with timer("produce_batch", step_timer_dict): - produce_result: ProduceBatchResult = asyncio_run( - self.agent_loop_manager.produce_batch( - self.train_batch_size, - train_step=train_step, - model_step=model_step, + try: + with timer("produce_batch", step_timer_dict): + produce_result: ProduceBatchResult = asyncio_run( + self.agent_loop_manager.produce_batch( + self.train_batch_size, + train_step=train_step, + model_step=model_step, + ) ) - ) + finally: + self.rl_health_manager.set_rollout_resources_available(False) + if XTUNER_DETERMINISTIC: produce_result.rollout_states = sort_rollout_state_for_deterministic(produce_result.rollout_states) train_batch = produce_result.rollout_states @@ -1887,33 +1910,29 @@ def _sync_weights_and_save(self, train_step: int, step_timer_dict: dict) -> bool timer_name = "sync_weight" if should_sync_weights else "switch_to_rollout" with timer(timer_name, step_timer_dict): if should_sync_weights: - ray.get( - self.rollout_controller.restart_inactive_workers.remote(), - timeout=RL_TRAINER_RAY_GET_TIMEOUT, - ) - bind_train_rollout( - train_controller=self.train_controller, - rollout_controller=self.rollout_controller, - rollout_config=self._rollout_config, - ) - - 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, + with self.rl_health_manager.weight_update_guard(): + bind_train_rollout( + train_controller=self.train_controller, + rollout_controller=self.rollout_controller, + rollout_config=self._rollout_config, ) - self.train_controller.weight_update(need_register=False, need_update=True) + 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") - 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") suspend_train_nccl = ( os.getenv( "XTUNER_SUSPEND_TRAIN_NCCL_AFTER_SYNC", diff --git a/xtuner/v1/utils/httpx_utils.py b/xtuner/v1/utils/httpx_utils.py index c2be88632a..ed2e0b1027 100644 --- a/xtuner/v1/utils/httpx_utils.py +++ b/xtuner/v1/utils/httpx_utils.py @@ -107,7 +107,9 @@ def __post_init__(self): ) if self.error_type == HttpRequestErrorType.REQUEST_ERROR and self.exception: if hasattr(self.exception, "__cause__") and self.exception.__cause__: - self.error_msg += f" __cause__: {self.exception.__cause__}" + # Truncate the chained exception message to prevent oversized error propagation. + cause = str(self.exception.__cause__) + self.error_msg += f" __cause__: {cause[:2048]}" @property def is_success(self) -> bool: