From 9eb2dccbb7c0abbd1567fc27a6cea915e5540001 Mon Sep 17 00:00:00 2001 From: Xun Sun Date: Sat, 25 Jul 2026 14:32:47 +0800 Subject: [PATCH] [Elastic EP] Fix recovery lifecycle and add manual coverage (#31744) Co-authored-by: Shangming Cai --- python/sglang/srt/elastic_ep/elastic_ep.py | 16 +- python/sglang/srt/managers/scheduler.py | 6 + python/sglang/srt/managers/tp_worker.py | 4 +- .../sglang/srt/model_executor/model_runner.py | 66 ++--- .../load_model_utils.py | 4 +- python/sglang/srt/server_args.py | 4 +- test/manual/ep/test_elastic_recover.py | 255 ++++++++++++++++++ 7 files changed, 310 insertions(+), 45 deletions(-) create mode 100644 test/manual/ep/test_elastic_recover.py diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 8ceb53a9e..40e277af2 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -11,9 +11,10 @@ from sglang.srt.distributed import get_world_group, parallel_state from sglang.srt.distributed.utils import get_global_tcp_store from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata from sglang.srt.managers.schedule_batch import ServerArgs -from sglang.srt.utils import broadcast_pyobj, is_cpu, is_cuda +from sglang.srt.utils import is_cpu, is_cuda if TYPE_CHECKING: + from sglang.srt.configs.model_config import ModelConfig from sglang.srt.eplb.eplb_manager import EPLBManager logger = logging.getLogger(__name__) @@ -462,7 +463,8 @@ def maybe_recover_ep_ranks( *, tp_group: parallel_state.GroupCoordinator, eplb_manager: EPLBManager, - random_seed: int, + model_config: ModelConfig, + moe_ep_rank: int, ) -> bool: # TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device # synchronization, and this function is on the forward-path. @@ -489,17 +491,13 @@ def maybe_recover_ep_ranks( if ranks_to_recover and try_recover_ranks(ranks_to_recover): eplb_manager.reset_generator() broadcast_global_expert_location_metadata( + model_config=model_config, + moe_ep_rank=moe_ep_rank, src_rank=get_healthy_expert_location_src_rank( invoked_in_elastic_ep_rejoin_path=False - ) + ), ) ElasticEPStateManager.instance().reset() - broadcast_pyobj( - [random_seed], - parallel_state.get_world_group().rank, - parallel_state.get_world_group().cpu_group, - src=parallel_state.get_world_group().ranks[0], - ) logger.info(f"recover ranks {ranks_to_recover} done") return True diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 1fb37645c..3fc35ee2e 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -885,6 +885,12 @@ class Scheduler( if model_runner.token_to_kv_pool.post_capture_active: model_runner.post_capture_resize_kv_pool() + if ( + self.server_args.elastic_ep_backend is not None + and self.server_args.ep_join_mode == "recover" + ): + model_runner.post_capture_elastic_ep_recover() + # Dispatch the model worker if self.spec_algorithm.is_none(): self.model_worker = self.tp_worker diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 6d09e8133..eec4e6455 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -342,8 +342,8 @@ class TpModelWorker(BaseTpWorker): self.world_group = get_world_group() # Sync random seed across TP workers. - # Scale joiners cannot enter the launch-time WORLD broadcast. - if server_args.is_ep_scale_joiner: + # Elastic joiners cannot enter the launch-time WORLD broadcast. + if server_args.is_ep_joiner: self.random_seed = server_args.random_seed else: self.random_seed = broadcast_pyobj( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 2d814f085..e6f0460e5 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -399,39 +399,27 @@ class ModelRunner: def _initialize_elastic_ep_joiner(self) -> None: if not ( self.server_args.elastic_ep_backend is not None - and self.server_args.is_ep_joiner + and self.server_args.is_ep_scale_joiner ): return - is_scale_join = self.server_args.ep_join_mode == "scale" - if is_scale_join: - join_effective_ep_size = ( - self.server_args.ep_join_rank_offset + self.ps.tp_size + join_effective_ep_size = self.server_args.ep_join_rank_offset + self.ps.tp_size + dist.barrier(group=self.tp_group.cpu_group) + if self.ps.tp_rank == 0: + register_scale_cohort( + self.server_args.ep_join_rank_offset, + join_effective_ep_size, ) - dist.barrier(group=self.tp_group.cpu_group) - if self.ps.tp_rank == 0: - register_scale_cohort( - self.server_args.ep_join_rank_offset, - join_effective_ep_size, - ) - join_scale_process_group() - self.server_args.override( - "elastic_ep.scale_join", ep_size=join_effective_ep_size - ) - else: - join_process_groups() + join_scale_process_group() + self.server_args.override( + "elastic_ep.scale_join", ep_size=join_effective_ep_size + ) global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset broadcast_global_expert_location_metadata( model_config=self.model_config, moe_ep_rank=global_ep_rank, - src_rank=( - 0 - if is_scale_join - else get_healthy_expert_location_src_rank( - invoked_in_elastic_ep_rejoin_path=True - ) - ), + src_rank=0, ) set_global_expert_distribution_recorder( ExpertDistributionRecorder.init_new( @@ -441,10 +429,6 @@ class ModelRunner: ) ) - if not is_scale_join: - ElasticEPStateManager.instance().reset() - return - from sglang.srt.layers.dp_attention import ( enable_joiner_all_gather, update_dp_attention_post_scale, @@ -839,6 +823,27 @@ class ModelRunner: resize.capped_max_running_requests ) + def post_capture_elastic_ep_recover(self): + join_process_groups() + + global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset + broadcast_global_expert_location_metadata( + model_config=self.model_config, + moe_ep_rank=global_ep_rank, + src_rank=get_healthy_expert_location_src_rank( + invoked_in_elastic_ep_rejoin_path=True + ), + ) + set_global_expert_distribution_recorder( + ExpertDistributionRecorder.init_new( + self.server_args, + get_global_expert_location_metadata(), + rank=global_ep_rank, + ) + ) + + ElasticEPStateManager.instance().reset() + def init_attention_backends(self): """Initialize attention backends only (no cuda graph capture).""" # Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture @@ -1089,7 +1094,7 @@ class ModelRunner: dist_barrier_after_load( elastic_ep_backend=self.server_args.elastic_ep_backend, tp_rank=self.ps.tp_rank, - is_ep_scale_joiner=self.server_args.is_ep_scale_joiner, + is_ep_joiner=self.server_args.is_ep_joiner, ) def maybe_init_dwdp(self): @@ -1822,7 +1827,8 @@ class ModelRunner: recovered = maybe_recover_ep_ranks( tp_group=self.tp_group, eplb_manager=self.eplb_manager, - random_seed=self.server_args.random_seed, + model_config=self.model_config, + moe_ep_rank=self._elastic_global_rank(), ) if recovered: self.forward_pass_id = 0 diff --git a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py index 855994aaf..f91339d5e 100644 --- a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py @@ -301,11 +301,11 @@ def dist_barrier_after_load( *, elastic_ep_backend: Optional[str], tp_rank: int, - is_ep_scale_joiner: bool = False, + is_ep_joiner: bool = False, ) -> None: if elastic_ep_backend == "mooncake": # Mooncake does not support `monitored_barrier` - if not is_ep_scale_joiner: + if not is_ep_joiner: dist.barrier(group=get_tp_group().cpu_group) else: # Handle the case where some ranks do not finish loading. diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 8db35a082..7866b590a 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -8984,7 +8984,7 @@ class PortArgs: # (no availability-based search). If incrementing would # overflow the valid TCP range, decrement instead. NUM_DERIVED_PORTS = 5 - if server_args.is_ep_scale_joiner: + if server_args.is_ep_joiner: port_base = server_args.port + ZMQ_TCP_PORT_DELTA if port_base + NUM_DERIVED_PORTS > 65535: port_base = server_args.port - ZMQ_TCP_PORT_DELTA @@ -9004,7 +9004,7 @@ class PortArgs: assert worker_ports is not None scheduler_input_port = worker_ports[dp_rank] - is_joiner = server_args.is_ep_scale_joiner + is_joiner = server_args.is_ep_joiner # Under SGLANG_DISTRIBUTED_INIT_METHOD_OVERRIDE, SGLang never binds # dist_init_port / nccl_port (rendezvous uses the externally-managed # store; see distributed/bootstrap.py:_resolve_dist_init_method), so diff --git a/test/manual/ep/test_elastic_recover.py b/test/manual/ep/test_elastic_recover.py new file mode 100644 index 000000000..beab45b92 --- /dev/null +++ b/test/manual/ep/test_elastic_recover.py @@ -0,0 +1,255 @@ +"""Manual single-host Elastic EP recovery test. + +Run: + + CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7 python -m pytest \ + test/manual/ep/test_elastic_recover.py -v -s +""" + +import os +import shlex +import subprocess +import time +import unittest +from pathlib import Path + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.server_fixtures.disaggregation_fixture import get_rdma_devices_args +from sglang.test.test_utils import ( + DEFAULT_MODEL_NAME_FOR_TEST_MLA, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + CustomTestCase, + try_cached_model, +) +from sglang.utils import wait_for_http_ready + +TEST_MODEL = os.environ.get( + "SGLANG_ELASTIC_RECOVER_TEST_MODEL", + try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST_MLA), +) +EP_SIZE = 8 +LOCAL_EP_SIZE = 4 +DIST_INIT_ADDR = os.environ.get("SGLANG_ELASTIC_RECOVER_DIST_INIT", "127.0.0.1:25555") +PRIMARY_PORT = int(os.environ.get("SGLANG_ELASTIC_RECOVER_PRIMARY_PORT", "21000")) +JOINER_PORT = int(os.environ.get("SGLANG_ELASTIC_RECOVER_JOINER_PORT", "22000")) +RECOVER_WAIT_SECONDS = float(os.environ.get("SGLANG_ELASTIC_RECOVER_WAIT_SECONDS", "5")) +RECOVER_TIMEOUT_SECONDS = float( + os.environ.get("SGLANG_ELASTIC_RECOVER_TIMEOUT_SECONDS", "300") +) +RANDOM_SEED = int(os.environ.get("SGLANG_ELASTIC_RECOVER_RANDOM_SEED", "42")) +ib_devices = get_rdma_devices_args() + + +def _visible_device_ids() -> list[str]: + visible = os.environ.get("CUDA_VISIBLE_DEVICES") + if visible: + return [device.strip() for device in visible.split(",") if device.strip()] + try: + import torch + + return [str(index) for index in range(torch.cuda.device_count())] + except Exception: + return [] + + +def _server_args(node_rank: int, port: int, recover: bool = False) -> list[str]: + args = [ + "sglang", + "serve", + "--model-path", + TEST_MODEL, + "--host", + "127.0.0.1", + "--port", + str(port), + "--device", + "cuda", + "--trust-remote-code", + "--tp", + str(EP_SIZE), + "--dp", + str(EP_SIZE), + "--nnodes", + "2", + "--node-rank", + str(node_rank), + "--dist-init-addr", + DIST_INIT_ADDR, + "--random-seed", + str(RANDOM_SEED), + "--enable-dp-attention", + "--enable-dp-lm-head", + "--elastic-ep-backend", + "mooncake", + "--mooncake-ib-device", + ib_devices, + "--moe-a2a-backend", + "mooncake", + "--deepep-mode", + "low_latency", + "--moe-dense-tp-size", + "1", + "--disable-custom-all-reduce", + "--enable-eplb", + "--ep-num-redundant-experts", + "72", + "--chunked-prefill-size", + "512", + "--cuda-graph-max-bs-decode", + "16", + "--mem-fraction-static", + "0.5", + ] + if recover: + args.extend(["--elastic-ep-join-mode", "recover"]) + extra_args = os.environ.get("SGLANG_ELASTIC_RECOVER_EXTRA_SERVER_ARGS", "") + return args + shlex.split(extra_args) + + +@unittest.skipUnless( + len(_visible_device_ids()) >= EP_SIZE, + "Elastic EP recovery E2E needs 8 visible GPUs.", +) +class TestElasticRecover4To4(CustomTestCase): + """Kill one four-rank node and recover it with a fresh process group.""" + + @classmethod + def setUpClass(cls): + cls.base_url = f"http://127.0.0.1:{PRIMARY_PORT}" + cls.processes: list[subprocess.Popen] = [] + cls.log_files = [] + cls.log_paths: dict[str, Path] = {} + visible_devices = _visible_device_ids() + + cls.primary = cls._launch( + node_rank=0, + port=PRIMARY_PORT, + visible_devices=visible_devices[:LOCAL_EP_SIZE], + name="primary", + ) + cls.initial_joiner = cls._launch( + node_rank=1, + port=JOINER_PORT, + visible_devices=visible_devices[LOCAL_EP_SIZE:EP_SIZE], + name="initial_joiner", + ) + wait_for_http_ready( + f"{cls.base_url}/health_generate", + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + process=cls.primary, + ) + + @classmethod + def _launch( + cls, + *, + node_rank: int, + port: int, + visible_devices: list[str], + name: str, + recover: bool = False, + ) -> subprocess.Popen: + log_path = Path(f"/tmp/elastic_ep_recover_{name}_{int(time.time())}.log") + log_file = open(log_path, "w") + env = os.environ.copy() + env["CUDA_VISIBLE_DEVICES"] = ",".join(visible_devices) + process = subprocess.Popen( + _server_args(node_rank, port, recover), + env=env, + stdout=log_file, + stderr=subprocess.STDOUT, + ) + cls.processes.append(process) + cls.log_files.append(log_file) + cls.log_paths[name] = log_path + print(f"Started {name}; log: {log_path}") + return process + + @classmethod + def tearDownClass(cls): + for process in reversed(getattr(cls, "processes", [])): + if process.poll() is None: + kill_process_tree(process.pid, wait_timeout=60) + for log_file in getattr(cls, "log_files", []): + log_file.close() + + def _generate(self, routed_dp_rank: int | None = None) -> requests.Response: + payload = { + "text": "The capital of France is", + "sampling_params": {"max_new_tokens": 4, "temperature": 0.0}, + } + if routed_dp_rank is not None: + payload["routed_dp_rank"] = routed_dp_rank + return requests.post(f"{self.base_url}/generate", json=payload, timeout=90) + + def _generate_ok(self, description: str, routed_dp_rank: int | None = None) -> None: + response = self._generate(routed_dp_rank) + self.assertEqual(response.status_code, 200, f"{description}: {response.text}") + payload = response.json() + generated_text = payload.get("text", "") + self.assertIn( + "paris", + generated_text.casefold(), + f"{description}: unexpected generation: {generated_text!r}", + ) + + def _wait_for_recover_capture(self) -> None: + deadline = time.monotonic() + RECOVER_TIMEOUT_SECONDS + log_path = self.log_paths["recover_joiner"] + marker = "Capture target decode CUDA graph end" + while time.monotonic() < deadline: + self.assertIsNone( + self.recover_joiner.poll(), + "Recover joiner exited during CUDA graph capture", + ) + if log_path.exists() and log_path.read_text(errors="replace").count( + marker + ) >= (EP_SIZE - LOCAL_EP_SIZE): + return + time.sleep(2) + self.fail(f"Timed out waiting for recover CUDA graph capture: {log_path}") + + def _wait_for_recovered_ranks(self) -> None: + self._wait_for_recover_capture() + self._generate_ok("recovery trigger") + deadline = time.monotonic() + RECOVER_TIMEOUT_SECONDS + marker = f"recover ranks {list(range(LOCAL_EP_SIZE, EP_SIZE))} done" + primary_log = self.log_paths["primary"] + while time.monotonic() < deadline: + self.assertIsNone( + self.recover_joiner.poll(), "Recover joiner exited before rejoining" + ) + if ( + primary_log.exists() + and primary_log.read_text(errors="replace").count(marker) + >= LOCAL_EP_SIZE + ): + for request_index in range(3): + self._generate_ok(f"post-recovery request {request_index + 1}") + return + time.sleep(2) + self.fail(f"Timed out waiting for recovery collective: {primary_log}") + + def test_recover_four_ranks(self): + self._generate_ok("initial service") + + kill_process_tree(self.initial_joiner.pid, wait_timeout=60) + # Give the terminated schedulers time to disappear before fault handling. + time.sleep(RECOVER_WAIT_SECONDS) + self._generate_ok("degraded service after node1 failure") + + visible_devices = _visible_device_ids() + self.recover_joiner = self._launch( + node_rank=1, + port=JOINER_PORT, + visible_devices=visible_devices[LOCAL_EP_SIZE:EP_SIZE], + name="recover_joiner", + recover=True, + ) + self._wait_for_recovered_ranks() + + +if __name__ == "__main__": + unittest.main()