[1/N] elastic-ep: Add runtime EP scale-up (#30164)
This commit is contained in:
@@ -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 {}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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())
|
||||
@@ -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":
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user