[1/N] elastic-ep: Add runtime EP scale-up (#30164)

This commit is contained in:
Yoray Zack
2026-07-16 15:53:44 -07:00
committed by GitHub
parent d28e35b1a1
commit 77d23a796e
34 changed files with 2549 additions and 338 deletions
+1 -1
View File
@@ -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 {}
+15 -4
View File
@@ -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,
+121 -26
View File
@@ -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,
)
+289 -58
View File
@@ -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()
@@ -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())
+6 -5
View File
@@ -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":
+14 -2
View File
@@ -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:
+16
View File
@@ -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
+18 -3
View File
@@ -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
+130 -65
View File
@@ -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
@@ -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:
+68 -14
View File
@@ -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
)
@@ -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
@@ -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
+22 -2
View File
@@ -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
@@ -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()
+25
View File
@@ -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
+160 -15
View File
@@ -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()
@@ -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(
@@ -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(
@@ -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
@@ -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(
+11 -7
View File
@@ -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
@@ -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
+309 -20
View File
@@ -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,
@@ -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:
+2
View File
@@ -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
+201 -9
View File
@@ -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")
+8 -1
View File
@@ -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
+479
View File
@@ -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()
@@ -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,
@@ -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()
@@ -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
+7 -2
View File
@@ -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: