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..1727996c3a --- /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.RECOVERING.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..010213907a 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -19,6 +19,7 @@ """ import asyncio +import threading import tempfile import unittest from pathlib import Path @@ -112,6 +113,11 @@ 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_resources_available = threading.Event() + trainer._rollout_weight_update_lock = threading.Lock() + trainer._pending_rollout_weight_update_stop_event = threading.Event() + trainer._pending_rollout_weight_update_thread: threading.Thread | None = None + trainer._rollout_config = SimpleNamespace(weight_transport_type='ipc') trainer.logger = MagicMock() trainer._total_train_steps = total_train_steps trainer._cur_step = 0 diff --git a/tests/rl/test_rollout_logic.py b/tests/rl/test_rollout_logic.py index e5a1a879b3..65058f84ee 100644 --- a/tests/rl/test_rollout_logic.py +++ b/tests/rl/test_rollout_logic.py @@ -651,6 +651,29 @@ def test_registry_filters_entrypoints_and_tracks_lifecycle(self): registry.set_group_recovery_result(claimed_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)) registry = RolloutWorkerRegistry(rollout_topology=runtime_layout) @@ -668,7 +691,6 @@ def test_registry_projects_weight_update_targets_from_topology_and_runtime_state 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): @@ -1228,7 +1250,10 @@ def test_run_once_does_not_log_error_when_last_active_worker_becomes_inactive(se 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, [()]) @@ -1325,7 +1350,7 @@ 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): + def test_restart_barrier_marks_recovered_group_pending_weights_after_success(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) worker_info = WorkerSnapshot( rank=0, @@ -1334,10 +1359,10 @@ def test_restart_barrier_notifies_recovered_group_after_success(self): session_url="http://session-0", lifecycle_state=WorkerLifecycleState.INACTIVE, ) - recovered_groups = [] + pending_weights_groups = [] listener = SimpleNamespace( on_worker_group_inactive=MagicMock(), - on_worker_group_recovered=recovered_groups.append, + on_worker_group_pending_weights=pending_weights_groups.append, ) manager, registry = self._build_manager( {0: worker_info}, @@ -1347,9 +1372,14 @@ def test_restart_barrier_notifies_recovered_group_after_success(self): 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)) + self.assertEqual(self._worker_by_rank(registry, 0).lifecycle_state, WorkerLifecycleState.PENDING_WEIGHTS) + self.assertEqual([group.ranks for group in pending_weights_groups], [(0,)]) + self.assertTrue( + all( + worker.lifecycle_state is WorkerLifecycleState.PENDING_WEIGHTS + for worker in pending_weights_groups[0].workers + ) + ) def test_restart_barrier_cleans_claimed_groups_when_stopping(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) @@ -1442,7 +1472,7 @@ 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): + def test_pending_weights_listener_runs_outside_lifecycle_operation_lock(self): actor = SimpleNamespace(check_health=_FakeAsyncRemoteMethod(True)) worker_info = WorkerSnapshot( rank=0, @@ -1453,7 +1483,7 @@ def test_recovered_listener_runs_outside_lifecycle_operation_lock(self): lock_acquired_by_listener = [] manager, _ = self._build_manager({0: worker_info}) - def on_worker_group_recovered(group): + def on_worker_group_pending_weights(group): acquired = manager._lifecycle_operation_lock.acquire(blocking=False) lock_acquired_by_listener.append(acquired) if acquired: @@ -1462,7 +1492,7 @@ def on_worker_group_recovered(group): manager._worker_lifecycle_listeners = ( SimpleNamespace( on_worker_group_inactive=MagicMock(), - on_worker_group_recovered=on_worker_group_recovered, + on_worker_group_pending_weights=on_worker_group_pending_weights, ), ) 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 88c7dde548..95584f3cee 100644 --- a/xtuner/v1/rl/agent_loop_manager/producer.py +++ b/xtuner/v1/rl/agent_loop_manager/producer.py @@ -216,6 +216,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:** @@ -233,6 +236,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, @@ -261,6 +265,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, is_valid_sample_fn=self.is_valid_sample_fn, should_continue_fn=self.should_continue_fn, ) @@ -338,6 +343,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, @@ -368,6 +374,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: @@ -434,7 +441,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/rollout/controller.py b/xtuner/v1/rl/rollout/controller.py index b2ab50e23c..859751940c 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 @@ -70,6 +70,28 @@ 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_pending_weight_update_targets(self) -> tuple[RolloutWeightUpdateTarget, ...]: + """Return recovered rollout endpoints waiting for weights.""" + return tuple( + target + for target in self.registry.weight_update_targets() + if target.lifecycle_state == WorkerLifecycleState.PENDING_WEIGHTS + ) + + 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: return @@ -120,6 +142,11 @@ async def generate(self, rollout_state: RolloutState) -> RolloutState: f"Rollout request timed out after {self.config.rollout_timeout * self.timeout_multiplier} seconds." ) return rollout_state + except Exception as e: + self.logger.exception(f"RolloutController.generate failed: session_id={session_id}") + rollout_state.status = Status.FAILED + rollout_state.error_msg = f"Rollout request failed: {type(e).__name__}: {str(e)[:1024]}" + return rollout_state def set_enable_partial_rollout(self, enable: bool) -> None: """Propagate enable_partial_rollout flag to all active workers.""" @@ -159,24 +186,48 @@ async def check_and_shutdown_inactive_workers(self): 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, ...]], + source_state: WorkerLifecycleState, + target_state: WorkerLifecycleState, + ) -> None: + """Move selected worker groups from source_state to target_state. + + Only groups whose current lifecycle state matches source_state are considered. When groups are moved to ACTIVE + or INACTIVE, the health manager is notified so routing and lifecycle listeners stay in sync. + """ + groups_by_ranks = {group.ranks: group for group in self.registry.get_target_state_worker_groups(source_state)} + 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_recovered(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 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.""" @@ -187,8 +238,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..a10ec29329 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,6 +37,8 @@ class RolloutWorkerLifecycleListener(Protocol): def on_worker_group_inactive(self, group: WorkerGroup) -> None: ... + def on_worker_group_pending_weights(self, group: WorkerGroup) -> None: ... + def on_worker_group_recovered(self, group: WorkerGroup) -> None: ... @@ -210,33 +212,14 @@ 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. + self._restart_inactive_workers() - 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, ...] = () - - 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() def check_and_shutdown_inactive_workers(self) -> None: """Fail-fast health-check active workers, mark failures inactive, and @@ -394,15 +377,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 +418,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 +490,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 +499,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 +615,50 @@ def _wait_worker_server_down(self, worker: WorkerSnapshot, *, max_wait_attempts: return False + def notify_worker_group_recovered(self, groups: Iterable[WorkerGroup]) -> None: + self._notify_worker_lifecycle_listeners( + groups, + event_name="recovered", + notify_listener=lambda listener, group: listener.on_worker_group_recovered(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.claim_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 () + + self._notify_worker_lifecycle_listeners( + pending_weights_groups, + event_name="pending_weights", + notify_listener=lambda listener, group: listener.on_worker_group_pending_weights(group), + ) + + 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..f0f19a0563 100644 --- a/xtuner/v1/rl/rollout/proxy_manager.py +++ b/xtuner/v1/rl/rollout/proxy_manager.py @@ -57,6 +57,10 @@ 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_pending_weights(self, group: "WorkerGroup") -> None: + """Pending workers are healthy but not ready for routed traffic yet.""" + del group + def on_worker_group_recovered(self, group: "WorkerGroup") -> None: """Register recovered request entrypoints to routed API proxy.""" for worker in group.workers: 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..92ea66281d 100644 --- a/xtuner/v1/rl/rollout/worker_registry.py +++ b/xtuner/v1/rl/rollout/worker_registry.py @@ -32,6 +32,8 @@ class WorkerLifecycleState(str, Enum): 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 +124,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,6 +189,10 @@ def inactive_worker_groups(self) -> tuple[WorkerGroup, ...]: ] return tuple(sorted(inactive_groups, key=lambda group: group.ranks)) + def pending_weights_worker_groups(self) -> tuple[WorkerGroup, ...]: + """Return lifecycle groups waiting for a rollout weight update.""" + return self.get_target_state_worker_groups(WorkerLifecycleState.PENDING_WEIGHTS) + def claim_inactive_groups_for_recovery(self) -> tuple[WorkerGroup, ...]: """Claim inactive worker groups by moving them to RECOVERING state.""" with self._lock: @@ -236,6 +258,33 @@ def set_group_recovery_result( ) return recorded_group + 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[RolloutWeightUpdateTarget, ...]: """Return weight-update targets resolved with current runtime state.""" from xtuner.v1.rl.weight_update.data import RolloutWeightUpdateTarget 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 6e8c5ab495..37fbad3ee3 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -331,6 +331,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..ed1f644ddf 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("[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 61ef1d1dbf..6cef93b78b 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -3,6 +3,7 @@ import os import random import re +import threading import time from dataclasses import asdict, dataclass from pathlib import Path @@ -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 @@ -62,6 +64,7 @@ # TODO: Move DEVICE to `xtuner.utils.device` PG_READY_TIMEOUT = 30 RL_TRAINER_RAY_GET_TIMEOUT = 3600 +PENDING_ROLLOUT_WORKER_CHECK_INTERVAL = 1.0 DEVICE = get_device() DEVICE_MODULE = get_torch_device_module() @@ -1607,6 +1610,11 @@ def __init__(self, cfg: RLColocateTrainerConfig): self._cpu_resource_manager.log_initial_snapshot() set_cpu_resource_manager(self._cpu_resource_manager) + 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 + if self._debug_rollout: if self._rollout_config.skip_load_weights: self.logger.info( @@ -1719,9 +1727,11 @@ def _fit(self): step_timer_dict = {} with timer("step", step_timer_dict): # 共卡一次调用内完成生产和消费。 + self._rollout_resources_available.set() self.logger.info( f"[Step {train_step}] start to generate rollout experience for train step {train_step} with model step {model_step}" ) + self._start_check_pending_rollout_worker_thread() with timer("produce_batch", step_timer_dict): produce_result: ProduceBatchResult = asyncio_run( self.agent_loop_manager.produce_batch( @@ -1730,6 +1740,9 @@ def _fit(self): model_step=model_step, ) ) + self._rollout_resources_available.clear() + self._stop_check_pending_rollout_worker_thread() + if XTUNER_DETERMINISTIC: produce_result.rollout_states = sort_rollout_state_for_deterministic(produce_result.rollout_states) train_batch = produce_result.rollout_states @@ -1738,6 +1751,7 @@ def _fit(self): ) if not self._debug_rollout: + self._rollout_resources_available.clear() train_log_info = self._train_one_batch( train_batch, train_step, @@ -1814,33 +1828,60 @@ 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") + with self._rollout_weight_update_lock: + # 重启成功的 inactive worker group 会变成 pending_weights ray.get( - self.rollout_controller.onload_weights.remote(), + self.rollout_controller.restart_inactive_workers.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, + bind_train_rollout( + train_controller=self.train_controller, + rollout_controller=self.rollout_controller, + rollout_config=self._rollout_config, ) - self.train_controller.weight_update() - self.train_controller.offload(target="model") - self.logger.info("Rollout workers update weights successfully in colocate mode") + if self._rollout_config.weight_transport_type == "checkpoint_engine": + pending_targets = ray.get( + self.rollout_controller.get_pending_weight_update_targets.remote(), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + pending_group_ranks = [(target.endpoint_rank,) for target in pending_targets] + self.train_controller.weight_update(need_register=True, need_update=False) + self.train_controller.offload(target="model") + if pending_group_ranks: + ray.get( + self.rollout_controller.onload_weights.remote( + target_state=WorkerLifecycleState.PENDING_WEIGHTS + ), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + 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) + if pending_group_ranks: + # 权重更新成功后将 pending状态变为activate状态 + ray.get( + self.rollout_controller.mark_worker_groups_lifecycle_state.remote( + pending_group_ranks, + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + target_state=WorkerLifecycleState.ACTIVE, + ), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + self.logger.info( + "Pending rollout workers promoted to active after sync-step " + f"Checkpoint Engine weight update: {pending_group_ranks}." + ) + 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", @@ -1857,8 +1898,147 @@ def _sync_weights_and_save(self, train_step: int, step_timer_dict: dict) -> bool self.train_controller.offload(target="model") ray.get(self.rollout_controller.onload_weights.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) ray.get(self.rollout_controller.onload_kvcache.remote(), timeout=RL_TRAINER_RAY_GET_TIMEOUT) + self._rollout_resources_available.set() return should_sync_weights + def _update_pending_rollout_weights_from_checkpoint_engine(self) -> tuple[tuple[int, ...], ...]: + """Synchronize recovered rollout workers with the latest registered + checkpoint in Checkpoint Engine parameter server. + + RolloutHealthManager only restarts failed rollout workers and marks them as + pending_weight_update. They must not receive rollout requests until the + trainer pushes the latest registered weight checkpoint and marks the worker groups active. + + Returns: + The ranks of worker groups that were successfully updated and promoted + to active. + """ + + if self._rollout_config.weight_transport_type != "checkpoint_engine": + return () + + if not self.train_controller.has_registered_weight_checkpoint(): + self.logger.info( + "Skip pending rollout checkpoint-engine update because no train checkpoint has been registered yet." + ) + return () + + pending_targets = ray.get( + self.rollout_controller.get_pending_weight_update_targets.remote(), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + if not pending_targets: + return () + + pending_group_ranks = [(target.endpoint_rank,) for target in pending_targets] + 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_TRAINER_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_TRAINER_RAY_GET_TIMEOUT, + ) + ray.get( + self.rollout_controller.mark_worker_groups_lifecycle_state.remote( + pending_group_ranks, + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + target_state=WorkerLifecycleState.ACTIVE, + ), + timeout=RL_TRAINER_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( + pending_group_ranks, + source_state=WorkerLifecycleState.PENDING_WEIGHTS, + target_state=WorkerLifecycleState.INACTIVE, + ), + timeout=RL_TRAINER_RAY_GET_TIMEOUT, + ) + return () + + def _start_check_pending_rollout_worker_thread(self) -> None: + """Start the background checker for searching pending weight updates + rollout worker.""" + if self._rollout_config.weight_transport_type != "checkpoint_engine": + return + if getattr(self, "_pending_rollout_weight_update_thread", None) 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_check_pending_rollout_worker_thread(self) -> None: + """Stop the background checker for searching pending weight updates + rollout worker.""" + + stop_event = getattr(self, "_pending_rollout_weight_update_stop_event", None) + if stop_event is None: + return + + stop_event.set() + thread = getattr(self, "_pending_rollout_weight_update_thread", None) + if thread is not None: + thread.join(timeout=PENDING_ROLLOUT_WORKER_CHECK_INTERVAL * 2) + if thread.is_alive(): + self.logger.warning("Pending rollout weight update thread did not stop before timeout.") + return + + self._pending_rollout_weight_update_thread = None + self.logger.info("Stopped pending rollout checkpoint-engine update thread.") + + def _pending_rollout_worker_weight_update_loop(self) -> None: + """Periodically update recovered rollout workers that are waiting for + weights. + + The loop waits between checks to avoid busy polling. It skips updates while rollout resources are unavailable, + and uses a non-blocking lock so it does not race with normal rollout weight updates. Any failure is logged and + the next interval will retry pending workers. + """ + while not self._pending_rollout_weight_update_stop_event.wait(PENDING_ROLLOUT_WORKER_CHECK_INTERVAL): + if not self._rollout_resources_available.is_set(): + self.logger.debug( + "Skip pending rollout checkpoint-engine update because rollout resources are unavailable." + ) + continue + if not self._rollout_weight_update_lock.acquire(blocking=False): + self.logger.debug("Skip pending rollout checkpoint-engine update because weight update lock is held.") + continue + try: + 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() + class RLDisaggregatedTrainer(BaseRLTrainer): _META_PATH = ".xtuner_rl_disaggregated_trainer"