From 77d23a796e94c37f8657bc261856db725fa35891 Mon Sep 17 00:00:00 2001 From: Yoray Zack Date: Fri, 17 Jul 2026 01:53:44 +0300 Subject: [PATCH] [1/N] elastic-ep: Add runtime EP scale-up (#30164) --- python/sglang/srt/arg_groups/overrides.py | 2 +- python/sglang/srt/distributed/bootstrap.py | 19 +- .../sglang/srt/distributed/parallel_state.py | 147 +++++- python/sglang/srt/elastic_ep/elastic_ep.py | 347 ++++++++++--- python/sglang/srt/entrypoints/elastic_ep.py | 87 ++++ python/sglang/srt/entrypoints/engine.py | 11 +- python/sglang/srt/entrypoints/http_server.py | 16 +- python/sglang/srt/eplb/eplb_manager.py | 16 + python/sglang/srt/eplb/expert_distribution.py | 21 +- python/sglang/srt/eplb/expert_location.py | 195 ++++--- .../sglang/srt/layers/attention/dsa/utils.py | 13 +- python/sglang/srt/layers/dp_attention.py | 82 ++- .../srt/layers/moe/fused_moe_triton/layer.py | 16 +- .../srt/layers/moe/token_dispatcher/nixl.py | 113 ++++- python/sglang/srt/managers/communicator.py | 24 +- .../srt/managers/data_parallel_controller.py | 225 ++++++-- python/sglang/srt/managers/io_struct.py | 25 + python/sglang/srt/managers/scheduler.py | 175 ++++++- .../managers/scheduler_components/dp_attn.py | 117 ++++- .../scheduler_components/request_receiver.py | 5 +- .../srt/managers/tokenizer_control_mixin.py | 23 +- .../sglang/srt/managers/tokenizer_manager.py | 83 ++- python/sglang/srt/managers/tp_worker.py | 18 +- .../srt/model_executor/forward_batch_info.py | 29 +- .../sglang/srt/model_executor/model_runner.py | 329 +++++++++++- .../load_model_utils.py | 10 +- python/sglang/srt/runtime_context.py | 2 + python/sglang/srt/server_args.py | 210 +++++++- python/sglang/srt/utils/common.py | 9 +- test/manual/ep/test_elastic_scale.py | 479 ++++++++++++++++++ .../entrypoints/test_v1_loads_aggregate.py | 1 + ...e_logical_to_rank_dispatch_physical_map.py | 24 +- .../managers/test_data_parallel_controller.py | 5 +- test/registered/unit/test_model_overrides.py | 9 +- 34 files changed, 2549 insertions(+), 338 deletions(-) create mode 100644 python/sglang/srt/entrypoints/elastic_ep.py create mode 100644 test/manual/ep/test_elastic_scale.py diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 4d175a47a..1ef908905 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -1887,7 +1887,7 @@ def _page_size_default(view: Any) -> dict: @register_post_process def _data_parallelism_defaults(view: Any) -> dict: - if view.dp_size == 1: + if view.dp_size == 1 and view.ep_join_mode != "scale": return {"enable_dp_attention": False, "enable_dp_lm_head": False} return {} diff --git a/python/sglang/srt/distributed/bootstrap.py b/python/sglang/srt/distributed/bootstrap.py index 7daacf9fa..ab7be1f9e 100644 --- a/python/sglang/srt/distributed/bootstrap.py +++ b/python/sglang/srt/distributed/bootstrap.py @@ -225,15 +225,24 @@ def _init_parallel_groups( moe_dp_size: int, dcp_size: int, ) -> None: + is_ep_joiner = server_args.is_ep_joiner + is_scale_joiner = server_args.is_ep_scale_joiner + rank_offset = server_args.ep_join_rank_offset if is_scale_joiner else 0 + world_size = ( + rank_offset + tp_size * pp_size if is_scale_joiner else tp_size * pp_size + ) + rank = rank_offset + tp_size * pp_rank + tp_rank + init_distributed_environment( backend=backend, - world_size=tp_size * pp_size, - rank=tp_size * pp_rank + tp_rank, + world_size=world_size, + rank=rank, local_rank=gpu_id, distributed_init_method=dist_init_method, timeout=server_args.dist_timeout, moe_a2a_backend=server_args.moe_a2a_backend, - recovered_rank=server_args.elastic_ep_rejoin, + recovered_rank=is_ep_joiner, + max_world_size=server_args.max_ep_size, ) initialize_model_parallel( tensor_model_parallel_size=tp_size, @@ -245,7 +254,9 @@ def _init_parallel_groups( decode_context_parallel_size=dcp_size, duplicate_tp_group=server_args.enable_pdmux, enable_symm_mem=server_args.enable_symm_mem, - recovered_rank=server_args.elastic_ep_rejoin, + recovered_rank=is_ep_joiner, + rank_offset=rank_offset, + max_world_size=server_args.max_ep_size, ) initialize_dp_attention( server_args=server_args, diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 42ab1b413..3be32d3f1 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -270,6 +270,8 @@ class GroupCoordinator: group_name: Optional[str] = None, gloo_timeout: timedelta = timedelta(seconds=120 * 60), recovered_rank: bool = False, + rank_offset: int = 0, + max_world_size: Optional[int] = None, ): # Set group info group_name = group_name or "anonymous" @@ -278,6 +280,9 @@ class GroupCoordinator: # Set rank info self.rank = torch.distributed.get_rank() + # Joiner group ranks are local; shift them into global rank space. + if rank_offset > 0: + group_ranks = [[r + rank_offset for r in ranks] for ranks in group_ranks] self.local_rank = local_rank self.device_group = None self.cpu_group = None @@ -299,25 +304,57 @@ class GroupCoordinator: self.device_module = torch.get_device_module(self.device) for ranks in group_ranks: - active_ranks = torch.ones(len(ranks), dtype=torch.int32, device=self.device) - active_ranks_cpu = torch.ones(len(ranks), dtype=torch.int32) subgroup_timeout = _MODEL_PARALLEL_GROUP_TIMEOUT if "mooncake" in torch_distributed_backend: from mooncake.ep import MooncakeBackendOptions + pg_active_size = len(ranks) + if not recovered_rank and max_world_size is not None: + assert max_world_size >= len(ranks), ( + f"max_world_size ({max_world_size}) must be >= " + f"group size ({len(ranks)})" + ) + pg_active_size = max_world_size + + pg_active_ranks = torch.zeros( + pg_active_size, dtype=torch.int32, device=self.device + ) + pg_active_ranks[: len(ranks)] = 1 + pg_active_ranks_cpu = torch.zeros(pg_active_size, dtype=torch.int32) + pg_active_ranks_cpu[: len(ranks)] = 1 + + if not recovered_rank and max_world_size is not None: + dev_opts = MooncakeBackendOptions( + pg_active_ranks, recovered_rank, max_world_size + ) + cpu_opts = MooncakeBackendOptions( + pg_active_ranks_cpu, recovered_rank, max_world_size + ) + else: + dev_opts = MooncakeBackendOptions(pg_active_ranks, recovered_rank) + cpu_opts = MooncakeBackendOptions( + pg_active_ranks_cpu, recovered_rank + ) + + active_ranks = pg_active_ranks[: len(ranks)] + active_ranks_cpu = pg_active_ranks_cpu[: len(ranks)] device_group = torch.distributed.new_group( ranks, backend="mooncake", - pg_options=MooncakeBackendOptions(active_ranks, recovered_rank), + pg_options=dev_opts, timeout=subgroup_timeout, ) cpu_group = torch.distributed.new_group( ranks, backend="mooncake-cpu", - pg_options=MooncakeBackendOptions(active_ranks_cpu, recovered_rank), + pg_options=cpu_opts, timeout=subgroup_timeout, ) else: + active_ranks = torch.ones( + len(ranks), dtype=torch.int32, device=self.device + ) + active_ranks_cpu = torch.ones(len(ranks), dtype=torch.int32) pg_options = get_torch_distributed_pg_options(group_name) device_group = torch.distributed.new_group( ranks, @@ -1657,6 +1694,8 @@ def init_model_parallel_group( use_mscclpp_allreduce: Optional[bool] = None, use_torch_symm_mem_allreduce: Optional[bool] = None, recovered_rank: bool = False, + rank_offset: int = 0, + max_world_size: Optional[int] = None, ) -> GroupCoordinator: if use_custom_allreduce is None: use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE @@ -1682,6 +1721,8 @@ def init_model_parallel_group( use_message_queue_broadcaster=use_message_queue_broadcaster, group_name=group_name, recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) @@ -1842,7 +1883,12 @@ def get_default_distributed_backend(device: str) -> str: return _DEVICE_TO_DISTRIBUTED_BACKEND.get(device, "gloo") -def _create_global_tcp_store(rank: int, world_size: int) -> None: +def _create_global_tcp_store( + rank: int, + world_size: int, + dist_init_addr: Optional[str] = None, + allow_dynamic_membership: bool = False, +) -> None: """Create a global TCPStore for coordination across ranks. This function creates a TCPStore that all ranks can use for coordination @@ -1850,42 +1896,51 @@ def _create_global_tcp_store(rank: int, world_size: int) -> None: """ from torch.distributed import TCPStore - master_ip = os.environ.get("MASTER_ADDR") + base_store_port = envs.SGLANG_TCP_STORE_PORT.get() + master_ip = os.environ.get("MASTER_ADDR") + if not master_ip and allow_dynamic_membership and dist_init_addr: + addr = dist_init_addr + if addr.startswith("tcp://"): + addr = addr[len("tcp://") :] + master_ip = addr.rsplit(":", 1)[0] if not master_ip: logger.warning( "Could not determine master IP for global TCPStore. " "Broadcasting from rank 0 to all ranks." ) - - base_store_port = envs.SGLANG_TCP_STORE_PORT.get() - - # Rank 0 gets its local IP and broadcasts it to all ranks - # Use broadcast_object_list which works with any backend (handles CPU/GPU automatically) - if not master_ip: if rank == 0: master_ip = get_local_ip_auto() ip_list = [master_ip] else: ip_list = [None] - torch.distributed.broadcast_object_list(ip_list, src=0) master_ip = ip_list[0] try: - tcp_store = TCPStore( - host_name=master_ip, - port=base_store_port, - world_size=world_size, - is_master=(rank == 0), - ) + if allow_dynamic_membership: + is_master = rank == 0 + tcp_store = TCPStore( + host_name=master_ip, + port=base_store_port, + is_master=is_master, + wait_for_workers=False, + ) + else: + is_master = rank == 0 + tcp_store = TCPStore( + host_name=master_ip, + port=base_store_port, + world_size=world_size, + is_master=is_master, + ) set_global_tcp_store(tcp_store) logger.info( - "Created global TCPStore at %s:%d (rank=%d, world_size=%d)", + "Created global TCPStore at %s:%d (rank=%d, is_master=%s)", master_ip, base_store_port, rank, - world_size, + is_master, ) except Exception as e: logger.warning( @@ -1906,6 +1961,7 @@ def init_distributed_environment( timeout: Optional[int] = None, moe_a2a_backend: Optional[str] = None, recovered_rank: bool = False, + max_world_size: Optional[int] = None, ): logger.debug( "world_size=%d rank=%d local_rank=%d " "distributed_init_method=%s backend=%s", @@ -1942,9 +1998,16 @@ def init_distributed_environment( if backend == "mooncake": from mooncake.ep import MooncakeBackendOptions - # Setting "cuda" as device here is safe, as it is guarded under the mooncake case - active_ranks = torch.ones(world_size, dtype=torch.int32, device="cuda") - pg_options = MooncakeBackendOptions(active_ranks, recovered_rank) + use_max_ws = max_world_size and max_world_size > world_size + ar_size = max_world_size if use_max_ws else world_size + active_ranks = torch.zeros(ar_size, dtype=torch.int32, device="cuda") + active_ranks[:world_size] = 1 + if use_max_ws: + pg_options = MooncakeBackendOptions( + active_ranks, recovered_rank, max_world_size + ) + else: + pg_options = MooncakeBackendOptions(active_ranks, recovered_rank) else: pg_options = get_torch_distributed_pg_options() @@ -1960,7 +2023,15 @@ def init_distributed_environment( # Create a global TCPStore for coordination (used by NIXL) if moe_a2a_backend == "nixl": - _create_global_tcp_store(rank, world_size) + _create_global_tcp_store( + rank, + world_size, + dist_init_addr=distributed_init_method, + allow_dynamic_membership=( + recovered_rank + or (max_world_size is not None and max_world_size > world_size) + ), + ) # set the local rank # local_rank is not available in torch ProcessGroup, @@ -1996,6 +2067,8 @@ def initialize_model_parallel( duplicate_tp_group: bool = False, enable_symm_mem: bool = False, recovered_rank: bool = False, + rank_offset: int = 0, + max_world_size: Optional[int] = None, ) -> None: """ Initialize model parallel groups. @@ -2048,9 +2121,15 @@ def initialize_model_parallel( """ # Get world size and rank. Ensure some consistencies. assert torch.distributed.is_initialized() - world_size: int = torch.distributed.get_world_size() backend = backend or torch.distributed.get_backend(get_world_group().device_group) + # Joiners construct their local TP/PP layout in global rank space. + world_size: int = ( + tensor_model_parallel_size * pipeline_model_parallel_size + if recovered_rank + else torch.distributed.get_world_size() + ) + if world_size != tensor_model_parallel_size * pipeline_model_parallel_size: raise RuntimeError( f"world_size ({world_size}) is not equal to " @@ -2096,6 +2175,8 @@ def initialize_model_parallel( use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), group_name="tp", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) if duplicate_tp_group: @@ -2110,6 +2191,8 @@ def initialize_model_parallel( use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), group_name="pdmux_prefill_tp", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) if _TP.pynccl_comm: _TP.pynccl_comm.disabled = False @@ -2172,6 +2255,8 @@ def initialize_model_parallel( use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), group_name="attn_cp", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) from sglang.srt.layers.sampler import SYNC_TOKEN_IDS_ACROSS_TP @@ -2208,6 +2293,8 @@ def initialize_model_parallel( use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), group_name="attention_tp", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) moe_ep_size = expert_model_parallel_size @@ -2239,6 +2326,8 @@ def initialize_model_parallel( backend, group_name="moe_dp", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) global _MOE_EP @@ -2267,6 +2356,8 @@ def initialize_model_parallel( use_custom_allreduce=False, group_name="moe_ep", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) global _MOE_TP @@ -2295,6 +2386,8 @@ def initialize_model_parallel( use_custom_allreduce=False, group_name="moe_tp", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) # Build the pipeline model-parallel groups. @@ -2315,6 +2408,8 @@ def initialize_model_parallel( use_custom_allreduce=False, group_name="pp", recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, ) diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index ece4fd31d..8ceb53a9e 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -3,11 +3,12 @@ from __future__ import annotations import logging import time from dataclasses import dataclass -from typing import TYPE_CHECKING, Iterator, List, Optional +from typing import TYPE_CHECKING, Callable, Iterator, List, Optional import torch 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 @@ -17,12 +18,39 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +_SCALE_COHORT_KEY_PREFIX = "elastic_ep/scale_cohort" + + +def register_scale_cohort(rank_offset: int, target_ep_size: int) -> None: + store = get_global_tcp_store() + if store is None: + raise RuntimeError("Elastic EP scale-up requires the global TCPStore.") + store.set(f"{_SCALE_COHORT_KEY_PREFIX}/{rank_offset}", str(target_ep_size).encode()) + + +def get_scale_cohort_target(rank_offset: int) -> Optional[int]: + store = get_global_tcp_store() + if store is None: + return None + key = f"{_SCALE_COHORT_KEY_PREFIX}/{rank_offset}" + if not store.check([key]): + return None + return int(store.get(key).decode()) + @dataclass class ElasticEPState: active_ranks: Optional[torch.Tensor] last_active_ranks: Optional[torch.Tensor] active_ranks_cpu: Optional[torch.Tensor] + effective_ep_size: int = 0 + pending_ep_size: Optional[int] = None + scale_phase: str = "idle" + last_error: Optional[str] = None + pending_since: Optional[float] = None + original_ep_size: int = 0 + has_scaled: bool = False + ep_join_rank_offset: int = 0 def is_active_equal_last(self) -> bool: return torch.equal(self.active_ranks, self.last_active_ranks) @@ -37,13 +65,16 @@ class ElasticEPState: def reset(self): if self.active_ranks is not None: - self.active_ranks.fill_(1) + # Reserved slots stay inactive until their ranks join. + self.active_ranks.zero_() + self.active_ranks[: self.effective_ep_size] = 1 self.snapshot_active_to_last() self.sync_active_to_cpu() class ElasticEPStateManager: _instance: Optional[ElasticEPState] = None + _on_scale: Optional[Callable[[int, int], None]] = None @classmethod def instance(cls) -> ElasticEPState: @@ -55,16 +86,53 @@ class ElasticEPStateManager: return cls._instance if server_args.elastic_ep_backend is not None: - cls._instance = cls._build_state(ep_size=None, device=None) - if server_args.elastic_ep_rejoin: - # Mask out peer ranks to perform cuda graph capture on its own - cls._instance.active_ranks.zero_() - cls._instance.active_ranks[torch.distributed.get_rank()] = 1 - cls._instance.snapshot_active_to_last() - cls._instance.sync_active_to_cpu() + world_size = torch.distributed.get_world_size() + active_rank_capacity = server_args.max_ep_size or world_size + assert active_rank_capacity >= world_size, ( + f"--max-ep-size ({active_rank_capacity}) must be >= " + f"world_size ({world_size})." + ) + + inst = cls._build_state(ep_size=active_rank_capacity, device=None) + inst.effective_ep_size = world_size + inst.original_ep_size = world_size + if active_rank_capacity > world_size: + inst.active_ranks[world_size:].zero_() + inst.snapshot_active_to_last() + inst.sync_active_to_cpu() + + if server_args.moe_a2a_backend == "nixl": + cls._on_scale = cls._on_scale_nixl + + inst.ep_join_rank_offset = server_args.ep_join_rank_offset + if server_args.is_ep_joiner: + cls._init_joiner_state(inst, server_args) + + cls._instance = inst return cls._instance + @classmethod + def _init_joiner_state(cls, inst: ElasticEPState, server_args: ServerArgs) -> None: + global_rank = torch.distributed.get_rank() + inst.active_ranks.zero_() + inst.active_ranks[global_rank] = 1 + inst.snapshot_active_to_last() + inst.sync_active_to_cpu() + + if server_args.ep_join_mode == "scale": + inst.effective_ep_size = ( + server_args.ep_join_rank_offset + server_args.tp_size + ) + inst.original_ep_size = ( + server_args.elastic_ep_initial_size or server_args.ep_join_rank_offset + ) + inst.has_scaled = True + else: + world_size = torch.distributed.get_world_size() + inst.effective_ep_size = world_size + inst.original_ep_size = world_size + @staticmethod def _select_device() -> torch.device: if is_cuda(): @@ -94,39 +162,197 @@ class ElasticEPStateManager: return torch.ones(size, dtype=torch.int32, device=dev) + @classmethod + def request_scale(cls, n: int) -> bool: + inst = cls._instance + if inst is None: + return False + if ( + inst.pending_ep_size is not None + or inst.scale_phase == "recovery_unsupported" + ): + return False + inst.pending_ep_size = n + inst.scale_phase = "waiting_for_cohort" + inst.last_error = None + inst.pending_since = time.monotonic() + return True -# --------------------------------------------------------------------------- -# Helpers for elastic EP recovery -# --------------------------------------------------------------------------- + @classmethod + def begin_scale(cls) -> bool: + inst = cls._instance + if ( + inst is None + or inst.pending_ep_size is None + or inst.scale_phase != "waiting_for_cohort" + ): + return False + inst.scale_phase = "pending" + return True + + @classmethod + def mark_joining(cls) -> None: + cls._mark_phase("joining") + + @classmethod + def mark_configuring_data_plane(cls) -> None: + cls._mark_phase("configuring_data_plane") + + @classmethod + def mark_syncing_new_world(cls) -> None: + cls._mark_phase("syncing_new_world") + + @classmethod + def _mark_phase(cls, phase: str) -> None: + inst = cls._instance + if inst is not None and inst.pending_ep_size is not None: + inst.scale_phase = phase + + @classmethod + def commit_scale(cls) -> None: + inst = cls._instance + if inst is None or inst.pending_ep_size is None: + return + inst.effective_ep_size = inst.pending_ep_size + inst.pending_ep_size = None + inst.has_scaled = True + inst.scale_phase = "serving_expanded" + inst.last_error = None + inst.pending_since = None + inst.reset() + + @classmethod + def fail_scale(cls, error: str) -> None: + inst = cls._instance + if inst is None: + return + inst.pending_ep_size = None + inst.scale_phase = "failed" + inst.last_error = error + inst.pending_since = None + inst.reset() + + @classmethod + def fail_recovery(cls, error: str) -> None: + inst = cls._instance + if inst is None: + return + inst.scale_phase = "recovery_unsupported" + inst.last_error = error + + @classmethod + def get_effective_ep_size(cls) -> int: + inst = cls._instance + assert inst is not None, "Elastic EP state is not initialized." + return inst.effective_ep_size + + @classmethod + def get_pending_ep_size(cls) -> Optional[int]: + inst = cls._instance + if inst is None: + return None + return inst.pending_ep_size + + @classmethod + def get_scale_phase(cls) -> str: + inst = cls._instance + if inst is None: + return "disabled" + return inst.scale_phase + + @classmethod + def get_last_error(cls) -> Optional[str]: + inst = cls._instance + if inst is None: + return None + return inst.last_error + + @classmethod + def get_ep_join_rank_offset(cls) -> int: + inst = cls._instance + if inst is None: + return 0 + return inst.ep_join_rank_offset + + @classmethod + def on_scale(cls, from_ep_size: int, to_ep_size: int) -> None: + if cls._on_scale is not None: + cls._on_scale(from_ep_size, to_ep_size) + + @staticmethod + def _on_scale_nixl(from_ep_size: int, to_ep_size: int) -> None: + from sglang.srt.layers.moe.token_dispatcher.nixl import NixlEPBuffer + + NixlEPBuffer.on_scale(from_ep_size, to_ep_size) + + @classmethod + def is_scaling(cls) -> bool: + """Return whether a scale or recovery operation is pending. + + The CPU snapshot is authoritative because rank polling uses it too. + """ + inst = cls._instance + if inst is None or inst.active_ranks_cpu is None: + return False + if inst.scale_phase == "recovery_unsupported": + return False + if inst.pending_ep_size is not None: + return True + active_count = int(inst.active_ranks_cpu[: inst.effective_ep_size].sum().item()) + return active_count < inst.effective_ep_size + + +def elastic_expanded_world_enabled() -> bool: + """Return whether execution uses ranks admitted after server launch. + + Launch-time TP groups exclude ranks admitted during scale-up. + """ + from sglang.srt.runtime_context import get_server_args + + inst = ElasticEPStateManager.instance() + if inst is None: + return False + sa = get_server_args() + if sa.max_ep_size is None: + return False + active_target_size = inst.effective_ep_size + if inst.pending_ep_size is not None and inst.scale_phase in ( + "configuring_data_plane", + "syncing_new_world", + ): + active_target_size = inst.pending_ep_size + + return active_target_size > inst.original_ep_size + + +def _refresh_ep_members() -> None: + from sglang.srt.layers.moe.token_dispatcher.mooncake import EPBuffer + + buffer = EPBuffer.get_existing_buffer() + if buffer is not None: + buffer.update_ep_member() _PEER_STATE_POLL_INTERVAL_SEC = 0.01 -def _get_process_group_backend(process_group, device: str): - return process_group - - def _iter_live_parallel_groups() -> Iterator[parallel_state.GroupCoordinator]: groups = [] for group_ref in parallel_state._groups.values(): group = group_ref() if group is not None: groups.append(group) - for group in sorted(groups, key=lambda x: x.unique_name): - yield group + yield from sorted(groups, key=lambda group: group.unique_name) def _map_global_to_group_local_ranks( group_ranks: List[int], global_ranks: List[int] ) -> List[int]: - rank_to_local = {rank: idx for idx, rank in enumerate(group_ranks)} + rank_to_local = {rank: index for index, rank in enumerate(group_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: - # Relaunched ranks become recoverable asynchronously, so we poll until the - # target backend reports all requested peers as ready. while not all(mooncake_ep.get_peer_state(backend, ranks)): time.sleep(_PEER_STATE_POLL_INTERVAL_SEC) @@ -142,66 +368,71 @@ def _maybe_create_message_queue(group) -> None: ) -def _refresh_ep_members() -> None: - from sglang.srt.layers.moe.token_dispatcher.mooncake import EPBuffer +def _try_recover_world(global_ranks: List[int]) -> bool: + from mooncake import ep as mooncake_ep - EPBuffer.get_existing_buffer().update_ep_member() + world_backend = torch.distributed.group.WORLD + if not all(mooncake_ep.get_peer_state(world_backend, global_ranks)): + return False + + mooncake_ep.recover_ranks(world_backend, global_ranks) + logger.debug("[Elastic EP][recover] WORLD recover_ranks(%s) done", global_ranks) + return True + + +def try_admit_scale_ranks(global_ranks: List[int]) -> bool: + """Admit append-only ranks into the expandable WORLD group.""" + if not _try_recover_world(global_ranks): + return False + + _refresh_ep_members() + return True def try_recover_ranks(global_ranks: List[int]) -> bool: - from mooncake import ep as mooncake_ep - - world_backend = _get_process_group_backend(torch.distributed.group.WORLD, "cuda") - if not all(mooncake_ep.get_peer_state(world_backend, global_ranks)): - # The relaunched ranks have not finished initializing yet. + """Recover ranks in WORLD and every launch-time parallel group.""" + if not _try_recover_world(global_ranks): return False - # Recover the world backend first, then recover each derived process group - # using ranks mapped into that group's local rank space. - mooncake_ep.recover_ranks(world_backend, global_ranks) + from mooncake import ep as mooncake_ep for group in _iter_live_parallel_groups(): - group_local_ranks = _map_global_to_group_local_ranks(group.ranks, global_ranks) - if not group_local_ranks: + local_ranks = _map_global_to_group_local_ranks(group.ranks, global_ranks) + if not local_ranks: continue - device_backend = _get_process_group_backend(group.device_group, "cuda") - _wait_for_peer_state(mooncake_ep, device_backend, group_local_ranks) - mooncake_ep.recover_ranks(device_backend, group_local_ranks) - - cpu_backend = _get_process_group_backend(group.cpu_group, "cpu") - _wait_for_peer_state(mooncake_ep, cpu_backend, group_local_ranks) - mooncake_ep.recover_ranks(cpu_backend, group_local_ranks) + _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) _maybe_create_message_queue(group) _refresh_ep_members() return True -def join_process_groups(): +def _join_world_group() -> None: from mooncake import ep as mooncake_ep - def join_backend(label: str, backend) -> None: - logger.info("Recovered rank joining Mooncake backend %s", label) - mooncake_ep.join_group(backend) + mooncake_ep.join_group(torch.distributed.group.WORLD) - join_backend( - "default_world", - _get_process_group_backend(torch.distributed.group.WORLD, "cuda"), - ) +def join_scale_process_group() -> None: + """Join the expandable WORLD group for an append-only scale operation.""" + _join_world_group() + _refresh_ep_members() + + +def join_process_groups() -> None: + """Rejoin WORLD and every launch-time parallel group after recovery.""" + from mooncake import ep as mooncake_ep + + _join_world_group() for group in _iter_live_parallel_groups(): if group.world_size <= 1: continue - - join_backend( - f"{group.unique_name}:device", - _get_process_group_backend(group.device_group, "cuda"), - ) - join_backend( - f"{group.unique_name}:cpu", - _get_process_group_backend(group.cpu_group, "cpu"), - ) + mooncake_ep.join_group(group.device_group) + mooncake_ep.join_group(group.cpu_group) _maybe_create_message_queue(group) _refresh_ep_members() diff --git a/python/sglang/srt/entrypoints/elastic_ep.py b/python/sglang/srt/entrypoints/elastic_ep.py new file mode 100644 index 000000000..34525f55f --- /dev/null +++ b/python/sglang/srt/entrypoints/elastic_ep.py @@ -0,0 +1,87 @@ +"""Elastic EP scaling HTTP endpoints for dp_attention deployments.""" + +import json +from http import HTTPStatus + +from fastapi import APIRouter, Request +from fastapi.responses import ORJSONResponse + +from sglang.srt.utils.auth import AuthLevel, auth_level + +router = APIRouter() + + +@router.post("/scale_elastic_ep") +@auth_level(AuthLevel.ADMIN_OPTIONAL) +async def scale_elastic_ep(raw_request: Request): + """Request an asynchronous EP scale-up.""" + try: + body = await raw_request.json() + except (json.JSONDecodeError, UnicodeDecodeError) as e: + return ORJSONResponse( + {"error": f"Invalid JSON: {e}"}, + status_code=HTTPStatus.BAD_REQUEST, + ) + + if not isinstance(body, dict): + return ORJSONResponse( + {"error": "request body must be a JSON object"}, + status_code=HTTPStatus.BAD_REQUEST, + ) + + new_ep_size = body.get("new_ep_size") + if ( + not isinstance(new_ep_size, int) + or isinstance(new_ep_size, bool) + or new_ep_size <= 0 + ): + return ORJSONResponse( + {"error": "new_ep_size must be a positive integer"}, + status_code=HTTPStatus.BAD_REQUEST, + ) + + from sglang.srt.entrypoints.http_server import _global_state + from sglang.srt.managers.io_struct import ScaleElasticEPReqInput + + if _global_state.tokenizer_manager.server_args.elastic_ep_backend is None: + return ORJSONResponse( + {"error": "elastic EP is not enabled (set --elastic-ep-backend)"}, + status_code=HTTPStatus.NOT_FOUND, + ) + + result = await _global_state.tokenizer_manager.scale_elastic_ep( + ScaleElasticEPReqInput(new_ep_size=new_ep_size) + ) + + if not result.success: + return ORJSONResponse( + {"error": result.message}, + status_code=( + HTTPStatus.CONFLICT + if result.pending_ep_size is not None + else HTTPStatus.BAD_REQUEST + ), + ) + + return ORJSONResponse( + { + "message": result.message, + "old_ep_size": result.old_ep_size, + "new_ep_size": result.new_ep_size, + } + ) + + +@router.get("/is_scaling_elastic_ep") +@auth_level(AuthLevel.ADMIN_OPTIONAL) +async def is_scaling_elastic_ep(raw_request: Request): + """Return the tokenizer's mirrored Elastic EP scale state.""" + from sglang.srt.entrypoints.http_server import _global_state + + if _global_state.tokenizer_manager.server_args.elastic_ep_backend is None: + return ORJSONResponse( + {"error": "elastic EP is not enabled (set --elastic-ep-backend)"}, + status_code=HTTPStatus.NOT_FOUND, + ) + + return ORJSONResponse(_global_state.tokenizer_manager.get_elastic_ep_state()) diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 6e822d6af..7fd810c47 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -603,8 +603,11 @@ class Engine(EngineScoreMixin, EngineBase): scheduler_procs is None for RayEngine (uses Ray actors instead). """ scheduler_procs = [] + use_dp_controller = ( + server_args.dp_size > 1 or server_args.ep_join_mode == "scale" + ) - if server_args.dp_size == 1: + if not use_dp_controller: # Launch tensor parallel scheduler processes memory_saver_adapter = TorchMemorySaverAdapter.create( enable=server_args.enable_memory_saver @@ -678,8 +681,7 @@ class Engine(EngineScoreMixin, EngineBase): def wait_for_ready(): infos = _wait_for_scheduler_ready(scheduler_pipe_readers, scheduler_procs) scheduler_infos.extend(infos) - # For dp_size > 1, collect child scheduler PIDs from the DP controller - if server_args.dp_size > 1: + if use_dp_controller: for info in infos: if SCHEDULER_PIDS_ARG in info: all_child_pids.extend(info[SCHEDULER_PIDS_ARG]) @@ -833,8 +835,7 @@ class Engine(EngineScoreMixin, EngineBase): run_expert_backup_manager(server_args, port_args) if server_args.node_rank >= 1: - # In multi-node cases, non-zero rank nodes do not need to run tokenizer or detokenizer, - # so they can just wait here. + # Non-zero-rank nodes do not run tokenizer processes. scheduler_init_result.wait_for_ready() if os.getenv("SGLANG_BLOCK_NONZERO_RANK_CHILDREN") == "0": diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index d7bae9633..e1dee6184 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -440,6 +440,10 @@ from sglang.srt.entrypoints.v1_loads import router as v1_loads_router app.include_router(v1_loads_router) +from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router + +app.include_router(elastic_ep_router) + def _anthropic_validation_message(raw_errors) -> str: """Render Pydantic-style errors for an Anthropic /v1/messages route. @@ -2213,8 +2217,16 @@ def _wait_and_warmup( if server_args.checkpoint_engine_wait_weights_before_ready: _wait_weights_ready() - # Send a warmup request - if not server_args.skip_server_warmup: + # Joiner schedulers are served through the primary after adoption. + skip_elastic_joiner_warmup = server_args.is_ep_scale_joiner + if skip_elastic_joiner_warmup: + logger.debug( + "[Elastic EP] Skipping server warmup for elastic joiner " + "(ep_join_mode=%s)", + server_args.ep_join_mode, + ) + + if not server_args.skip_server_warmup and not skip_elastic_joiner_warmup: if not execute_warmup_func(server_args): return else: diff --git a/python/sglang/srt/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index 1d6c7ff06..360ac93b7 100644 --- a/python/sglang/srt/eplb/eplb_manager.py +++ b/python/sglang/srt/eplb/eplb_manager.py @@ -52,6 +52,8 @@ class EPLBManager: self._server_args.eplb_rebalance_layers_per_chunk ) self._rebalance_num_iterations = self._server_args.eplb_rebalance_num_iterations + self._rebalance_disabled_reason = None + self._rebalance_disabled_logged = False # Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented. assert ( @@ -74,6 +76,11 @@ class EPLBManager: def reset_generator(self): self._main_generator = self._entrypoint() + def disable_rebalance(self, reason: str): + self._rebalance_disabled_reason = reason + self._rebalance_disabled_logged = False + self.reset_generator() + # can be more complex if needed def _entrypoint(self): while True: @@ -83,6 +90,15 @@ class EPLBManager: yield from self.rebalance() def rebalance(self): + if self._rebalance_disabled_reason is not None: + if not self._rebalance_disabled_logged: + logger.debug( + "[EPLBManager] rebalance disabled: %s", + self._rebalance_disabled_reason, + ) + self._rebalance_disabled_logged = True + return + logger.info("[EPLBManager] rebalance start") enable_timing = self._rebalance_layers_per_chunk is None diff --git a/python/sglang/srt/eplb/expert_distribution.py b/python/sglang/srt/eplb/expert_distribution.py index faed8e44c..d23703f7c 100644 --- a/python/sglang/srt/eplb/expert_distribution.py +++ b/python/sglang/srt/eplb/expert_distribution.py @@ -331,7 +331,9 @@ class _SinglePassGatherer(ABC): return _SelectExpertsSinglePassGatherer(expert_location_metadata, rank) elif server_args.deepep_mode == "low_latency": return _DeepepLowLatencySinglePassGatherer( - expert_location_metadata, rank + expert_location_metadata, + rank, + elastic_ep_enabled=server_args.elastic_ep_backend is not None, ) else: raise NotImplementedError @@ -574,13 +576,26 @@ class _DeepepNormalSinglePassGatherer(_LayerBasedCpuSinglePassGatherer): class _DeepepLowLatencySinglePassGatherer(_LayerBasedGpuSinglePassGatherer): - def __init__(self, *args, **kwargs): + def __init__(self, *args, elastic_ep_enabled: bool = False, **kwargs): super().__init__(*args, **kwargs, enable_global_physical_experts=False) + self._elastic_ep_enabled = elastic_ep_enabled def on_deepep_dispatch_low_latency( self, layer_idx: int, local_physical_count_of_layer: torch.Tensor ): - # Most naive implementation, can optimize later + if local_physical_count_of_layer.shape[0] != self._data.shape[1]: + if not self._elastic_ep_enabled: + self._data[layer_idx, :] += local_physical_count_of_layer + return + + n = self._data.shape[1] + if local_physical_count_of_layer.shape[0] > n: + local_physical_count_of_layer = local_physical_count_of_layer[:n] + else: + local_physical_count_of_layer = torch.nn.functional.pad( + local_physical_count_of_layer, + (0, n - local_physical_count_of_layer.shape[0]), + ) self._data[layer_idx, :] += local_physical_count_of_layer diff --git a/python/sglang/srt/eplb/expert_location.py b/python/sglang/srt/eplb/expert_location.py index 78b38e7f4..e4f5de8d1 100644 --- a/python/sglang/srt/eplb/expert_location.py +++ b/python/sglang/srt/eplb/expert_location.py @@ -32,6 +32,25 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +def _prefer_same_node_experts(server_args: ServerArgs) -> bool: + from sglang.srt.elastic_ep.elastic_ep import elastic_expanded_world_enabled + + return server_args.ep_join_mode != "scale" and not elastic_expanded_world_enabled() + + +def _compute_elastic_expert_layout( + base_num_physical_experts: int, + initial_ep_size: int, + effective_ep_size: int, +) -> tuple[int, int]: + assert base_num_physical_experts % initial_ep_size == 0 + num_local_physical_experts = base_num_physical_experts // initial_ep_size + return ( + num_local_physical_experts * effective_ep_size, + num_local_physical_experts, + ) + + @dataclass class ExpertLocationMetadata: physical_to_logical_map: torch.Tensor # (layers, num_physical_experts) @@ -39,6 +58,7 @@ class ExpertLocationMetadata: logical_to_all_physical_map: torch.Tensor # (layers, num_logical_experts, X) logical_to_all_physical_map_cpu: torch.Tensor # CPU copy for performance logical_to_all_physical_map_num_valid: torch.Tensor # (layers, num_logical_experts) + ep_size: int # (layers, num_logical_experts) logical_to_rank_dispatch_physical_map: Optional[torch.Tensor] @@ -62,11 +82,6 @@ class ExpertLocationMetadata: def num_logical_experts(self) -> int: return self.logical_to_all_physical_map.shape[1] - @property - def ep_size(self): - # TODO change when EP size != world size - return torch.distributed.get_world_size() - def __post_init__(self): num_layers_0, num_physical_experts_0 = self.physical_to_logical_map.shape num_layers_1, num_logical_experts_0, num_physical_experts_1 = ( @@ -96,10 +111,16 @@ class ExpertLocationMetadata: num_layers = model_config_for_expert_location.num_layers num_logical_experts = model_config_for_expert_location.num_logical_experts + base_num_physical_experts = common["base_num_physical_experts"] physical_to_logical_map = ( - torch.arange(0, num_physical_experts).repeat(num_layers, 1) + torch.arange(0, base_num_physical_experts).repeat(num_layers, 1) % num_logical_experts ) + physical_to_logical_map = append_trivial_expert_slots( + physical_to_logical_map, + num_physical_experts - base_num_physical_experts, + num_logical_experts, + ) return ExpertLocationMetadata.init_by_mapping( server_args, @@ -125,6 +146,15 @@ class ExpertLocationMetadata: return None model_config_for_expert_location = common["model_config_for_expert_location"] + if common["num_physical_experts"] > common["base_num_physical_experts"]: + if physical_to_logical_map.shape[-1] == common["base_num_physical_experts"]: + physical_to_logical_map = append_trivial_expert_slots( + physical_to_logical_map, + common["num_physical_experts"] + - common["base_num_physical_experts"], + model_config_for_expert_location.num_logical_experts, + ) + assert physical_to_logical_map.shape[-1] == common["num_physical_experts"] logical_to_all_physical_map = _compute_logical_to_all_physical_map( server_args=server_args, physical_to_logical_map=physical_to_logical_map, @@ -138,6 +168,7 @@ class ExpertLocationMetadata: ep_size=common["ep_size"], physical_to_logical_map=physical_to_logical_map, logical_to_all_physical_map=logical_to_all_physical_map, + moe_ep_rank=moe_ep_rank, ) @staticmethod @@ -195,16 +226,33 @@ class ExpertLocationMetadata: if model_config_for_expert_location is None: return None - num_physical_experts = ( + base_num_physical_experts = ( model_config_for_expert_location.num_logical_experts + server_args.ep_num_redundant_experts ) ep_size = server_args.ep_size - assert num_physical_experts % ep_size == 0 - num_local_physical_experts = num_physical_experts // ep_size + num_physical_experts = base_num_physical_experts + initial_ep_size = server_args.elastic_ep_initial_size + if initial_ep_size is not None: + if server_args.ep_join_mode == "scale": + ep_size = max( + ep_size, + server_args.ep_join_rank_offset + server_args.tp_size, + ) + num_physical_experts, num_local_physical_experts = ( + _compute_elastic_expert_layout( + base_num_physical_experts, + initial_ep_size, + ep_size, + ) + ) + else: + assert num_physical_experts % ep_size == 0 + num_local_physical_experts = num_physical_experts // ep_size return dict( model_config_for_expert_location=model_config_for_expert_location, + base_num_physical_experts=base_num_physical_experts, num_physical_experts=num_physical_experts, num_local_physical_experts=num_local_physical_experts, ep_size=ep_size, @@ -216,6 +264,7 @@ class ExpertLocationMetadata: ep_size: int, physical_to_logical_map: torch.Tensor, logical_to_all_physical_map: torch.Tensor, + moe_ep_rank: Optional[int] = None, ): _, num_physical_experts = physical_to_logical_map.shape @@ -235,14 +284,18 @@ class ExpertLocationMetadata: logical_to_all_physical_map=logical_to_all_physical_map_padded, logical_to_all_physical_map_cpu=logical_to_all_physical_map_padded.cpu(), logical_to_all_physical_map_num_valid=logical_to_all_physical_map_num_valid, + ep_size=ep_size, logical_to_rank_dispatch_physical_map=( compute_logical_to_rank_dispatch_physical_map( server_args=server_args, logical_to_all_physical_map=logical_to_all_physical_map, ep_size=ep_size, num_physical_experts=num_physical_experts, - # TODO improve when we have real EP rank - ep_rank=torch.distributed.get_rank() % ep_size, + ep_rank=( + moe_ep_rank + if moe_ep_rank is not None + else torch.distributed.get_rank() % ep_size + ), ) if server_args.ep_dispatch_algorithm == "static" else None @@ -425,58 +478,57 @@ def get_global_expert_location_metadata(): return get_resources().expert_location_metadata -def set_global_expert_location_metadata(value): +def set_global_expert_location_metadata(value, allow_overwrite=False): from sglang.srt.runtime_context import get_resources resources = get_resources() - assert resources.expert_location_metadata is None + if not allow_overwrite: + assert resources.expert_location_metadata is None resources.expert_location_metadata = value +def append_trivial_expert_slots( + physical_to_logical_map: torch.Tensor, + count: int, + num_logical_experts: int, + start: int = 0, +) -> torch.Tensor: + if count <= 0: + return physical_to_logical_map + new_slots = torch.arange( + start, + start + count, + dtype=physical_to_logical_map.dtype, + device=physical_to_logical_map.device, + ).unsqueeze(0) + new_slots = new_slots.expand(physical_to_logical_map.shape[0], -1) + return torch.cat([physical_to_logical_map, new_slots % num_logical_experts], dim=1) + + def broadcast_global_expert_location_metadata( - src_rank: int = 0, group: Optional[torch.distributed.ProcessGroup] = None -): - """Broadcast the global ExpertLocationMetadata from src_rank to all ranks. + model_config: ModelConfig, + moe_ep_rank: int, + src_rank: int = 0, + group: Optional[torch.distributed.ProcessGroup] = None, +) -> ExpertLocationMetadata: + from sglang.srt.runtime_context import get_server_args - This is used in Elastic EP rank recovery to ensure that all ranks (including - newly recovered ones) share exactly the same expert location metadata. - - Note: The caller must ensure src_rank is a healthy rank. In recovery scenarios, - this function is called after try_recover_ranks succeeds, at which point all - ranks (including src_rank=0) have recovered and are ready. - """ + server_args = get_server_args() metadata = get_global_expert_location_metadata() assert metadata is not None - # Ensure device tensors are contiguous before broadcasting in-place metadata.physical_to_logical_map = metadata.physical_to_logical_map.contiguous() - metadata.logical_to_all_physical_map = ( - metadata.logical_to_all_physical_map.contiguous() + torch.distributed.broadcast( + metadata.physical_to_logical_map, src=src_rank, group=group ) - metadata.logical_to_all_physical_map_num_valid = ( - metadata.logical_to_all_physical_map_num_valid.contiguous() - ) - if metadata.logical_to_rank_dispatch_physical_map is not None: - metadata.logical_to_rank_dispatch_physical_map = ( - metadata.logical_to_rank_dispatch_physical_map.contiguous() - ) - - device_tensors = [ + metadata = ExpertLocationMetadata.init_by_mapping( + server_args, + model_config, metadata.physical_to_logical_map, - metadata.logical_to_all_physical_map, - metadata.logical_to_all_physical_map_num_valid, - ] - if metadata.logical_to_rank_dispatch_physical_map is not None: - device_tensors.append(metadata.logical_to_rank_dispatch_physical_map) - - for tensor in device_tensors: - torch.distributed.broadcast(tensor, src=src_rank, group=group) - - # After broadcasting device tensors, refresh corresponding CPU copies - metadata.physical_to_logical_map_cpu = metadata.physical_to_logical_map.cpu() - metadata.logical_to_all_physical_map_cpu = ( - metadata.logical_to_all_physical_map.cpu() + moe_ep_rank=moe_ep_rank, ) + set_global_expert_location_metadata(metadata, allow_overwrite=True) + return metadata def _compute_logical_to_all_physical_map( @@ -506,10 +558,15 @@ def _compute_logical_to_all_physical_map( # Replace by the physical expert on local GPU or node if possible if moe_ep_rank is not None: - num_gpus_per_node = server_args.ep_size // server_args.nnodes num_local_gpu_physical_experts = num_physical_experts // ep_size + prefer_same_node = _prefer_same_node_experts(server_args) + num_gpus_per_node = ( + server_args.ep_size // server_args.nnodes if prefer_same_node else None + ) num_local_node_physical_experts = ( num_local_gpu_physical_experts * num_gpus_per_node + if num_gpus_per_node is not None + else None ) for layer_id in range(num_layers): for logical_expert_id in range(num_logical_experts): @@ -563,8 +620,15 @@ def compute_logical_to_rank_dispatch_physical_map( logical_to_all_physical_map = logical_to_all_physical_map.cpu() num_local_gpu_physical_experts = num_physical_experts // ep_size - num_gpus_per_node = server_args.ep_size // server_args.nnodes - num_local_node_physical_experts = num_local_gpu_physical_experts * num_gpus_per_node + prefer_same_node = _prefer_same_node_experts(server_args) + num_gpus_per_node = ( + server_args.ep_size // server_args.nnodes if prefer_same_node else None + ) + num_local_node_physical_experts = ( + num_local_gpu_physical_experts * num_gpus_per_node + if num_gpus_per_node is not None + else None + ) num_layers, num_logical_experts, _ = logical_to_all_physical_map.shape dtype = logical_to_all_physical_map.dtype @@ -633,8 +697,8 @@ def _find_nearest_expert( candidate_physical_expert_ids: List[int], num_local_gpu_physical_experts: int, moe_ep_rank: int, - num_gpus_per_node: int, - num_local_node_physical_experts: int, + num_gpus_per_node: Optional[int], + num_local_node_physical_experts: Optional[int], ) -> int: # 1. If only one candidate, return it directly if len(candidate_physical_expert_ids) == 1: @@ -652,18 +716,19 @@ def _find_nearest_expert( if len(same_gpu_physical_expert_ids) > 0: return same_gpu_physical_expert_ids[0] - # 3. Otherwise, prefer same-node experts - node_rank = moe_ep_rank // num_gpus_per_node - same_node_physical_expert_ids = [ - physical_expert_id - for physical_expert_id in candidate_physical_expert_ids - if _compute_node_id_of_physical_expert( - physical_expert_id, num_local_node_physical_experts - ) - == node_rank - ] - if len(same_node_physical_expert_ids) > 0: - return same_node_physical_expert_ids[0] + # Prefer same-node experts only when it narrows the candidate set. + if num_gpus_per_node is not None and num_local_node_physical_experts is not None: + node_rank = moe_ep_rank // num_gpus_per_node + same_node_physical_expert_ids = [ + physical_expert_id + for physical_expert_id in candidate_physical_expert_ids + if _compute_node_id_of_physical_expert( + physical_expert_id, num_local_node_physical_experts + ) + == node_rank + ] + if 0 < len(same_node_physical_expert_ids) < len(candidate_physical_expert_ids): + return same_node_physical_expert_ids[0] # 4. At last, leave it as -1 to indicate not found. return -1 diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index 62f47a57b..402d152b9 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -5,9 +5,7 @@ import torch import triton from sglang.srt.environ import envs -from sglang.srt.layers.dp_attention import ( - DpPaddingMode, -) +from sglang.srt.layers.dp_attention import DpPaddingMode from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( is_in_breakable_cuda_graph, ) @@ -164,9 +162,12 @@ def cal_padded_tokens(forward_batch: "ForwardBatch"): cp_align_size = get_cp_padding_align_size() for i in range(sync_group_size): global_num_tokens[i] = ceil_align(global_num_tokens[i], cp_align_size) - dp_padding_mode = DpPaddingMode.get_dp_padding_mode( - forward_batch.is_extend_in_batch, global_num_tokens - ) + # Reuse the mode selected when the DP buffer was prepared. + dp_padding_mode = forward_batch.dp_padding_mode + if dp_padding_mode is None: + dp_padding_mode = DpPaddingMode.get_dp_padding_mode( + forward_batch.is_extend_in_batch, global_num_tokens + ) if dp_padding_mode.is_max_len(): tokens = max(global_num_tokens) elif len(global_num_tokens) > 1: diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index eeed1d8a9..01c542d65 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -39,6 +39,29 @@ if TYPE_CHECKING: _ATTN_DP_RANK: Optional[int] = None _ATTN_DP_SIZE: Optional[int] = None + +def world_dp_gather_enabled() -> bool: + """Whether DP gathers should use expanded WORLD after joiner admission.""" + dp = get_flags().dp + return dp.use_world_group_for_gather and not dp.joiner_skip_all_gather + + +def enable_joiner_all_gather(): + get_flags().dp.joiner_skip_all_gather = False + + +def update_dp_attention_post_scale(new_dp_size: int, new_dp_rank: int): + global _ATTN_DP_SIZE, _ATTN_DP_RANK + _ATTN_DP_SIZE = new_dp_size + _ATTN_DP_RANK = new_dp_rank + get_flags().dp.use_world_group_for_gather = True + logger.debug( + "[Elastic EP] dp_attention switched to WORLD: dp_size=%d dp_rank=%d", + new_dp_size, + new_dp_rank, + ) + + _is_hip = is_hip() _USE_ROCM700A_WA = _is_hip and get_bool_env_var("SGLANG_USE_ROCM700A") @@ -287,7 +310,6 @@ def initialize_dp_attention( ) enable_dp_attention = server_args.enable_dp_attention dp_size = server_args.dp_size - moe_dense_tp_size = server_args.moe_dense_tp_size attn_cp_size = server_args.attn_cp_size dp.enabled = enable_dp_attention @@ -300,6 +322,11 @@ def initialize_dp_attention( ) _ATTN_DP_SIZE = dp_size if enable_dp_attention else 1 + if server_args.elastic_ep_backend is not None and server_args.max_ep_size: + _ATTN_DP_RANK = tp_rank + server_args.ep_join_rank_offset + if server_args.is_ep_scale_joiner: + dp.joiner_skip_all_gather = True + _DpGatheredBufferWrapper.set_metadata( hidden_size=model_config.hidden_size, dtype=model_config.dtype, @@ -408,17 +435,24 @@ def _dp_gather_via_all_reduce( ) # Input IDs are in int 32. We should use inplace_all_reduce for local case because of custom all reduce. - NUM_GPUS_PER_NODE = 8 - if ( - not local_tokens.dtype.is_floating_point - and get_tensor_model_parallel_world_size() <= NUM_GPUS_PER_NODE - ): - from sglang.srt.distributed.parallel_state import inplace_all_reduce - - inplace_all_reduce(global_tokens, group_name=get_tp_group().unique_name) - + if world_dp_gather_enabled(): + torch.distributed.all_reduce( + global_tokens, + op=torch.distributed.ReduceOp.SUM, + group=torch.distributed.group.WORLD, + ) else: - global_tokens[:] = tensor_model_parallel_all_reduce(global_tokens) + NUM_GPUS_PER_NODE = 8 + if ( + not local_tokens.dtype.is_floating_point + and get_tensor_model_parallel_world_size() <= NUM_GPUS_PER_NODE + ): + from sglang.srt.distributed.parallel_state import inplace_all_reduce + + inplace_all_reduce(global_tokens, group_name=get_tp_group().unique_name) + + else: + global_tokens[:] = tensor_model_parallel_all_reduce(global_tokens) def _dp_gather_via_all_gather( @@ -427,8 +461,17 @@ def _dp_gather_via_all_gather( forward_batch: ForwardBatch, is_partial: bool, ): + use_world = world_dp_gather_enabled() + if get_attn_tensor_model_parallel_world_size() == 1: - get_tp_group().all_gather_into_tensor(global_tokens, local_tokens) + if use_world: + torch.distributed.all_gather_into_tensor( + global_tokens, + local_tokens, + group=torch.distributed.group.WORLD, + ) + else: + get_tp_group().all_gather_into_tensor(global_tokens, local_tokens) return if not is_partial: @@ -438,7 +481,14 @@ def _dp_gather_via_all_gather( get_attn_tensor_model_parallel_world_size() )[get_attn_tensor_model_parallel_rank()] get_attn_tp_group().reduce_scatter_tensor(scattered_local_tokens, local_tokens) - get_tp_group().all_gather_into_tensor(global_tokens, scattered_local_tokens) + if use_world: + torch.distributed.all_gather_into_tensor( + global_tokens, + scattered_local_tokens, + group=torch.distributed.group.WORLD, + ) + else: + get_tp_group().all_gather_into_tensor(global_tokens, scattered_local_tokens) # Variable-length DP-MoE gather (reference https://github.com/ROCm/ATOM/pull/930): instead of padding every @@ -461,6 +511,7 @@ def is_dp_gatherv_active() -> bool: dp_reduce_scatter_tensor) consistent.""" return ( _USE_DP_GATHERV + and not world_dp_gather_enabled() and get_attn_tensor_model_parallel_world_size() == 1 and get_tensor_model_parallel_world_size() == get_attention_dp_size() and not _DpGatheredBufferWrapper.is_dp_max_padding() @@ -541,7 +592,10 @@ def _dp_gather( global_tokens, local_tokens, forward_batch, is_partial, _gatherv_sizes ) return - if forward_batch.dp_padding_mode.is_max_len(): + if ( + forward_batch.dp_padding_mode is not None + and forward_batch.dp_padding_mode.is_max_len() + ): _dp_gather_via_all_gather( global_tokens, local_tokens, forward_batch, is_partial ) diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index b326663e3..09f79c106 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -229,9 +229,19 @@ class FusedMoE(torch.nn.Module): else: num_shared_slots = num_fused_shared_experts - assert (num_experts - num_shared_slots) % self.moe_ep_size == 0 self._num_global_routed = num_experts - num_shared_slots - self._num_local_routed = self._num_global_routed // self.moe_ep_size + server_args = get_server_args() + if server_args.ep_join_mode == "scale": + storage_ep_size = server_args.elastic_ep_initial_size + assert storage_ep_size is not None + self._expert_storage_rank = ( + server_args.ep_join_rank_offset + self.moe_ep_rank + ) + else: + storage_ep_size = self.moe_ep_size + self._expert_storage_rank = self.moe_ep_rank + assert self._num_global_routed % storage_ep_size == 0 + self._num_local_routed = self._num_global_routed // storage_ep_size self.num_local_experts = self._num_local_routed + num_fused_shared_experts self._has_fused_shared = num_fused_shared_experts > 0 self._pending_fp8_shared_weights: dict[tuple[int, str], torch.Tensor] = {} @@ -712,7 +722,7 @@ class FusedMoE(torch.nn.Module): expert_data.copy_(loaded_weight) def _map_global_expert_id_to_local_expert_id(self, expert_id: int) -> int: - start_idx = self.moe_ep_rank * self._num_local_routed + start_idx = self._expert_storage_rank * self._num_local_routed end_idx = start_idx + self._num_local_routed if start_idx <= expert_id < end_idx: return expert_id - start_idx diff --git a/python/sglang/srt/layers/moe/token_dispatcher/nixl.py b/python/sglang/srt/layers/moe/token_dispatcher/nixl.py index dd04f7d09..090eb05fa 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/nixl.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/nixl.py @@ -56,10 +56,46 @@ class NixlEPBuffer: num_max_dispatch_tokens_per_rank=None, num_experts=None, num_local_experts=None, + connected_ep_size=None, + scale_to=None, + dispatch_ep_size=None, ) buffers["nixl_ep_state"] = state return state + @classmethod + def on_scale(cls, from_ep_size: int, to_ep_size: int) -> None: + """Schedule connections for newly admitted ranks.""" + state = cls._state() + state.scale_to = to_ep_size + state.dispatch_ep_size = to_ep_size + logger.debug( + "[Elastic EP][nixl] scheduling rank connections: old_ep_size=%d " + "new_ep_size=%d", + from_ep_size, + to_ep_size, + ) + + @classmethod + def _connect_ranks(cls, state, ranks: list, *, tag: str) -> None: + current_store = get_global_tcp_store() + if current_store is not None: + state.buffer.set_tcp_store_group(current_store) + + state.buffer.connect_ranks(ranks) + logger.debug( + "[Elastic EP][nixl] connect (%s) ranks=%s group_size=%s", + tag, + ranks, + state.buffer.group_size, + ) + + @classmethod + def _update_connections(cls, state, scale_to: int) -> None: + new_ranks = list(range(state.connected_ep_size, scale_to)) + cls._connect_ranks(state, new_ranks, tag="update") + state.connected_ep_size = scale_to + @classmethod def get_nixl_buffer( cls, @@ -72,6 +108,12 @@ class NixlEPBuffer: ): state = cls._state() if state.buffer is not None: + if ( + state.scale_to is not None + and state.connected_ep_size is not None + and state.scale_to > state.connected_ep_size + ): + cls._update_connections(state, state.scale_to) return state.buffer state.hidden_size = hidden_size @@ -79,23 +121,31 @@ class NixlEPBuffer: state.num_experts = num_experts state.num_local_experts = num_local_experts + rank = dist.get_rank(group) + world_size = dist.get_world_size(group) + # Joiner-local ranks are offset into the expanded global rank space. + offset = ElasticEPStateManager.get_ep_join_rank_offset() + global_rank = rank + offset + + from sglang.srt.runtime_context import get_server_args + + max_ep_size = get_server_args().max_ep_size or world_size + nixl_max_ranks = max_ep_size + num_rdma_bytes = 0 if deepep_mode.enable_normal(): raise NotImplementedError("Normal mode is not supported for Nixl EP yet.") if deepep_mode.enable_low_latency(): assert num_max_dispatch_tokens_per_rank != -1 - assert num_experts != -1 and num_experts % group.size() == 0 + assert num_experts > 0 and num_local_experts > 0 + max_num_global_experts = nixl_max_ranks * num_local_experts num_rdma_bytes = Buffer.get_rdma_size_hint( num_max_dispatch_tokens_per_rank, hidden_size, - group.size(), - num_experts, + nixl_max_ranks, + max_num_global_experts, ) - rank = dist.get_rank(group) - world_size = dist.get_world_size(group) - - # Get the global TCPStore for coordination tcp_store = get_global_tcp_store() if tcp_store is None: raise RuntimeError( @@ -104,23 +154,28 @@ class NixlEPBuffer: ) logger.info( - f"Using NIXL EP (world_size={world_size}, rank={rank}, " - f"num_experts={state.num_experts}, num_experts_per_rank={state.num_local_experts}) " + f"Using NIXL EP (world_size={world_size}, max_ep_size={max_ep_size}, " + f"rank={rank}, global_rank={global_rank}, offset={offset}, " + f"num_experts={state.num_experts}, " + f"num_experts_per_rank={state.num_local_experts}) " ) state.buffer = Buffer( - rank=rank, + rank=global_rank, tcp_store_group=tcp_store, ) state.buffer.update_memory_buffers( - num_ranks=world_size, + num_ranks=nixl_max_ranks, num_experts_per_rank=state.num_local_experts, num_rdma_bytes=num_rdma_bytes, ) - all_ranks = list(range(world_size)) - state.buffer.connect_ranks(all_ranks) - + initial_ep_size = offset + world_size + scale_to = max(initial_ep_size, state.scale_to or 0) + cls._connect_ranks(state, list(range(scale_to)), tag="initial") + state.connected_ep_size = scale_to + state.scale_to = scale_to + state.dispatch_ep_size = scale_to return state.buffer @classmethod @@ -170,8 +225,12 @@ class _NixlEPDispatcherImplBase: self.active_ranks = ( elastic_state.active_ranks if elastic_state is not None else None ) + self._active_world_size = dist.get_world_size(group) + from sglang.srt.runtime_context import get_server_args + + _max_ep = get_server_args().max_ep_size or self._active_world_size self._mask_buffer = ( - torch.zeros_like(self.active_ranks) + torch.zeros(_max_ep, dtype=torch.int32, device="cuda") if self.active_ranks is not None else None ) @@ -232,14 +291,21 @@ class _NixlEPDispatcherImpl(_NixlEPDispatcherImplBase): buffer = self._get_buffer() topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids topk_ids = topk_ids.to(torch.int64) + state = NixlEPBuffer._state() + dispatch_ep_size = state.dispatch_ep_size + num_local_experts = state.num_local_experts + assert dispatch_ep_size is not None and num_local_experts is not None + num_dispatch_experts = num_local_experts * dispatch_ep_size expected_m = ( - hidden_states.shape[0] * buffer.group_size * topk_ids.shape[1] - + self.num_experts - ) // self.num_experts + hidden_states.shape[0] * dispatch_ep_size * topk_ids.shape[1] + + num_dispatch_experts + ) // num_dispatch_experts + hidden_states, masked_m, event, hook = self._dispatch_core( hidden_states, topk_ids, ) + return ( hidden_states, topk_ids, @@ -289,12 +355,17 @@ class _NixlEPDispatcherImpl(_NixlEPDispatcherImplBase): use_fp8 = not envs.SGLANG_NIXL_EP_BF16_DISPATCH.get() buffer = self._get_buffer() + state = NixlEPBuffer._state() + dispatch_ep_size = state.dispatch_ep_size + num_local_experts = state.num_local_experts + assert dispatch_ep_size is not None and num_local_experts is not None + nixl_num_experts = num_local_experts * dispatch_ep_size packed_recv_hidden, self.packed_recv_count, self.handle, event, hook = ( buffer.dispatch( hidden_states, topk_idx, self.num_max_dispatch_tokens_per_rank, - self.num_experts, + nixl_num_experts, use_fp8=use_fp8, async_finish=not self.return_recv_hook, return_recv_hook=self.return_recv_hook, @@ -341,7 +412,9 @@ class _NixlEPDispatcherImpl(_NixlEPDispatcherImplBase): ) if self._mask_buffer is not None: buffer.query_mask_buffer(self._mask_buffer) - self.active_ranks.copy_(1 - self._mask_buffer) + + n = ElasticEPStateManager.get_effective_ep_size() + self.active_ranks[:n].copy_(1 - self._mask_buffer[:n]) self.packed_recv_count = self.handle = None return combined_hidden_states, event, hook diff --git a/python/sglang/srt/managers/communicator.py b/python/sglang/srt/managers/communicator.py index 4e255c6e2..e56ed66d9 100644 --- a/python/sglang/srt/managers/communicator.py +++ b/python/sglang/srt/managers/communicator.py @@ -2,8 +2,11 @@ from __future__ import annotations import asyncio import copy +import logging from typing import Callable, Generic, List, Optional, TypeVar +logger = logging.getLogger(__name__) + T = TypeVar("T") @@ -30,6 +33,7 @@ class FanOutCommunicator(Generic[T]): self._mode = mode self._result_event: Optional[asyncio.Event] = None self._result_values: Optional[List[T]] = None + self._result_fan_out: Optional[int] = None self._queueing_lock = asyncio.Lock() assert mode in ["queueing", "watching"] @@ -45,10 +49,11 @@ class FanOutCommunicator(Generic[T]): self._result_event = asyncio.Event() self._result_values = [] + self._result_fan_out = self._fan_out await self._result_event.wait() result_values = self._result_values self._result_event = self._result_values = None - + self._result_fan_out = None return result_values async def watching_call(self, obj): @@ -56,6 +61,7 @@ class FanOutCommunicator(Generic[T]): assert self._result_values is None self._result_values = [] self._result_event = asyncio.Event() + self._result_fan_out = self._fan_out if obj is not None: self._send(obj) @@ -69,6 +75,7 @@ class FanOutCommunicator(Generic[T]): result_values = copy.deepcopy(values) if self._result_event is event: self._result_event = self._result_values = None + self._result_fan_out = None return result_values async def __call__(self, obj): @@ -77,9 +84,22 @@ class FanOutCommunicator(Generic[T]): else: return await self.watching_call(obj) + def set_fan_out(self, fan_out: int): + self._fan_out = fan_out + def handle_recv(self, recv_obj: T): + if ( + self._result_values is None + or self._result_event is None + or self._result_fan_out is None + ): + logger.debug( + "Dropping communicator response without active waiter: %s", + type(recv_obj).__name__, + ) + return self._result_values.append(recv_obj) - if len(self._result_values) == self._fan_out: + if len(self._result_values) == self._result_fan_out: self._result_event.set() @staticmethod diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 2353ac51f..c43a4fc1b 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -33,6 +33,7 @@ from sglang.srt.managers.io_struct import ( BatchTokenizedEmbeddingReqInput, BatchTokenizedGenerateReqInput, BlockReqInput, + ElasticScaleUpdateReq, ProfileReq, TokenizedEmbeddingReqInput, TokenizedGenerateReqInput, @@ -164,7 +165,17 @@ class DataParallelController: LoadBalanceMethod.TOTAL_TOKENS, ) - # Load balance budget + self.launch_dp_size: int = server_args.dp_size + self.max_dp_size: int = server_args.max_ep_size or server_args.dp_size + assert self.max_dp_size >= self.launch_dp_size, ( + f"--max-ep-size ({self.max_dp_size}) must be >= " + f"--dp ({self.launch_dp_size})." + ) + + self.dp_active: List[bool] = [True] * self.launch_dp_size + [False] * ( + self.max_dp_size - self.launch_dp_size + ) + self.dp_budget = DPBudget(server_args.dp_size) self.load_snapshot_reader = create_load_snapshot_reader( server_args, @@ -178,8 +189,10 @@ class DataParallelController: # Launch data parallel workers self.scheduler_procs = [] - self.workers: List[zmq.Socket] = [None] * server_args.dp_size - self.status: List[bool] = [True] * server_args.dp_size + self.workers: List[Optional[zmq.Socket]] = [None] * self.max_dp_size + self.status: List[bool] = list(self.dp_active) + self._active_workers: List[int] = list(range(self.launch_dp_size)) + self._active_count_cache: int = self.launch_dp_size if server_args.enable_dp_attention: self.launch_dp_attention_schedulers(server_args, port_args) @@ -208,16 +221,80 @@ class DataParallelController: def send_to_all_workers(self, obj): for i, worker in enumerate(self.workers): - if self.status[i]: + if worker is not None and self.status[i]: sock_send(worker, obj) def send_control_message(self, obj): - # Send control messages to first worker of tp group - for worker in self.workers[:: self.control_message_step]: - sock_send(worker, obj) + for i in self._active_workers[:: self.control_message_step]: + worker = self.workers[i] + if worker is not None: + sock_send(worker, obj) def update_active_ranks(self, ranks: ActiveRanksOutput): - self.status = ranks.status + if self.server_args.elastic_ep_backend is not None: + if len(ranks.status) != self.max_dp_size: + logger.warning( + "[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; " + "ignoring update", + len(ranks.status), + self.max_dp_size, + ) + return + self.status = [ + self.dp_active[i] and bool(ranks.status[i]) + for i in range(self.max_dp_size) + ] + self._refresh_active_workers() + return + if len(ranks.status) != self.max_dp_size: + logger.warning( + "[DPC] update_active_ranks: status len=%d != max_dp_size=%d; " + "ignoring update", + len(ranks.status), + self.max_dp_size, + ) + return + self.status = list(ranks.status) + + def add_elastic_workers(self, slot_offset: int, slot_count: int): + """Activate a range of pre-bound worker slots.""" + end = slot_offset + slot_count + if end > self.max_dp_size: + raise ValueError( + f"[Elastic EP] add_elastic_workers: slot_offset={slot_offset} + " + f"slot_count={slot_count} exceeds max_dp_size={self.max_dp_size}. " + f"Restart with a larger --max-ep-size." + ) + + for slot in range(slot_offset, end): + if self.dp_active[slot]: + logger.debug( + "[Elastic EP] add_elastic_workers: slot %d already active; " + "skipping", + slot, + ) + continue + assert self.workers[slot] is not None, ( + f"[Elastic EP] add_elastic_workers: slot {slot} was not " + f"pre-bound at launch; expected a primary-bound PUSH socket." + ) + self.dp_active[slot] = True + self.status[slot] = True + + self._refresh_active_workers() + logger.debug( + "[Elastic EP] DataParallelController activated slots %s " + "(active=%d / max=%d)", + list(range(slot_offset, end)), + self._active_count_cache, + self.max_dp_size, + ) + + def _refresh_active_workers(self) -> None: + self._active_workers = [ + i for i, active in enumerate(self.dp_active) if active and self.status[i] + ] + self._active_count_cache = len(self._active_workers) def refresh_load_budget(self): # Throttle to at most once per 20ms. When a burst of requests @@ -272,6 +349,12 @@ class DataParallelController: (BlockReqInput, self.send_to_all_workers), (ProfileReq, self.send_to_all_workers), (ActiveRanksOutput, self.update_active_ranks), + ( + ElasticScaleUpdateReq, + lambda msg: self.add_elastic_workers( + msg.slot_offset, msg.slot_count + ), + ), ] ) self._request_dispatcher.add_fallback_fn(self.send_control_message) @@ -355,8 +438,8 @@ class DataParallelController: Returns: List of worker ports (same on all nodes after broadcast). """ - # Determine the endpoint for inter-node communication - if server_args.dist_init_addr is None: + is_joiner = server_args.is_ep_scale_joiner + if server_args.dist_init_addr is None or is_joiner: na = NetworkAddress( server_args.host or "127.0.0.1", server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA, @@ -411,11 +494,10 @@ class DataParallelController: ).start() def _reply_ports_as_server(self, rep_socket: zmq.Socket, worker_ports: List[int]): - """ - Runs as a background thread to broadcast worker ports for recovered EP ranks - """ + """Background thread: serve the pre-bound worker-port list to + late-arriving elastic joiners. Publishes port numbers only; the primary + keeps ownership of every socket.""" while True: - # Wait for client handshake try: client_rank = sock_recv(rep_socket) except Exception: @@ -453,6 +535,12 @@ class DataParallelController: finally: req_socket.close() + def _joiner_local_tp_span(self, server_args: ServerArgs) -> int: + return server_args.tp_size + + def _joiner_slot_offset(self, server_args: ServerArgs) -> int: + return server_args.ep_join_rank_offset + def launch_dp_attention_schedulers( self, server_args: ServerArgs, port_args: PortArgs ): @@ -461,25 +549,42 @@ class DataParallelController: else: bind_host = NetworkAddress.parse(server_args.dist_init_addr).host - # Pre-allocate worker ports on node 0 to avoid conflicts worker_ports = [] - if server_args.node_rank == 0: - for dp_rank in range(server_args.dp_size): + if server_args.is_ep_scale_joiner: + # Scale joiners connect to their pre-bound primary worker sockets. + primary = NetworkAddress.parse(server_args.dist_init_addr) + primary_endpoint = NetworkAddress( + primary.host, primary.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA + ).to_tcp() + all_ports = self._receive_ports_as_client( + primary_endpoint, server_args.node_rank + ) + offset = self._joiner_slot_offset(server_args) + local_tp_span = self._joiner_local_tp_span(server_args) + broadcasted_ports = all_ports[offset : offset + local_tp_span] + elif server_args.node_rank == 0: + # Elastic primaries reserve sockets for the maximum DP size. + bind_count = ( + self.max_dp_size + if server_args.elastic_ep_backend is not None + else server_args.dp_size + ) + for slot in range(bind_count): worker_port, worker_socket = get_zmq_socket_on_host( self.context, zmq.PUSH, host=bind_host ) worker_ports.append(worker_port) - self.workers[dp_rank] = worker_socket + self.workers[slot] = worker_socket logger.debug( - "Assigned port %s to worker %s on host %s", + "Assigned port %s to worker slot %s on host %s", worker_port, - dp_rank, + slot, bind_host, ) + broadcasted_ports = self._broadcast_worker_ports(server_args, worker_ports) + else: + broadcasted_ports = self._broadcast_worker_ports(server_args, None) - broadcasted_ports = self._broadcast_worker_ports( - server_args, worker_ports if worker_ports else None - ) self.launch_tensor_parallel_group( server_args, port_args, 0, None, broadcasted_ports ) @@ -510,10 +615,15 @@ class DataParallelController: nnodes_per_tp_group = nnodes_per_pp_rank tp_size_per_node = server_args.tp_size // nnodes_per_tp_group - tp_rank_range = range( - tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group), - tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1), - ) + if server_args.is_ep_scale_joiner: + # Scale joiners enumerate their full local TP span. + tp_rank_range = range(server_args.tp_size) + tp_size_per_node = server_args.tp_size + else: + tp_rank_range = range( + tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group), + tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1), + ) attn_cp_rank = 0 moe_dp_rank = 0 @@ -534,6 +644,16 @@ class DataParallelController: rank_port_args = PortArgs.init_new( server_args, dp_rank, worker_ports ) + if server_args.is_ep_scale_joiner: + # Scale-joiner outputs return through the primary tokenizer. + primary_addr = NetworkAddress.parse(server_args.dist_init_addr) + primary_port_base = primary_addr.port + 1 + rank_port_args.tokenizer_ipc_name = NetworkAddress( + primary_addr.host, primary_port_base + ).to_tcp() + rank_port_args.detokenizer_ipc_name = NetworkAddress( + primary_addr.host, primary_port_base + 1 + ).to_tcp() # Data parallelism reuses the tensor parallelism group, # so all dp ranks should use the same nccl port. rank_port_args.nccl_port = port_args.nccl_port @@ -570,6 +690,12 @@ class DataParallelController: ) ) + # Scheduler internals use local ranks; logs use global ranks. + offset = server_args.ep_join_rank_offset + display_tp_rank = tp_rank + offset + display_moe_ep_rank = moe_ep_rank + offset + display_dp_rank = dp_rank + offset if dp_rank is not None else None + with self.env_lock, maybe_reindex_device_id(gpu_id) as gpu_id: proc = mp.Process( target=self.run_scheduler_process_func, @@ -584,6 +710,9 @@ class DataParallelController: pp_rank, dp_rank, writer, + display_tp_rank, + display_dp_rank, + display_moe_ep_rank, ), ) with ( @@ -604,8 +733,16 @@ class DataParallelController: def maybe_external_dp_rank_routing(self, req: Req): if req.routed_dp_rank is not None: - logger.debug(f"Direct routing to DP rank {req.routed_dp_rank}") - sock_send(self.workers[req.routed_dp_rank], req) + rank = req.routed_dp_rank + if ( + rank < 0 + or rank >= len(self.workers) + or rank not in self._active_workers + or self.workers[rank] is None + ): + raise ValueError(f"DP rank {rank} is not active.") + logger.debug(f"Direct routing to DP rank {rank}") + sock_send(self.workers[rank], req) return True return False @@ -613,17 +750,22 @@ class DataParallelController: if self.maybe_external_dp_rank_routing(req): return - while True: - if self.status[self.round_robin_counter]: - logger.debug(f"Choose worker {self.round_robin_counter}") - sock_send(self.workers[self.round_robin_counter], req) - self.round_robin_counter = (self.round_robin_counter + 1) % len( - self.workers - ) - break - self.round_robin_counter = (self.round_robin_counter + 1) % len( - self.workers - ) + active = self._active_workers + if not active: + raise RuntimeError("No active DP workers are available for routing.") + attempts = 0 + while attempts < len(active): + slot = active[self.round_robin_counter % len(active)] + self.round_robin_counter = (self.round_robin_counter + 1) % len(active) + if self.status[slot]: + logger.debug(f"Choose worker {slot}") + sock_send(self.workers[slot], req) + return + attempts += 1 + raise RuntimeError( + f"Cannot route request: all {len(active)} active DP workers " + "are unavailable." + ) def follow_bootstrap_room_scheduler(self, req: Req): if self.maybe_external_dp_rank_routing(req): @@ -702,7 +844,8 @@ def run_data_parallel_controller_process( SCHEDULER_PIDS_ARG: scheduler_pids, } ) - if server_args.node_rank == 0: + # The primary owns routing for the expanded scheduler set. + if server_args.node_rank == 0 and not server_args.is_ep_scale_joiner: controller.event_loop() for proc in controller.scheduler_procs: proc.join() diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 5029c8b13..6f900c7e0 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1809,6 +1809,31 @@ class ActiveRanksOutput(BaseReq, kw_only=True): status: List[bool] +class ElasticScaleUpdateReq(BaseReq, kw_only=True): + """Report asynchronous Elastic EP scale completion or failure.""" + + success: bool + effective_ep_size: int + slot_offset: int = 0 + slot_count: int = 0 + error: Optional[str] = None + + +class ScaleElasticEPReqInput(BaseReq, kw_only=True): + """Request to scale EP by changing the effective EP size (dp_attention mode).""" + + new_ep_size: int + + +class ScaleElasticEPReqOutput(BaseReq, kw_only=True): + success: bool + message: str + old_ep_size: int = 0 + new_ep_size: int = 0 + pending_ep_size: Optional[int] = None + scale_phase: str = "idle" + + class GetInternalStateReq(BaseReq, kw_only=True): pass diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 02581abe0..8d025d71e 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -127,6 +127,8 @@ from sglang.srt.managers.io_struct import ( ResumeMemoryOccupationReqInput, RpcReqInput, RpcReqOutput, + ScaleElasticEPReqInput, + ScaleElasticEPReqOutput, SendWeightsToRemoteInstanceReqInput, SendWeightsToRemoteInstanceReqOutput, SetInternalStateReq, @@ -973,6 +975,7 @@ class Scheduler( is_fully_idle=self.is_fully_idle, ipc_channels=self.ipc_channels, ) + self._last_logged_elastic_radix_namespace: Optional[str] = None self.session_controller = SessionController(self.tree_cache) self.forward_sleep_time = None self._engine_paused = False @@ -1408,6 +1411,7 @@ class Scheduler( (PauseGenerationReqInput, self.pause_generation), (ContinueGenerationReqInput, self.continue_generation), (ConfigureLoggingReq, self.configure_logging), + (ScaleElasticEPReqInput, self.handle_scale_elastic_ep), (DumperControlReqInput, self.handle_dumper_control), (AddExternalCorpusReqInput, self.add_external_corpus), ( @@ -1994,6 +1998,33 @@ class Scheduler( mm.mrope_positions = mrope_positions mm.mrope_position_delta = mrope_position_delta + def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None: + if ( + self.server_args.elastic_ep_backend is None + or self.disable_radix_cache + or not self.tree_cache.is_tree_cache() + ): + return + + from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager + + inst = ElasticEPStateManager.instance() + if inst is None: + return + + namespace = f"elastic_ep_size={ElasticEPStateManager.get_effective_ep_size()}" + if req.extra_key: + req.extra_key = f"{req.extra_key}|{namespace}" + else: + req.extra_key = namespace + + if self._last_logged_elastic_radix_namespace != namespace: + self._last_logged_elastic_radix_namespace = namespace + logger.debug( + "[Elastic EP][scale] radix cache namespace is now %s", + namespace, + ) + def _maybe_clear_mm_inputs(self, batch: ScheduleBatch) -> None: for req in batch.reqs: if not req.finished() or not (mm_inputs := req.multimodal_inputs): @@ -2134,6 +2165,8 @@ class Scheduler( self._add_request_to_queue(req) return + self._maybe_namespace_elastic_radix_cache(req) + if self.spec_algorithm.is_dflash_family(): error_msg = validate_dflash_request(req, self.enable_overlap) if error_msg is not None: @@ -2471,6 +2504,7 @@ class Scheduler( multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices, ) req.tokenizer = self.tokenizer + self._maybe_namespace_elastic_radix_cache(req) # Handle multimodal inputs if recv_req.mm_inputs is not None: @@ -3427,14 +3461,24 @@ class Scheduler( self.enable_dp_attention and self.server_args.elastic_ep_backend is not None ): return - # Get the tensors indicating rank activeness - tp_active_ranks = self.tp_group.active_ranks.detach().cpu().numpy() - tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy() - tp_active_ranks &= tp_active_ranks_cpu - dp_active_ranks = tp_active_ranks.reshape(self.ps.dp_size, -1).prod(axis=1) - self.ipc_channels.send_to_tokenizer.send_output( - ActiveRanksOutput(status=dp_active_ranks.tolist()) - ) + from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager + + inst = ElasticEPStateManager.instance() + if inst is not None and inst.active_ranks_cpu is not None: + self.ipc_channels.send_to_tokenizer.send_output( + ActiveRanksOutput( + status=[bool(x) for x in inst.active_ranks_cpu.tolist()] + ) + ) + else: + logger.debug("[Elastic EP] active rank state is unavailable") + return + + model_runner = self.tp_worker.model_runner + pending = model_runner._pending_elastic_scale_update + if pending is not None: + self.ipc_channels.send_to_tokenizer.send_output(pending) + model_runner._pending_elastic_scale_update = None def _relay_forward_payload( self, future_indices: torch.Tensor, batch_result: GenerationBatchResult @@ -3808,6 +3852,15 @@ class Scheduler( } ret["effective_max_running_requests_per_dp"] = self.max_running_requests + if self.server_args.elastic_ep_backend is not None: + from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager + + ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling() + ret["effective_ep_size"] = ElasticEPStateManager.get_effective_ep_size() + ret["pending_ep_size"] = ElasticEPStateManager.get_pending_ep_size() + ret["scale_phase"] = ElasticEPStateManager.get_scale_phase() + ret["elastic_ep_last_error"] = ElasticEPStateManager.get_last_error() + if ( not self.spec_algorithm.is_none() and self.metrics_reporter.spec_total_num_forward_ct > 0 @@ -4170,6 +4223,84 @@ class Scheduler( self.disagg_decode_prealloc_queue.enqueue_held_rebootstrap() self._engine_paused = False + def handle_scale_elastic_ep( + self, recv_req: ScaleElasticEPReqInput + ) -> ScaleElasticEPReqOutput: + """Begin a pending elastic EP scale-up request.""" + from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager + + old_ep_size = ElasticEPStateManager.get_effective_ep_size() + new_ep_size = recv_req.new_ep_size + max_ep_size = self.server_args.max_ep_size or old_ep_size + + logger.debug( + "[Elastic EP][scale] request received: new_ep_size=%d " + "old_ep_size=%d max_ep_size=%d", + new_ep_size, + old_ep_size, + max_ep_size, + ) + + if new_ep_size <= old_ep_size: + return ScaleElasticEPReqOutput( + success=False, + message=( + f"new_ep_size ({new_ep_size}) must be greater than current " + f"effective_ep_size ({old_ep_size})." + ), + old_ep_size=old_ep_size, + new_ep_size=new_ep_size, + ) + if new_ep_size > max_ep_size: + return ScaleElasticEPReqOutput( + success=False, + message=( + f"new_ep_size ({new_ep_size}) exceeds --max-ep-size " + f"({max_ep_size}). Restart with a larger --max-ep-size." + ), + old_ep_size=old_ep_size, + new_ep_size=new_ep_size, + ) + if ElasticEPStateManager.is_scaling(): + return ScaleElasticEPReqOutput( + success=False, + message=( + "A previous scale operation has not completed yet. Wait until " + "all pending ranks have joined before issuing another scale." + ), + old_ep_size=old_ep_size, + new_ep_size=new_ep_size, + pending_ep_size=ElasticEPStateManager.get_pending_ep_size(), + scale_phase=ElasticEPStateManager.get_scale_phase(), + ) + + if not ElasticEPStateManager.request_scale(new_ep_size): + return ScaleElasticEPReqOutput( + success=False, + message=( + "Failed to queue elastic EP scale: no elastic state or " + "scale already pending." + ), + old_ep_size=old_ep_size, + new_ep_size=new_ep_size, + pending_ep_size=ElasticEPStateManager.get_pending_ep_size(), + scale_phase=ElasticEPStateManager.get_scale_phase(), + ) + logger.debug( + "[Elastic EP][scale] scale requested: target_ep_size=%d; " + "waiting for a joining cohort", + new_ep_size, + ) + + return ScaleElasticEPReqOutput( + success=True, + message=f"Scaling initiated from {old_ep_size} to {new_ep_size}", + old_ep_size=old_ep_size, + new_ep_size=new_ep_size, + pending_ep_size=ElasticEPStateManager.get_pending_ep_size(), + scale_phase=ElasticEPStateManager.get_scale_phase(), + ) + def load_lora_adapter( self, recv_req: LoadLoRAAdapterReqInput ) -> LoadLoRAAdapterReqOutput: @@ -4334,11 +4465,13 @@ def configure_scheduler_process( moe_ep_rank: int, pp_rank: int, dp_rank: Optional[int], + display_tp_rank: Optional[int] = None, + display_dp_rank: Optional[int] = None, + display_moe_ep_rank: Optional[int] = None, ) -> Optional[int]: - """Configure scheduler worker: logging, process title, etc. + """Configure scheduler worker logging and process title. - Returns: - dp_rank + display_* ranks are cosmetic; runtime ranks stay local. """ kill_itself_when_parent_died() @@ -4347,9 +4480,15 @@ def configure_scheduler_process( # [For Router] if env var "SGLANG_DP_RANK" exist, set dp_rank to the value of the env var dp_rank = int(os.environ["SGLANG_DP_RANK"]) + shown_dp = display_dp_rank if display_dp_rank is not None else dp_rank + shown_tp = display_tp_rank if display_tp_rank is not None else tp_rank + shown_moe_ep = ( + display_moe_ep_rank if display_moe_ep_rank is not None else moe_ep_rank + ) + prefix = "" - if dp_rank is not None: - prefix += f" DP{dp_rank}" + if shown_dp is not None: + prefix += f" DP{shown_dp}" if server_args.pp_size > 1: prefix += f" PP{pp_rank}" if server_args.attn_cp_size > 1: @@ -4357,9 +4496,9 @@ def configure_scheduler_process( if server_args.moe_dp_size > 1: prefix += f" MOE_DP{moe_dp_rank}" if server_args.tp_size > 1: - prefix += f" TP{tp_rank}" + prefix += f" TP{shown_tp}" if server_args.ep_size > 1: - prefix += f" EP{moe_ep_rank}" + prefix += f" EP{shown_moe_ep}" # Config the process setproctitle.setproctitle(f"sglang::scheduler{prefix.replace(' ', '_')}") @@ -4393,6 +4532,9 @@ def run_scheduler_process( pp_rank: int, dp_rank: Optional[int], pipe_writer, + display_tp_rank: Optional[int] = None, + display_dp_rank: Optional[int] = None, + display_moe_ep_rank: Optional[int] = None, ): # Load plugins so hooks can override Scheduler and its dependencies. load_plugins() @@ -4405,6 +4547,9 @@ def run_scheduler_process( moe_ep_rank, pp_rank, dp_rank, + display_tp_rank=display_tp_rank, + display_dp_rank=display_dp_rank, + display_moe_ep_rank=display_moe_ep_rank, ) parent_process = psutil.Process().parent() diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index 0851f7605..7219aaf9b 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -10,6 +10,7 @@ from sglang.srt.configs.model_config import ModelConfig from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs +from sglang.srt.layers.dp_attention import world_dp_gather_enabled from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler_components.recv_skipper import ( SchedulerRecvSkipper, @@ -36,6 +37,43 @@ if TYPE_CHECKING: _ENABLE_METRICS_DP_ATTENTION = envs.SGLANG_ENABLE_METRICS_DP_ATTENTION.get() +def _resolve_elastic_world_dp_size( + dp_size: int, + *, + group: torch.distributed.ProcessGroup, + local_num_tokens: int, + local_forward_mode: int, +) -> int: + if not world_dp_gather_enabled(): + return dp_size + + from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager + from sglang.srt.layers.dp_attention import get_attention_dp_size + + live_dp_size = get_attention_dp_size() + effective_ep_size = ElasticEPStateManager.get_effective_ep_size() + world_size = torch.distributed.get_world_size(group) + + if live_dp_size != effective_ep_size: + raise RuntimeError( + "[Elastic EP] WORLD MLP sync dp_size is out of sync: " + f"rank={torch.distributed.get_rank(group)} " + f"live_dp_size={live_dp_size} effective_ep_size={effective_ep_size} " + f"world_size={world_size} server_args_dp_size={dp_size} " + f"local_num_tokens={local_num_tokens} " + f"local_forward_mode={local_forward_mode}" + ) + if live_dp_size > world_size: + raise RuntimeError( + "[Elastic EP] WORLD MLP sync dp_size exceeds WORLD size: " + f"rank={torch.distributed.get_rank(group)} " + f"live_dp_size={live_dp_size} world_size={world_size} " + f"effective_ep_size={effective_ep_size}" + ) + + return live_dp_size + + @dataclass class MLPSyncBatchInfo: dp_size: int @@ -88,27 +126,56 @@ class MLPSyncBatchInfo: dtype=dtype, ) - def all_gather(self, device, group: torch.distributed.ProcessGroup): + def all_gather( + self, + device, + group: torch.distributed.ProcessGroup, + use_all_reduce: bool = False, + ): local_info_tensor = self._get_local_tensor(device=device) - global_info_tensor = torch.empty( - (self.dp_size, self.tp_size * self.cp_size, 7), - dtype=torch.int64, - device=device, - ) + fallback_tensor = self._get_fallback_tensor(device=device) + info_width = local_info_tensor.numel() + # Inactive max_world_size slots must decode as IDLE. + global_info_tensor = fallback_tensor.expand( + self.dp_size, self.tp_size * self.cp_size, info_width + ).contiguous() - torch.distributed.all_gather_into_tensor( - global_info_tensor.flatten(), - local_info_tensor, - group=group, + if use_all_reduce: + # Admission can expose different WORLD sizes; use fixed global slots. + global_info_tensor.zero_() + flat_info = global_info_tensor.view(-1, info_width) + rank = torch.distributed.get_rank(group) + if 0 <= rank < flat_info.shape[0]: + flat_info[rank] = local_info_tensor + torch.distributed.all_reduce( + global_info_tensor, + op=torch.distributed.ReduceOp.SUM, + group=group, + ) + missing = flat_info.abs().sum(dim=1) == 0 + flat_info[missing] = fallback_tensor + else: + torch.distributed.all_gather_into_tensor( + global_info_tensor.flatten(), + local_info_tensor, + group=group, + ) + + tp_info = global_info_tensor.view( + self.dp_size * self.tp_size * self.cp_size, info_width ) + num_ranks_in_tp_info = tp_info.shape[0] if device == "cpu": tp_active_ranks = get_tp_group().active_ranks_cpu else: tp_active_ranks = get_tp_group().active_ranks - - # Set fallback values for inactive ranks - tp_info = global_info_tensor.view(self.dp_size * self.tp_size * self.cp_size, 7) - tp_info[tp_active_ranks == 0] = self._get_fallback_tensor(device=device) + if tp_active_ranks.shape[0] < num_ranks_in_tp_info: + tp_active_ranks = torch.ones( + num_ranks_in_tp_info, + dtype=tp_active_ranks.dtype, + device=tp_active_ranks.device, + ) + tp_info[tp_active_ranks[:num_ranks_in_tp_info] == 0] = fallback_tensor tp0_info = global_info_tensor[:, 0, :] self.tp0_info = tp0_info @@ -206,7 +273,14 @@ def prepare_mlp_sync_batch_raw( local_batch.is_extend_in_batch = is_extend_in_batch tbo_preparer = TboDPAttentionPreparer() - if len(offload_tags) == 0 and ( + use_world_group = world_dp_gather_enabled() + if use_world_group: + from sglang.srt.distributed.parallel_state import get_world_group + + world = get_world_group() + group = torch.distributed.group.WORLD + device = world.device + elif len(offload_tags) == 0 and ( disable_overlap_schedule or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get() ): @@ -217,6 +291,13 @@ def prepare_mlp_sync_batch_raw( device = "cpu" local_can_run_tbo, local_forward_mode = tbo_preparer.prepare_all_gather(local_batch) + if use_world_group: + dp_size = _resolve_elastic_world_dp_size( + dp_size, + group=group, + local_num_tokens=num_tokens, + local_forward_mode=local_forward_mode, + ) mlp_sync_info = MLPSyncBatchInfo( dp_size=dp_size, @@ -232,7 +313,11 @@ def prepare_mlp_sync_batch_raw( ) if not skip_all_gather: - mlp_sync_info.all_gather(device=device, group=group) + mlp_sync_info.all_gather( + device=device, + group=group, + use_all_reduce=use_world_group, + ) mlp_sync_info.tbo_split_seq_index, mlp_sync_info.global_forward_mode = ( tbo_preparer.compute_output( diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index a27664adc..fccd74c1d 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -167,7 +167,10 @@ class SchedulerRequestReceiver: # controller, so we broadcast within attn_tp_group + attn_cp_group # instead of the full tp_group. This avoids an expensive # all-ranks gloo sync. - _local_ctrl = self.server_args.enable_dp_attention_local_control_broadcast + _local_ctrl = ( + self.server_args.enable_dp_attention_local_control_broadcast + or self.server_args.is_ep_scale_joiner + ) if _local_ctrl: if self.ps.attn_tp_size != 1: control_reqs = broadcast_pyobj( diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 9cbec3725..86b7b378f 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -57,6 +57,7 @@ from sglang.srt.managers.io_struct import ( RemoveExternalCorpusReqOutput, ResumeMemoryOccupationReqInput, ResumeMemoryOccupationReqOutput, + ScaleElasticEPReqOutput, SendWeightsToRemoteInstanceReqInput, SendWeightsToRemoteInstanceReqOutput, SetInternalStateReq, @@ -118,6 +119,7 @@ _COMMUNICATOR_SPECS = [ ("expert_distribution", ExpertDistributionReqOutput), ("update_lora_adapter", LoRAUpdateOutput), ("dumper_control", DumperControlReqOutput), + ("scale_elastic_ep", ScaleElasticEPReqOutput), ] @@ -141,6 +143,23 @@ class TokenizerControlMixin: dispatch_pairs.append((resp_type, comm.handle_recv)) self._result_dispatcher += TypeBasedDispatcher(dispatch_pairs) + def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int): + primary_group_control = ( + self.server_args.enable_dp_attention + and not self.server_args.enable_dp_attention_local_control_broadcast + ) + if primary_group_control: + control_fan_out = ( + worker_count + self.server_args.tp_size - 1 + ) // self.server_args.tp_size + else: + control_fan_out = worker_count + + for spec in _COMMUNICATOR_SPECS: + getattr(self, f"{spec[0]}_communicator").set_fan_out(worker_count) + + self.get_internal_state_communicator.set_fan_out(control_fan_out) + async def add_external_corpus( self: TokenizerManager, obj: AddExternalCorpusReqInput ) -> AddExternalCorpusReqOutput: @@ -821,7 +840,9 @@ class TokenizerControlMixin: List of LoadSnapshot, one per scheduler (filtered by dp_rank if specified) """ self.auto_create_handle_loop() - if dp_rank is not None and (dp_rank < 0 or dp_rank >= self.server_args.dp_size): + if dp_rank is not None and ( + dp_rank < 0 or dp_rank >= self.elastic_worker_count + ): return [] reader = self.load_snapshot_reader diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index f96660fb0..60b35ea02 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -65,6 +65,7 @@ from sglang.srt.managers.io_struct import ( BatchTokenizedGenerateReqInput, ConfigureLoggingReq, ContinueGenerationReqInput, + ElasticScaleUpdateReq, EmbeddingReqInput, FreezeGCReq, GenerateReqInput, @@ -72,6 +73,8 @@ from sglang.srt.managers.io_struct import ( LoadLoRAAdapterReqInput, OpenSessionReqOutput, PauseGenerationReqInput, + ScaleElasticEPReqInput, + ScaleElasticEPReqOutput, SessionParams, ShutdownReq, TokenizedEmbeddingReqInput, @@ -279,6 +282,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ): # Parse args self.server_args = server_args + self.elastic_worker_count = server_args.dp_size + self.elastic_pending_ep_size = None + self.elastic_scale_phase = "idle" + self.elastic_last_error = None self.enable_metrics = server_args.enable_metrics self.incremental_streaming_output = server_args.incremental_streaming_output self.enable_lora = server_args.enable_lora @@ -491,6 +498,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self.model_update_result: Optional[Awaitable[UpdateWeightFromDiskReqOutput]] = ( None ) + self.model_update_expected_workers = self.elastic_worker_count + self.model_update_tmp: List[UpdateWeightFromDiskReqOutput] = [] self.is_pause = False self.is_pause_cond = asyncio.Condition() @@ -604,6 +613,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Same skip-detokenizer forwarding case as above. (ConfigureLoggingReq, lambda x: None), (ActiveRanksOutput, self.update_active_ranks), + (ElasticScaleUpdateReq, self.forward_elastic_scale_update), ] ) self.init_communicators(self.server_args) @@ -623,7 +633,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self._set_default_priority(obj) if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None: - dp_size = self.server_args.dp_size + dp_size = self.elastic_worker_count if dp_size <= 1 and obj.routed_dp_rank == 0: logger.debug( f"routed_dp_rank={obj.routed_dp_rank} is ignored because dp_size={dp_size}" @@ -1781,15 +1791,17 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): async def _wait_for_model_update_from_disk( self, obj: UpdateWeightFromDiskReqInput ) -> Tuple[bool, str]: - self._dispatch_to_scheduler(obj) + expected_workers = self.elastic_worker_count + self.model_update_expected_workers = expected_workers + self.model_update_tmp = [] self.model_update_result = asyncio.Future() - if self.server_args.dp_size == 1: + self._dispatch_to_scheduler(obj) + if expected_workers == 1: result = await self.model_update_result if result.success: self._update_model_path_info(obj.model_path, obj.load_format) return result.success, result.message, result.num_paused_requests - else: # self.server_args.dp_size > 1 - self.model_update_tmp = [] + else: result = await self.model_update_result all_success = all([r.success for r in result]) @@ -2820,6 +2832,60 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): def update_active_ranks(self, ranks: ActiveRanksOutput): self._dispatch_to_scheduler(ranks) + def forward_elastic_scale_update(self, msg: ElasticScaleUpdateReq): + if not msg.success: + self.elastic_pending_ep_size = None + self.elastic_scale_phase = "failed" + self.elastic_last_error = msg.error + return + + self._dispatch_to_scheduler(msg) + self.elastic_worker_count = msg.effective_ep_size + self.elastic_pending_ep_size = None + self.elastic_scale_phase = "serving_expanded" + self.elastic_last_error = None + self.update_control_communicator_fan_out(msg.effective_ep_size) + + def get_elastic_ep_state(self): + return { + "is_scaling_elastic_ep": self.elastic_pending_ep_size is not None, + "effective_ep_size": self.elastic_worker_count, + "pending_ep_size": self.elastic_pending_ep_size, + "scale_phase": self.elastic_scale_phase, + "last_error": self.elastic_last_error, + } + + async def scale_elastic_ep( + self, obj: ScaleElasticEPReqInput + ) -> ScaleElasticEPReqOutput: + """Send a scale request to every DP scheduler.""" + if self.elastic_pending_ep_size is not None: + return ScaleElasticEPReqOutput( + success=False, + message=( + "A previous scale operation has not completed yet. Wait until " + "all pending ranks have joined before issuing another scale." + ), + old_ep_size=self.elastic_worker_count, + new_ep_size=obj.new_ep_size, + pending_ep_size=self.elastic_pending_ep_size, + scale_phase=self.elastic_scale_phase, + ) + self.auto_create_handle_loop() + responses: List[ScaleElasticEPReqOutput] = ( + await self.scale_elastic_ep_communicator(obj) + ) + for res in responses: + if not res.success: + self.elastic_scale_phase = res.scale_phase + self.elastic_pending_ep_size = res.pending_ep_size + self.elastic_last_error = res.message + return res + self.elastic_pending_ep_size = responses[0].pending_ep_size + self.elastic_scale_phase = responses[0].scale_phase + self.elastic_last_error = None + return responses[0] + def _handle_open_session_req_output(self, recv_obj): future = self.session_futures.get(recv_obj.session_id) if future is None: @@ -2832,12 +2898,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): future.set_result(recv_obj.session_id if recv_obj.success else None) def _handle_update_weights_from_disk_req_output(self, recv_obj): - if self.server_args.dp_size == 1: + if self.model_update_expected_workers == 1: self.model_update_result.set_result(recv_obj) - else: # self.server_args.dp_size > 1 + else: self.model_update_tmp.append(recv_obj) - # set future if the all results are received - if len(self.model_update_tmp) == self.server_args.dp_size: + if len(self.model_update_tmp) == self.model_update_expected_workers: self.model_update_result.set_result(self.model_update_tmp) async def _validate_and_resolve_lora( diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 485a897f8..eee3e4fb1 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -310,13 +310,17 @@ class TpModelWorker(BaseTpWorker): self.pp_group = get_pp_group() self.world_group = get_world_group() - # Sync random seed across TP workers - self.random_seed = broadcast_pyobj( - [server_args.random_seed], - self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank, - self.world_group.cpu_group, - src=self.world_group.ranks[0], - )[0] + # Sync random seed across TP workers. + # Scale joiners cannot enter the launch-time WORLD broadcast. + if server_args.is_ep_scale_joiner: + self.random_seed = server_args.random_seed + else: + self.random_seed = broadcast_pyobj( + [server_args.random_seed], + self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank, + self.world_group.cpu_group, + src=self.world_group.ranks[0], + )[0] set_random_seed(self.random_seed) self.enable_overlap = not server_args.disable_overlap_schedule diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index caf45e74e..85204c0d0 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -46,6 +46,7 @@ from sglang.srt.layers.dp_attention import ( DpPaddingMode, set_dp_buffer_len, set_is_extend_in_batch, + world_dp_gather_enabled, ) from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import ( ForwardBatchDeepSeekMHAMixin, @@ -75,6 +76,25 @@ _skip_attn_backend_init_warned = False _is_npu = is_npu() +def _elastic_should_preserve_local_token_counts( + *, + model_runner: ModelRunner, + dp_padding_mode: DpPaddingMode, + global_num_tokens: List[int], +) -> bool: + if not getattr(model_runner, "enable_elastic_ep", False): + return False + if not world_dp_gather_enabled(): + return False + if not dp_padding_mode.is_max_len(): + return False + if len(global_num_tokens) <= 1: + return False + + uneven_token_count = len(set(global_num_tokens)) > 1 + return uneven_token_count + + class ForwardMode(IntEnum): # Extend a sequence. The KV cache of the beginning part of the sequence is already computed (e.g., system prompt). # It is also called "prefill" in common terminology. @@ -1147,7 +1167,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): assert self.global_num_tokens_for_logprob_cpu is not None self._original_batch_size = self.batch_size - global_num_tokens = self.global_num_tokens_cpu + global_num_tokens = list(self.global_num_tokens_cpu) sync_group_size = len(global_num_tokens) attn_tp_size = get_parallel().attn_tp_size @@ -1168,6 +1188,13 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): dp_padding_mode = DpPaddingMode.get_dp_padding_mode( self.is_extend_in_batch, global_num_tokens ) + if _elastic_should_preserve_local_token_counts( + model_runner=model_runner, + dp_padding_mode=dp_padding_mode, + global_num_tokens=global_num_tokens, + ): + # Joined ranks require real token counts instead of MAX_LEN padding. + dp_padding_mode = DpPaddingMode.SUM_LEN # Prefill breakable CUDA graph requires every DP rank to run the SAME # captured shape. Under SUM_LEN each rank pads to its own local token # count and can select a different capture bucket, so the in-graph DP diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index eae69fd72..c622a33de 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -23,6 +23,7 @@ from dataclasses import dataclass from typing import Optional, Union import torch +import torch.distributed as dist from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.model_config import ( @@ -41,9 +42,13 @@ from sglang.srt.dllm.config import DllmConfig from sglang.srt.elastic_ep.elastic_ep import ( ElasticEPStateManager, get_healthy_expert_location_src_rank, + get_scale_cohort_target, join_process_groups, + join_scale_process_group, maybe_rebalance_after_rank_fault, maybe_recover_ep_ranks, + register_scale_cohort, + try_admit_scale_ranks, ) from sglang.srt.elastic_ep.expert_backup_client import ExpertBackupClient from sglang.srt.environ import envs @@ -55,6 +60,8 @@ from sglang.srt.eplb.expert_distribution import ( set_global_expert_distribution_recorder, ) from sglang.srt.eplb.expert_location import ( + ExpertLocationMetadata, + append_trivial_expert_slots, broadcast_global_expert_location_metadata, compute_initial_expert_location_metadata, format_expert_location_layout, @@ -276,6 +283,7 @@ class ModelRunner: self.attention_chunk_size = model_config.attention_chunk_size self.enable_elastic_ep = server_args.elastic_ep_backend is not None self.forward_pass_id = 0 + self._pending_elastic_scale_update = None self.init_new_workspace = False self.draft_model_idx = draft_model_idx self.enable_hisparse = server_args.enable_hisparse @@ -349,17 +357,7 @@ class ModelRunner: self.initialize() self.check_quantized_moe_compatibility() - if ( - self.server_args.elastic_ep_backend is not None - and self.server_args.elastic_ep_rejoin - ): - join_process_groups() - broadcast_global_expert_location_metadata( - src_rank=get_healthy_expert_location_src_rank( - invoked_in_elastic_ep_rejoin_path=True - ) - ) - ElasticEPStateManager.instance().reset() + self._initialize_elastic_ep_joiner() if self.is_multimodal: sanity_check_mm_pad_shift_value(self.model_config.vocab_size) @@ -378,6 +376,87 @@ class ModelRunner: self.init_weight_updater() self.init_weight_exporter() + 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 + ): + 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 + ) + 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() + + 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 + ) + ), + ) + set_global_expert_distribution_recorder( + ExpertDistributionRecorder.init_new( + self.server_args, + get_global_expert_location_metadata(), + rank=global_ep_rank, + ) + ) + + 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, + ) + + enable_joiner_all_gather() + update_dp_attention_post_scale( + new_dp_size=join_effective_ep_size, + new_dp_rank=global_ep_rank, + ) + self.server_args.override( + "elastic_ep.scale_join", dp_size=join_effective_ep_size + ) + if self.eplb_manager is not None: + self.eplb_manager.disable_rebalance( + "EPLB rebalance is disabled after elastic EP scale-up" + ) + + state = ElasticEPStateManager.instance() + if state is not None: + state.active_ranks.zero_() + state.active_ranks[:join_effective_ep_size] = 1 + state.snapshot_active_to_last() + state.sync_active_to_cpu() + state.scale_phase = "syncing_new_world" + self._elastic_scale_ready_barrier( + target_size=join_effective_ep_size, + log_tag="JOINER", + ) + if state is not None: + state.scale_phase = "serving_expanded" + def init_msprobe(self): self.msprobe_debugger = misc_utils.create_msprobe_debugger(self.server_args) @@ -526,11 +605,16 @@ class ModelRunner: def maybe_init_expert_location_metadata(self): if self.is_draft_worker: return + expert_rank = self.ps.moe_ep_rank + ( + self.server_args.ep_join_rank_offset + if self.server_args.is_ep_scale_joiner + else 0 + ) set_global_expert_location_metadata( compute_initial_expert_location_metadata( server_args=self.server_args, model_config=self.model_config, - moe_ep_rank=self.ps.moe_ep_rank, + moe_ep_rank=expert_rank, ) ) if self.ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get(): @@ -542,7 +626,7 @@ class ModelRunner: ExpertDistributionRecorder.init_new( self.server_args, get_global_expert_location_metadata(), - rank=self.ps.tp_rank, + rank=expert_rank, ) ) @@ -920,6 +1004,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, ) def init_lora_manager(self): @@ -1224,13 +1309,7 @@ class ModelRunner: self.msprobe_debugger.step() if self.server_args.elastic_ep_backend is not None: - recovered = maybe_recover_ep_ranks( - tp_group=self.tp_group, - eplb_manager=self.eplb_manager, - random_seed=self.server_args.random_seed, - ) - if recovered: - self.forward_pass_id = 0 + self.maybe_join_ep_ranks() return output @@ -1471,6 +1550,216 @@ class ModelRunner: action=action, allow_quant_error=allow_quant_error ) + def _expand_eplb_metadata_for_scale( + self, + from_ep_size: int, + effective_size: int, + ) -> None: + metadata = get_global_expert_location_metadata() + if metadata is None: + return + old_num_physical = metadata.num_physical_experts + num_local = old_num_physical // from_ep_size + added = num_local * effective_size - old_num_physical + if added <= 0: + return + + initial_ep_size = self.server_args.elastic_ep_initial_size + assert initial_ep_size is not None + self.server_args.override("elastic_ep.scale", ep_size=effective_size) + + expanded_p2l = append_trivial_expert_slots( + metadata.physical_to_logical_map, + added, + metadata.num_logical_experts, + start=old_num_physical - num_local * initial_ep_size, + ) + new_metadata = ExpertLocationMetadata.init_by_mapping( + self.server_args, + self.model_config, + physical_to_logical_map=expanded_p2l, + moe_ep_rank=self._elastic_global_rank(), + ) + set_global_expert_location_metadata(new_metadata, allow_overwrite=True) + + def _elastic_global_rank(self) -> int: + return self.ps.tp_rank + self.server_args.ep_join_rank_offset + + def _report_elastic_scale_failure(self, error: str, effective_size: int) -> None: + if self.ps.tp_rank != 0 or self.server_args.is_ep_scale_joiner: + return + from sglang.srt.managers.io_struct import ElasticScaleUpdateReq + + self._pending_elastic_scale_update = ElasticScaleUpdateReq( + success=False, + effective_ep_size=effective_size, + error=error, + ) + + def _elastic_scale_ready_barrier(self, target_size: int, log_tag: str) -> None: + if self.ps.tp_rank == 0: + logger.debug( + "[Elastic EP][scale] %s entering post-scale WORLD barrier " + "(target_ep_size=%d)", + log_tag, + target_size, + ) + dist.barrier(group=dist.group.WORLD) + if self.ps.tp_rank == 0: + logger.debug( + "[Elastic EP][scale] %s passed post-scale WORLD barrier " + "(target_ep_size=%d)", + log_tag, + target_size, + ) + + def _finalize_scale_up( + self, + ranks_to_join: list[int], + target_size: int, + effective_size: int, + ) -> None: + self.forward_pass_id = 0 + ElasticEPStateManager.mark_configuring_data_plane() + + state = ElasticEPStateManager.instance() + for rank in ranks_to_join: + state.active_ranks[rank] = 1 + state.snapshot_active_to_last() + state.sync_active_to_cpu() + if self.eplb_manager is not None: + self.eplb_manager.reset_generator() + + self._expand_eplb_metadata_for_scale( + from_ep_size=effective_size, + effective_size=target_size, + ) + broadcast_global_expert_location_metadata( + model_config=self.model_config, + moe_ep_rank=self._elastic_global_rank(), + src_rank=0, + ) + + ElasticEPStateManager.on_scale(effective_size, target_size) + set_global_expert_distribution_recorder( + ExpertDistributionRecorder.init_new( + self.server_args, + get_global_expert_location_metadata(), + rank=self._elastic_global_rank(), + ) + ) + + if self.eplb_manager is not None: + self.eplb_manager.disable_rebalance( + "EPLB rebalance is disabled after elastic EP scale-up" + ) + + from sglang.srt.layers.dp_attention import update_dp_attention_post_scale + + update_dp_attention_post_scale( + new_dp_size=target_size, + new_dp_rank=self._elastic_global_rank(), + ) + self.server_args.override("elastic_ep.scale", dp_size=target_size) + + ElasticEPStateManager.mark_syncing_new_world() + self._elastic_scale_ready_barrier( + target_size=target_size, + log_tag="JOINER" if self.server_args.is_ep_scale_joiner else "PRIMARY", + ) + ElasticEPStateManager.commit_scale() + + if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + from sglang.srt.managers.io_struct import ElasticScaleUpdateReq + + self._pending_elastic_scale_update = ElasticScaleUpdateReq( + success=True, + effective_ep_size=target_size, + slot_offset=effective_size, + slot_count=target_size - effective_size, + ) + logger.info( + "[Elastic EP] Scale completed: old_ep_size=%d " + "new_ep_size=%d joined_ranks=%s", + effective_size, + target_size, + ranks_to_join, + ) + + def maybe_join_ep_ranks(self) -> None: + if not ElasticEPStateManager.is_scaling(): + return + + state = ElasticEPStateManager.instance() + effective_size = ElasticEPStateManager.get_effective_ep_size() + pending_size = ElasticEPStateManager.get_pending_ep_size() + + if pending_size is None: + if state is not None and state.has_scaled: + error = ( + "Elastic EP rank recovery is unsupported after runtime scale-up. " + "Restart the expanded deployment." + ) + ElasticEPStateManager.fail_recovery(error) + self._report_elastic_scale_failure(error, effective_size) + if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + logger.error("[Elastic EP] %s", error) + return + + recovered = maybe_recover_ep_ranks( + tp_group=self.tp_group, + eplb_manager=self.eplb_manager, + random_seed=self.server_args.random_seed, + ) + if recovered: + self.forward_pass_id = 0 + return + + local_timeout = ( + state.pending_since is not None + and time.monotonic() - state.pending_since + > self.server_args.elastic_ep_scale_timeout + ) + timeout = state.active_ranks.new_tensor(int(local_timeout)) + dist.all_reduce(timeout, op=dist.ReduceOp.MAX, group=dist.group.WORLD) + if timeout.item(): + error = f"Timed out waiting for ranks to join target EP size {pending_size}" + ElasticEPStateManager.fail_scale(error) + self._report_elastic_scale_failure(error, effective_size) + if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + logger.error("[Elastic EP] %s", error) + return + + if state.scale_phase == "waiting_for_cohort": + cohort_target = get_scale_cohort_target(effective_size) + if cohort_target is None: + return + if cohort_target != pending_size: + error = ( + f"Requested target EP size {pending_size} does not match " + f"joining cohort target {cohort_target}" + ) + ElasticEPStateManager.fail_scale(error) + self._report_elastic_scale_failure(error, effective_size) + if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: + logger.error("[Elastic EP] %s", error) + return + if not ElasticEPStateManager.begin_scale(): + return + + ranks_to_join = list(range(effective_size, pending_size)) + if not ranks_to_join: + return + + current_platform.synchronize() + ElasticEPStateManager.mark_joining() + if try_admit_scale_ranks(ranks_to_join): + self._finalize_scale_up( + ranks_to_join=ranks_to_join, + target_size=pending_size, + effective_size=effective_size, + ) + def _maybe_rebalance_after_rank_fault( self, output: ModelRunnerOutput, 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 90e08e6a9..c55ba8007 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 @@ -248,10 +248,16 @@ def load_model_with_memory_saver( ) -def dist_barrier_after_load(*, elastic_ep_backend: Optional[str], tp_rank: int) -> None: +def dist_barrier_after_load( + *, + elastic_ep_backend: Optional[str], + tp_rank: int, + is_ep_scale_joiner: bool = False, +) -> None: if elastic_ep_backend == "mooncake": # Mooncake does not support `monitored_barrier` - dist.barrier(group=get_tp_group().cpu_group) + if not is_ep_scale_joiner: + dist.barrier(group=get_tp_group().cpu_group) else: # Handle the case where some ranks do not finish loading. try: diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 3b932e4eb..0ef936c09 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -345,6 +345,8 @@ class DpFlags(_FlagGroupBase): migrates them.""" enabled: bool = False + use_world_group_for_gather: bool = False + joiner_skip_all_gather: bool = False # Hybrid-SSM models materialize idle ranks via the MAX_LEN fabricated-row # conversion (set when hf_config has hybrid_override_pattern). max_len_with_idle: bool = False diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 0e7532857..15272cc28 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1998,9 +1998,40 @@ class ServerArgs: bool, "Enable Waterfill: dispatch the fused shared expert as an extra routed expert slot to the least-loaded EP rank. Supports DeepEP and MegaMOE MoE A2A backends, implicitly enables shared-expert fusion, and supports --deepep-mode auto, normal, or low_latency when used with DeepEP. Use auto or low_latency for production DeepEP decode so CUDA graph remains enabled. Supported on DeepSeek-V3/R1 with EP >= 2.", ] = False + ep_join_mode: A[ + Optional[Literal["scale", "recover"]], + Arg( + help="Join mode for elastic EP. 'recover' rejoins an existing slot after a fault. 'scale' joins as a new rank beyond the original group size and requires --node-rank 1.", + cli_name="--elastic-ep-join-mode", + choices=["scale", "recover"], + ), + ] = None + ep_join_rank_offset: A[ + int, + Arg( + help=( + "Global rank offset of an elastic EP joining group. Scale " + "joiners must set this to the current effective EP size." + ), + cli_name="--elastic-ep-join-rank-offset", + ), + ] = 0 + elastic_ep_initial_size: A[ + Optional[int], + "EP size used to define the immutable per-rank expert storage layout. " + "Scale joiners must use the primary deployment's launch-time EP size.", + ] = None + max_ep_size: A[ + Optional[int], + "Maximum EP size the server can scale to at runtime. Pre-allocates active-rank state and backend buffers to this size. Defaults to the launch-time world size.", + ] = None + elastic_ep_scale_timeout: A[ + float, + "Timeout in seconds for a pending elastic EP scale operation.", + ] = 600 elastic_ep_rejoin: A[ bool, - "Indicates that this process is a relaunched elastic EP rank that should rejoin an existing process group.", + "[Deprecated] Alias for --elastic-ep-join-mode recover.", ] = False disable_flashinfer_cutlass_moe_fp4_allgather: A[ bool, @@ -5634,10 +5665,21 @@ class ServerArgs: ): self.ep_dispatch_algorithm = "static" - if self.enable_eplb: + if self.enable_eplb and self.ep_join_mode != "scale": assert self._resolved().ep_size > 1 def _handle_elastic_ep(self): + if self.elastic_ep_rejoin: + if self.ep_join_mode is None: + logger.warning( + "--elastic-ep-rejoin is deprecated, use --elastic-ep-join-mode recover instead." + ) + self.ep_join_mode = "recover" + else: + assert self.ep_join_mode == "recover", ( + "--elastic-ep-rejoin (deprecated) conflicts with " + f"--elastic-ep-join-mode {self.ep_join_mode}." + ) if self.elastic_ep_backend is not None: if self.enable_eplb: if self.eplb_algorithm == "auto": @@ -5653,10 +5695,144 @@ class ServerArgs: self.mooncake_ib_device = self._validate_ib_devices( self.mooncake_ib_device ) - if self.elastic_ep_rejoin: + if self.ep_join_mode is not None: assert ( self.elastic_ep_backend is not None - ), "Elastic EP rejoin requires elastic_ep_backend to be set." + ), "--elastic-ep-join-mode requires --elastic-ep-backend to be set." + if self.ep_join_mode == "scale": + assert self.node_rank == 1, ( + "Elastic EP scale-up requires one joining TP group at " + f"--node-rank 1 (got {self.node_rank})." + ) + assert self.ep_join_rank_offset > 0, ( + "Elastic EP scale joiners require " + "--elastic-ep-join-rank-offset set to the current " + "effective EP size." + ) + if self.ep_join_rank_offset != 0: + assert self.ep_join_mode == "scale", ( + "--elastic-ep-join-rank-offset is only valid with " + "--elastic-ep-join-mode scale." + ) + assert ( + self.ep_join_rank_offset >= 0 + ), "elastic EP join rank offset must be >= 0." + if self.max_ep_size is not None: + assert ( + self.elastic_ep_backend is not None + ), "--max-ep-size requires --elastic-ep-backend to be set." + assert self.max_ep_size > 0, "--max-ep-size must be a positive integer." + + scaling_active = ( + self.elastic_ep_backend is not None + and self.max_ep_size is not None + and self.max_ep_size > self.tp_size + ) + if self.elastic_ep_initial_size is not None: + assert scaling_active, ( + "--elastic-ep-initial-size is only valid for an Elastic EP " + "deployment with --max-ep-size larger than its local TP size." + ) + if scaling_active: + resolved = self._resolved() + assert ( + self.elastic_ep_scale_timeout > 0 + ), "--elastic-ep-scale-timeout must be greater than zero." + assert self.tokenizer_worker_num == 1, ( + "Elastic EP runtime scale-up currently requires " + "--tokenizer-worker-num 1." + ) + assert ( + not self.use_ray + ), "Elastic EP runtime scale-up does not support --use-ray." + assert not self.enable_elastic_expert_backup, ( + "Elastic EP runtime scale-up does not support " + "--enable-elastic-expert-backup." + ) + self.enable_dp_attention_local_control_broadcast = True + if self.ep_join_mode == "scale": + assert self.elastic_ep_initial_size is not None, ( + "Elastic EP scale joiners require --elastic-ep-initial-size " + "set to the primary deployment's launch-time EP size." + ) + assert self.elastic_ep_initial_size <= self.ep_join_rank_offset, ( + "--elastic-ep-initial-size cannot exceed the current EP size " + f"(initial={self.elastic_ep_initial_size}, " + f"current={self.ep_join_rank_offset})." + ) + join_target = self.ep_join_rank_offset + self.tp_size + assert join_target <= self.max_ep_size, ( + "Elastic EP joining group exceeds --max-ep-size " + f"(join_target={join_target}, max_ep_size={self.max_ep_size})." + ) + if self.tp_size == 1: + assert self.moe_dense_tp_size == 1, ( + "A single-rank Elastic EP joining group requires " + "--moe-dense-tp-size 1." + ) + else: + if self.elastic_ep_initial_size is None: + self.elastic_ep_initial_size = self.tp_size + assert self.elastic_ep_initial_size == self.tp_size, ( + "The primary --elastic-ep-initial-size must equal its " + f"launch-time TP size ({self.tp_size})." + ) + assert self.elastic_ep_initial_size > 0 + assert self.load_balance_method == "round_robin", ( + "Elastic EP scale-up requires --load-balance-method round_robin; " + "load-aware methods " + "require global-rank load snapshots after scale " + f"(got {self.load_balance_method})." + ) + assert self.elastic_ep_backend == "mooncake", ( + "Elastic EP runtime scale-up requires --elastic-ep-backend " + f"mooncake (got elastic_ep_backend={self.elastic_ep_backend})." + ) + assert self.pp_size == 1, ( + "Elastic EP scale-up requires --pp-size 1 " + f"(got pp_size={self.pp_size}); WORLD must not span PP stages." + ) + + decode_cuda_graph_disabled = ( + self.cuda_graph_config.decode.backend == Backend.DISABLED + ) + prefill_cuda_graph_disabled = ( + self.cuda_graph_config.prefill.backend == Backend.DISABLED + ) + assert decode_cuda_graph_disabled and prefill_cuda_graph_disabled, ( + "Elastic EP runtime scale-up requires decode and prefill CUDA " + "graphs to be disabled." + ) + assert resolved.enable_dp_attention, ( + "Elastic EP scale-up requires --enable-dp-attention; without it " + "the TP group is not equivalent to WORLD and the post-scale " + "collective path is invalid." + ) + assert resolved.enable_dp_lm_head, ( + "Elastic EP scale-up requires --enable-dp-lm-head so output " + "projection does not depend on the joining group's TP size." + ) + assert resolved.attn_cp_size == 1, ( + "Elastic EP scale-up requires --attn-cp-size 1 " + f"(got attn_cp_size={resolved.attn_cp_size})." + ) + assert self.moe_dp_size == 1, ( + "Elastic EP scale-up requires --moe-dp-size 1 " + f"(got moe_dp_size={self.moe_dp_size})." + ) + assert resolved.ep_size == self.tp_size, ( + "Elastic EP scale-up requires ep_size == tp_size " + f"(got ep_size={resolved.ep_size}, tp_size={self.tp_size}); EP, TP " + "and the attention DP group must all coincide with WORLD." + ) + assert self.dp_size == self.tp_size, ( + "Elastic EP scale-up requires dp_size == tp_size " + f"(got dp_size={self.dp_size}, tp_size={self.tp_size})." + ) + assert resolved.moe_a2a_backend == "nixl", ( + "Elastic EP scale-up requires --moe-a2a-backend nixl " + f"(got moe_a2a_backend={resolved.moe_a2a_backend})." + ) def _handle_expert_distribution_metrics(self): if self.enable_expert_distribution_metrics and ( @@ -6917,6 +7093,15 @@ class ServerArgs: def engine_info_bootstrap_url(self): return self.url(port=self.engine_info_bootstrap_port) + @property + def is_ep_joiner(self) -> bool: + """True for processes launched as elastic-EP joiners.""" + return self.ep_join_mode in ("scale", "recover") + + @property + def is_ep_scale_joiner(self) -> bool: + return self.ep_join_mode == "scale" + def ssl_verify(self): """Return the value for the requests library's verify= parameter. @@ -7113,9 +7298,10 @@ class ServerArgs: def check_server_args(self): # Check parallel size constraints - assert ( - self.tp_size * self.pp_size - ) % self.nnodes == 0, "tp_size must be divisible by number of nodes" + if self.ep_join_mode != "scale": + assert ( + self.tp_size * self.pp_size + ) % self.nnodes == 0, "tp_size must be divisible by number of nodes" assert ( self.pp_max_micro_batch_size is None or self.pp_max_micro_batch_size >= 1 @@ -7870,7 +8056,11 @@ class PortArgs: # (no availability-based search). If incrementing would # overflow the valid TCP range, decrement instead. NUM_DERIVED_PORTS = 5 - if dist_init_port + NUM_DERIVED_PORTS > 65535: + if server_args.is_ep_scale_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 + elif dist_init_port + NUM_DERIVED_PORTS > 65535: port_base = dist_init_port - NUM_DERIVED_PORTS - 1 else: port_base = dist_init_port + 1 @@ -7886,9 +8076,11 @@ class PortArgs: assert worker_ports is not None scheduler_input_port = worker_ports[dp_rank] + is_joiner = server_args.is_ep_scale_joiner try: if dp_rank is None: - wait_port_available(dist_init_port, "dist_init_port") + if not is_joiner: + wait_port_available(dist_init_port, "dist_init_port") wait_port_available(port_base, "port_base") wait_port_available(detokenizer_port, "detokenizer_port") wait_port_available(nccl_port, "nccl_port") diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index d9931af35..578d55007 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -2401,7 +2401,7 @@ def _get_fastapi_request_path(request) -> Tuple[str, bool]: for route in request.app.routes: match, child_scope = route.matches(request.scope) if match == Match.FULL: - return route.path, True + return getattr(route, "path", request.url.path), True return request.url.path, False @@ -3475,6 +3475,13 @@ def require_mlp_tp_gather(server_args: ServerArgs): if server_args.enable_dp_attention: assert server_args.dp_size > 1, "dp_size must be greater than 1" + if server_args.elastic_ep_backend is not None: + from sglang.srt.elastic_ep.elastic_ep import ( + elastic_expanded_world_enabled, + ) + + if elastic_expanded_world_enabled(): + return True if ( server_args.moe_dense_tp_size is None ): # TODO(ch-wan): some MoE models do not have dense layers diff --git a/test/manual/ep/test_elastic_scale.py b/test/manual/ep/test_elastic_scale.py new file mode 100644 index 000000000..237526dd0 --- /dev/null +++ b/test/manual/ep/test_elastic_scale.py @@ -0,0 +1,479 @@ +"""Manual tests for elastic EP scale-up. + +Test classes: + TestElasticScaleUp4To6 primary + joiner scale-up (6 GPUs) + TestElasticScaleUp4To5To6 two consecutive single-rank scale-ups + TestElasticScaleUp4To8 full primary + joiner scale-up (8 GPUs) + +Run (8-GPU full scale-up): + + CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7 python -m pytest \\ + test/manual/ep/test_elastic_scale.py::TestElasticScaleUp4To8 \\ + -v -s +""" + +import os +import subprocess +import time +import unittest +from types import SimpleNamespace + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.run_eval import run_eval +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, + popen_launch_server, +) + +TEST_MODEL = os.environ.get("NIXL_EP_TEST_MODEL", DEFAULT_MODEL_NAME_FOR_TEST_MLA) +os.environ.setdefault("SGLANG_NIXL_EP_NUM_MAX_DISPATCH_TOKENS_PER_RANK", "1024") + +ib_devices = get_rdma_devices_args() + + +def _extra_server_args() -> list[str]: + """Extra `--flag [value]` tokens appended to every spawned server. + + Set via ``SGLANG_ELASTIC_EXTRA_SERVER_ARGS`` as a single space-separated + string, e.g. ``--disable-overlap-schedule``. + """ + raw = os.environ.get("SGLANG_ELASTIC_EXTRA_SERVER_ARGS", "").strip() + return raw.split() if raw else [] + + +DISABLED_CUDA_GRAPH_ARGS = [ + "--cuda-graph-backend-decode", + "disabled", + "--cuda-graph-backend-prefill", + "disabled", +] + + +def _assert_generate_logprob_ok(testcase: unittest.TestCase, base_url: str) -> None: + response = requests.post( + f"{base_url}/generate", + json={ + "text": "The answer is", + "sampling_params": {"max_new_tokens": 1, "temperature": 0.0}, + "return_logprob": True, + "top_logprobs_num": 1, + "logprob_start_len": 0, + }, + timeout=60, + ) + testcase.assertEqual(response.status_code, 200, response.text) + input_logprobs = response.json()["meta_info"]["input_token_logprobs"] + testcase.assertGreater(len(input_logprobs), 0) + + +def _count_visible_gpus() -> int: + env = os.environ.get("CUDA_VISIBLE_DEVICES") + if env: + return len([x for x in env.split(",") if x.strip()]) + try: + import torch + + return torch.cuda.device_count() if torch.cuda.is_available() else 0 + except Exception: + return 0 + + +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()] + return [str(index) for index in range(_count_visible_gpus())] + + +LAUNCH_EP_SIZE = 4 +MAX_EP_SIZE = 8 + +DIST_INIT_ADDR = os.environ.get("SGLANG_ELASTIC_SCALE_DIST_INIT", "127.0.0.1:24555") +PORT_A = int(os.environ.get("SGLANG_ELASTIC_SCALE_PORT_A", "21000")) +PORT_B = int(os.environ.get("SGLANG_ELASTIC_SCALE_PORT_B", "10000")) +PORT_C = int(os.environ.get("SGLANG_ELASTIC_SCALE_PORT_C", "11000")) +HOST_A = os.environ.get("SGLANG_ELASTIC_SCALE_HOST_A", "127.0.0.1") +BASE_URL_A = f"http://{HOST_A}:{PORT_A}" +PRE_SCALE_JOINER_DELAY_SEC = float( + os.environ.get("SGLANG_ELASTIC_PRE_SCALE_JOINER_DELAY_SEC", "0") +) + + +def _scale_up_common_args( + dist_init_addr: str, + tp_size: int, + nnodes: int, + node_rank: int, + cuda_graph_args: list[str], + moe_dense_tp_size: int | None, +) -> list[str]: + args = [ + "--trust-remote-code", + "--moe-a2a-backend", + "nixl", + "--deepep-mode", + "low_latency", + "--tp", + str(tp_size), + "--dp", + str(tp_size), + "--enable-dp-attention", + "--enable-dp-lm-head", + "--elastic-ep-backend", + "mooncake", + "--mooncake-ib-device", + ib_devices, + "--enable-eplb", + "--ep-num-redundant-experts", + "24", + "--elastic-ep-initial-size", + str(LAUNCH_EP_SIZE), + "--max-ep-size", + str(MAX_EP_SIZE), + "--mem-fraction-static", + "0.5", + "--chunked-prefill-size", + "1024", + "--nnodes", + str(nnodes), + "--node-rank", + str(node_rank), + "--dist-init-addr", + dist_init_addr, + ] + if moe_dense_tp_size is not None: + args.extend(["--moe-dense-tp-size", str(moe_dense_tp_size)]) + return args + cuda_graph_args + _extra_server_args() + + +class _ElasticScaleUpEndToEndBase(CustomTestCase): + """Shared scale-up E2E plumbing. Subclasses set JOIN_TP/JOIN_NNODES/JOIN_NODE_RANK.""" + + JOIN_TP: int + JOIN_NNODES: int + JOIN_NODE_RANK: int + TARGET_EP_SIZE: int + CUDA_GRAPH_ARGS: list[str] + MOE_DENSE_TP_SIZE: int | None = 1 + + def setUp(self): + if ( + not hasattr(type(self), "JOIN_TP") + or type(self) is _ElasticScaleUpEndToEndBase + ): + self.skipTest("Abstract base — run a concrete subclass instead") + + @classmethod + def setUpClass(cls): + if cls is _ElasticScaleUpEndToEndBase: + raise unittest.SkipTest("Abstract base") + cls.model = TEST_MODEL + cls.base_url = BASE_URL_A + cls._joining_procs = [] + cls._joining_log_fhs = [] + + primary_args = _scale_up_common_args( + DIST_INIT_ADDR, + tp_size=LAUNCH_EP_SIZE, + nnodes=1, + node_rank=0, + cuda_graph_args=cls.CUDA_GRAPH_ARGS, + moe_dense_tp_size=cls.MOE_DENSE_TP_SIZE, + ) + primary_env = os.environ.copy() + visible_devices = _visible_device_ids() + if len(visible_devices) < LAUNCH_EP_SIZE: + raise RuntimeError( + f"Scale-up requires {LAUNCH_EP_SIZE} visible GPUs, got " + f"{len(visible_devices)}" + ) + primary_env["CUDA_VISIBLE_DEVICES"] = ",".join(visible_devices[:LAUNCH_EP_SIZE]) + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=primary_args, + env=primary_env, + ) + + @classmethod + def _launch_joining_group( + cls, + *, + rank_offset: int, + join_tp: int, + port: int, + ) -> subprocess.Popen: + cmd = [ + "sglang", + "serve", + "--model-path", + cls.model, + *_scale_up_common_args( + DIST_INIT_ADDR, + tp_size=join_tp, + nnodes=cls.JOIN_NNODES, + node_rank=cls.JOIN_NODE_RANK, + cuda_graph_args=cls.CUDA_GRAPH_ARGS, + moe_dense_tp_size=cls.MOE_DENSE_TP_SIZE, + ), + "--elastic-ep-join-mode", + "scale", + "--elastic-ep-join-rank-offset", + str(rank_offset), + "--host", + "127.0.0.1", + "--port", + str(port), + "--device", + "cuda", + ] + env = os.environ.copy() + visible_devices = _visible_device_ids() + join_end = rank_offset + join_tp + if join_end > len(visible_devices): + raise RuntimeError( + f"Scale-up requires {join_end} visible GPUs, got " + f"{len(visible_devices)}" + ) + env["CUDA_VISIBLE_DEVICES"] = ",".join(visible_devices[rank_offset:join_end]) + base_joining_log = os.environ.get( + "SGLANG_ELASTIC_SCALE_JOINING_LOG", + f"/tmp/elastic_scale_joining_nnodes{cls.JOIN_NNODES}_{int(time.time())}.log", + ) + if cls._joining_procs: + root, ext = os.path.splitext(base_joining_log) + joining_log = f"{root}_step{len(cls._joining_procs) + 1}{ext}" + else: + joining_log = base_joining_log + joining_log_fh = open(joining_log, "w") + joining_proc = subprocess.Popen( + cmd, + env=env, + stdout=joining_log_fh, + stderr=subprocess.STDOUT, + ) + cls._joining_procs.append(joining_proc) + cls._joining_log_fhs.append(joining_log_fh) + return joining_proc + + @classmethod + def tearDownClass(cls): + processes = [ + *reversed(cls._joining_procs), + getattr(cls, "process", None), + ] + for proc in processes: + if proc is None: + continue + try: + kill_process_tree(proc.pid) + except Exception: + pass + for proc in processes: + if proc is None: + continue + try: + proc.wait(timeout=15) + except Exception: + pass + for fh in cls._joining_log_fhs: + try: + fh.close() + except Exception: + pass + time.sleep(2) + + def _post(self, path: str, **kwargs) -> requests.Response: + return requests.post(f"{self.base_url}{path}", timeout=60, **kwargs) + + def _generate_ok(self, msg_suffix: str, routed_dp_rank: int | None = None) -> None: + payload = { + "text": "Hello", + "sampling_params": {"max_new_tokens": 4, "temperature": 0.0}, + } + if routed_dp_rank is not None: + payload["routed_dp_rank"] = routed_dp_rank + resp = self._post( + "/generate", + json=payload, + ) + self.assertEqual( + resp.status_code, + 200, + f"/generate {msg_suffix} failed: {resp.text}", + ) + + def _generate_logprob_ok(self, msg_suffix: str) -> None: + try: + _assert_generate_logprob_ok(self, self.base_url) + except AssertionError as exc: + raise AssertionError( + f"/generate logprob {msg_suffix} failed: {exc}" + ) from exc + + def _scale_once( + self, + *, + old_ep_size: int, + target_ep_size: int, + join_tp: int, + port: int, + ) -> None: + joining_proc = self._launch_joining_group( + rank_offset=old_ep_size, + join_tp=join_tp, + port=port, + ) + self.assertIsNone( + joining_proc.poll(), + "Joining group exited before scale request; see joining log", + ) + if PRE_SCALE_JOINER_DELAY_SEC > 0: + time.sleep(PRE_SCALE_JOINER_DELAY_SEC) + self.assertIsNone( + joining_proc.poll(), + "Joining group exited before scale request; see joining log", + ) + + resp = self._post("/scale_elastic_ep", json={"new_ep_size": target_ep_size}) + self.assertEqual(resp.status_code, 200, resp.text) + body = resp.json() + self.assertEqual(body["old_ep_size"], old_ep_size) + self.assertEqual(body["new_ep_size"], target_ep_size) + + deadline = time.time() + 300 + while time.time() < deadline: + resp = requests.get(f"{self.base_url}/is_scaling_elastic_ep", timeout=60) + state = resp.json() if resp.ok else None + if state is not None and not state.get("is_scaling_elastic_ep", True): + self.assertEqual(state.get("effective_ep_size"), target_ep_size) + self.assertEqual(state.get("scale_phase"), "serving_expanded") + self.assertIsNone(state.get("last_error")) + self._generate_ok( + "on newest joiner", + routed_dp_rank=target_ep_size - 1, + ) + return + try: + self._post( + "/generate", + json={ + "text": "ping", + "sampling_params": {"max_new_tokens": 1, "temperature": 0.0}, + }, + ) + except Exception: + pass + time.sleep(2) + self.fail("Timed out waiting for scaling to complete (300s)") + + def _run_post_scale_gsm8k(self) -> None: + metrics = run_eval( + SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=256, + num_threads=50, + ) + ) + print(f"[TEST] Post-scale GSM8K accuracy: {metrics['score']:.2%}") + self.assertGreater( + metrics["score"], + 0.50, + f"Post-scale GSM8K accuracy too low: {metrics['score']:.2%}", + ) + + def test_scale_up_on_demand(self): + """Scale the primary group to the configured target.""" + self._generate_ok("pre-scale") + + self._scale_once( + old_ep_size=LAUNCH_EP_SIZE, + target_ep_size=self.TARGET_EP_SIZE, + join_tp=self.JOIN_TP, + port=PORT_B, + ) + + self._generate_ok("post-scale") + self._generate_logprob_ok("post-scale") + + self._run_post_scale_gsm8k() + + +@unittest.skipUnless( + _count_visible_gpus() >= 6, + "4-to-6 scale-up E2E needs 6 GPUs.", +) +class TestElasticScaleUp4To6(_ElasticScaleUpEndToEndBase): + """Scale from four to six ranks.""" + + JOIN_TP = 2 + JOIN_NNODES = 2 + JOIN_NODE_RANK = 1 + TARGET_EP_SIZE = 6 + CUDA_GRAPH_ARGS = DISABLED_CUDA_GRAPH_ARGS + MOE_DENSE_TP_SIZE = None + + +@unittest.skipUnless( + _count_visible_gpus() >= 6, + "4-to-5-to-6 scale-up E2E needs 6 GPUs.", +) +class TestElasticScaleUp4To5To6(_ElasticScaleUpEndToEndBase): + """Scale from four to five and then from five to six ranks.""" + + JOIN_TP = 1 + JOIN_NNODES = 2 + JOIN_NODE_RANK = 1 + TARGET_EP_SIZE = 6 + CUDA_GRAPH_ARGS = DISABLED_CUDA_GRAPH_ARGS + + def test_scale_up_on_demand(self): + self._generate_ok("pre-scale") + + self._scale_once( + old_ep_size=4, + target_ep_size=5, + join_tp=1, + port=PORT_B, + ) + self._generate_ok("after first scale") + self._generate_logprob_ok("after first scale") + + self._scale_once( + old_ep_size=5, + target_ep_size=6, + join_tp=1, + port=PORT_C, + ) + self._generate_ok("after second scale") + self._generate_logprob_ok("after second scale") + self._run_post_scale_gsm8k() + + +@unittest.skipUnless( + _count_visible_gpus() >= MAX_EP_SIZE, + f"Full scale-up E2E needs {MAX_EP_SIZE} GPUs.", +) +class TestElasticScaleUp4To8(_ElasticScaleUpEndToEndBase): + """Scale from four to eight ranks.""" + + JOIN_TP = LAUNCH_EP_SIZE + JOIN_NNODES = 2 + JOIN_NODE_RANK = 1 + TARGET_EP_SIZE = MAX_EP_SIZE + CUDA_GRAPH_ARGS = DISABLED_CUDA_GRAPH_ARGS + MOE_DENSE_TP_SIZE = None + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/entrypoints/test_v1_loads_aggregate.py b/test/registered/unit/entrypoints/test_v1_loads_aggregate.py index 8b325c890..e4e9240d8 100644 --- a/test/registered/unit/entrypoints/test_v1_loads_aggregate.py +++ b/test/registered/unit/entrypoints/test_v1_loads_aggregate.py @@ -42,6 +42,7 @@ def _temp_path() -> str: class _FakeTokenizerManager(TokenizerControlMixin): def __init__(self, reader, dp_size: int): self.load_snapshot_reader = reader + self.elastic_worker_count = dp_size self.server_args = SimpleNamespace( dp_size=dp_size, enable_dp_attention=False, diff --git a/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py b/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py index 7b9530485..ea73a7e0d 100644 --- a/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py +++ b/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py @@ -10,14 +10,16 @@ import unittest import torch from sglang.srt.eplb.expert_location import ( + _compute_logical_to_all_physical_map, + append_trivial_expert_slots, compute_logical_to_rank_dispatch_physical_map, ) from sglang.test.test_utils import CustomTestCase def _make_server_args(ep_size: int, nnodes: int): - """Minimal server_args stub — only ep_size and nnodes are used.""" - return types.SimpleNamespace(ep_size=ep_size, nnodes=nnodes) + """Minimal server_args stub for expert placement tests.""" + return types.SimpleNamespace(ep_size=ep_size, nnodes=nnodes, ep_join_mode=None) def _make_logical_to_all_physical_map( @@ -194,6 +196,24 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase): self.assertEqual(result.shape, (self.NUM_LAYERS, 1)) self.assertTrue(torch.all(result >= 0)) + def test_scale_joiner_maps_appended_expert_slots(self): + physical_to_logical = torch.arange(64).unsqueeze(0) + physical_to_logical = append_trivial_expert_slots( + physical_to_logical, count=16, num_logical_experts=64 + ) + server_args = _make_server_args(ep_size=5, nnodes=1) + server_args.ep_join_mode = "scale" + + logical_to_physical = _compute_logical_to_all_physical_map( + server_args=server_args, + physical_to_logical_map=physical_to_logical, + num_logical_experts=64, + ep_size=5, + moe_ep_rank=4, + ) + + self.assertEqual(logical_to_physical[0, :16, 0].tolist(), list(range(64, 80))) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/managers/test_data_parallel_controller.py b/test/registered/unit/managers/test_data_parallel_controller.py index 13c0917d6..6ab3cfc53 100644 --- a/test/registered/unit/managers/test_data_parallel_controller.py +++ b/test/registered/unit/managers/test_data_parallel_controller.py @@ -5,8 +5,8 @@ test/registered/disaggregation/test_disaggregation_dp_attention.py; its tie-break on `total_requests` transitively covers that state. Fragility: scheduler tests bypass `DataParallelController.__init__` via -`__new__` and inject only the attrs the schedulers read (`workers`, -`status`, `round_robin_counter`, `dp_budget`). Update `_make_controller` +`__new__` and inject only the attrs the schedulers read (`workers`, `status`, +`_active_workers`, `round_robin_counter`, `dp_budget`). Update `_make_controller` if a scheduler starts reading another attr. `maybe_external_dp_rank_routing` is exercised as the real method, no mock. """ @@ -48,6 +48,7 @@ def _make_controller(dp_size: int) -> DataParallelController: ctl = DataParallelController.__new__(DataParallelController) ctl.workers = [MagicMock(name=f"worker_{i}") for i in range(dp_size)] ctl.status = [True] * dp_size + ctl._active_workers = list(range(dp_size)) ctl.round_robin_counter = 0 ctl.dp_budget = DPBudget(dp_size=dp_size) return ctl diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 7ba5ed464..471b64884 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -1812,11 +1812,16 @@ class TestGoldenModelOverrides(_IsolatedPublish): ) self.assertEqual( - _data_parallelism_defaults(ResolvedView(SimpleNamespace(dp_size=1))), + _data_parallelism_defaults( + ResolvedView(SimpleNamespace(dp_size=1, ep_join_mode=None)) + ), {"enable_dp_attention": False, "enable_dp_lm_head": False}, ) self.assertEqual( - _data_parallelism_defaults(ResolvedView(SimpleNamespace(dp_size=2))), {} + _data_parallelism_defaults( + ResolvedView(SimpleNamespace(dp_size=2, ep_join_mode=None)) + ), + {}, ) with patch("sglang.srt.environ.envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE") as e: