From a2d1003b185fd060d7375ada08a3894c9d07f6fa Mon Sep 17 00:00:00 2001 From: Xun Sun Date: Mon, 3 Aug 2026 16:16:41 +0800 Subject: [PATCH] [Mooncake] Fix ProcessGroup API imports (#32403) --- .../sglang/srt/distributed/parallel_state.py | 4 +-- python/sglang/srt/elastic_ep/elastic_ep.py | 32 ++++++++++--------- 2 files changed, 19 insertions(+), 17 deletions(-) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 9347cde21..c90dc0a96 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -311,7 +311,7 @@ class GroupCoordinator: for ranks in group_ranks: subgroup_timeout = _MODEL_PARALLEL_GROUP_TIMEOUT if "mooncake" in torch_distributed_backend: - from mooncake.ep import MooncakeBackendOptions + from mooncake.pg import MooncakeBackendOptions pg_active_size = len(ranks) if not recovered_rank and max_world_size is not None: @@ -2114,7 +2114,7 @@ def init_distributed_environment( _MODEL_PARALLEL_GROUP_TIMEOUT = timeout if backend == "mooncake": - from mooncake.ep import MooncakeBackendOptions + from mooncake.pg import MooncakeBackendOptions use_max_ws = max_world_size and max_world_size > world_size ar_size = max_world_size if use_max_ws else world_size diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 61037ded8..d261101cb 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -351,8 +351,10 @@ def _map_global_to_group_local_ranks( return [rank_to_local[rank] for rank in global_ranks if rank in rank_to_local] -def _wait_for_peer_state(mooncake_ep, backend, ranks: List[int]) -> None: - while not all(mooncake_ep.get_peer_state(backend, ranks)): +def _wait_for_peer_state(backend, ranks: List[int]) -> None: + from mooncake.pg import get_peer_state + + while not all(get_peer_state(backend, ranks)): time.sleep(_PEER_STATE_POLL_INTERVAL_SEC) @@ -368,13 +370,13 @@ def _maybe_create_message_queue(group) -> None: def _try_recover_world(global_ranks: List[int]) -> bool: - from mooncake import ep as mooncake_ep + from mooncake.pg import get_peer_state, recover_ranks world_backend = torch.distributed.group.WORLD - if not all(mooncake_ep.get_peer_state(world_backend, global_ranks)): + if not all(get_peer_state(world_backend, global_ranks)): return False - mooncake_ep.recover_ranks(world_backend, global_ranks) + recover_ranks(world_backend, global_ranks) logger.debug("[Elastic EP][recover] WORLD recover_ranks(%s) done", global_ranks) return True @@ -393,17 +395,17 @@ def try_recover_ranks(global_ranks: List[int]) -> bool: if not _try_recover_world(global_ranks): return False - from mooncake import ep as mooncake_ep + from mooncake.pg import recover_ranks for group in _iter_live_parallel_groups(): local_ranks = _map_global_to_group_local_ranks(group.ranks, global_ranks) if not local_ranks: continue - _wait_for_peer_state(mooncake_ep, group.device_group, local_ranks) - mooncake_ep.recover_ranks(group.device_group, local_ranks) - _wait_for_peer_state(mooncake_ep, group.cpu_group, local_ranks) - mooncake_ep.recover_ranks(group.cpu_group, local_ranks) + _wait_for_peer_state(group.device_group, local_ranks) + recover_ranks(group.device_group, local_ranks) + _wait_for_peer_state(group.cpu_group, local_ranks) + recover_ranks(group.cpu_group, local_ranks) _maybe_create_message_queue(group) _refresh_ep_members() @@ -411,9 +413,9 @@ def try_recover_ranks(global_ranks: List[int]) -> bool: def _join_world_group() -> None: - from mooncake import ep as mooncake_ep + from mooncake.pg import join_group - mooncake_ep.join_group(torch.distributed.group.WORLD) + join_group(torch.distributed.group.WORLD) def join_scale_process_group() -> None: @@ -424,14 +426,14 @@ def join_scale_process_group() -> None: def join_process_groups() -> None: """Rejoin WORLD and every launch-time parallel group after recovery.""" - from mooncake import ep as mooncake_ep + from mooncake.pg import join_group _join_world_group() for group in _iter_live_parallel_groups(): if group.world_size <= 1: continue - mooncake_ep.join_group(group.device_group) - mooncake_ep.join_group(group.cpu_group) + join_group(group.device_group) + join_group(group.cpu_group) _maybe_create_message_queue(group) _refresh_ep_members()