[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
|
@register_post_process
|
||||||
def _data_parallelism_defaults(view: Any) -> dict:
|
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 {"enable_dp_attention": False, "enable_dp_lm_head": False}
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|||||||
@@ -225,15 +225,24 @@ def _init_parallel_groups(
|
|||||||
moe_dp_size: int,
|
moe_dp_size: int,
|
||||||
dcp_size: int,
|
dcp_size: int,
|
||||||
) -> None:
|
) -> 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(
|
init_distributed_environment(
|
||||||
backend=backend,
|
backend=backend,
|
||||||
world_size=tp_size * pp_size,
|
world_size=world_size,
|
||||||
rank=tp_size * pp_rank + tp_rank,
|
rank=rank,
|
||||||
local_rank=gpu_id,
|
local_rank=gpu_id,
|
||||||
distributed_init_method=dist_init_method,
|
distributed_init_method=dist_init_method,
|
||||||
timeout=server_args.dist_timeout,
|
timeout=server_args.dist_timeout,
|
||||||
moe_a2a_backend=server_args.moe_a2a_backend,
|
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(
|
initialize_model_parallel(
|
||||||
tensor_model_parallel_size=tp_size,
|
tensor_model_parallel_size=tp_size,
|
||||||
@@ -245,7 +254,9 @@ def _init_parallel_groups(
|
|||||||
decode_context_parallel_size=dcp_size,
|
decode_context_parallel_size=dcp_size,
|
||||||
duplicate_tp_group=server_args.enable_pdmux,
|
duplicate_tp_group=server_args.enable_pdmux,
|
||||||
enable_symm_mem=server_args.enable_symm_mem,
|
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(
|
initialize_dp_attention(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
|
|||||||
@@ -270,6 +270,8 @@ class GroupCoordinator:
|
|||||||
group_name: Optional[str] = None,
|
group_name: Optional[str] = None,
|
||||||
gloo_timeout: timedelta = timedelta(seconds=120 * 60),
|
gloo_timeout: timedelta = timedelta(seconds=120 * 60),
|
||||||
recovered_rank: bool = False,
|
recovered_rank: bool = False,
|
||||||
|
rank_offset: int = 0,
|
||||||
|
max_world_size: Optional[int] = None,
|
||||||
):
|
):
|
||||||
# Set group info
|
# Set group info
|
||||||
group_name = group_name or "anonymous"
|
group_name = group_name or "anonymous"
|
||||||
@@ -278,6 +280,9 @@ class GroupCoordinator:
|
|||||||
|
|
||||||
# Set rank info
|
# Set rank info
|
||||||
self.rank = torch.distributed.get_rank()
|
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.local_rank = local_rank
|
||||||
self.device_group = None
|
self.device_group = None
|
||||||
self.cpu_group = None
|
self.cpu_group = None
|
||||||
@@ -299,25 +304,57 @@ class GroupCoordinator:
|
|||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
|
|
||||||
for ranks in group_ranks:
|
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
|
subgroup_timeout = _MODEL_PARALLEL_GROUP_TIMEOUT
|
||||||
if "mooncake" in torch_distributed_backend:
|
if "mooncake" in torch_distributed_backend:
|
||||||
from mooncake.ep import MooncakeBackendOptions
|
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(
|
device_group = torch.distributed.new_group(
|
||||||
ranks,
|
ranks,
|
||||||
backend="mooncake",
|
backend="mooncake",
|
||||||
pg_options=MooncakeBackendOptions(active_ranks, recovered_rank),
|
pg_options=dev_opts,
|
||||||
timeout=subgroup_timeout,
|
timeout=subgroup_timeout,
|
||||||
)
|
)
|
||||||
cpu_group = torch.distributed.new_group(
|
cpu_group = torch.distributed.new_group(
|
||||||
ranks,
|
ranks,
|
||||||
backend="mooncake-cpu",
|
backend="mooncake-cpu",
|
||||||
pg_options=MooncakeBackendOptions(active_ranks_cpu, recovered_rank),
|
pg_options=cpu_opts,
|
||||||
timeout=subgroup_timeout,
|
timeout=subgroup_timeout,
|
||||||
)
|
)
|
||||||
else:
|
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)
|
pg_options = get_torch_distributed_pg_options(group_name)
|
||||||
device_group = torch.distributed.new_group(
|
device_group = torch.distributed.new_group(
|
||||||
ranks,
|
ranks,
|
||||||
@@ -1657,6 +1694,8 @@ def init_model_parallel_group(
|
|||||||
use_mscclpp_allreduce: Optional[bool] = None,
|
use_mscclpp_allreduce: Optional[bool] = None,
|
||||||
use_torch_symm_mem_allreduce: Optional[bool] = None,
|
use_torch_symm_mem_allreduce: Optional[bool] = None,
|
||||||
recovered_rank: bool = False,
|
recovered_rank: bool = False,
|
||||||
|
rank_offset: int = 0,
|
||||||
|
max_world_size: Optional[int] = None,
|
||||||
) -> GroupCoordinator:
|
) -> GroupCoordinator:
|
||||||
if use_custom_allreduce is None:
|
if use_custom_allreduce is None:
|
||||||
use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE
|
use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE
|
||||||
@@ -1682,6 +1721,8 @@ def init_model_parallel_group(
|
|||||||
use_message_queue_broadcaster=use_message_queue_broadcaster,
|
use_message_queue_broadcaster=use_message_queue_broadcaster,
|
||||||
group_name=group_name,
|
group_name=group_name,
|
||||||
recovered_rank=recovered_rank,
|
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")
|
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.
|
"""Create a global TCPStore for coordination across ranks.
|
||||||
|
|
||||||
This function creates a TCPStore that all ranks can use for coordination
|
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
|
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:
|
if not master_ip:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Could not determine master IP for global TCPStore. "
|
"Could not determine master IP for global TCPStore. "
|
||||||
"Broadcasting from rank 0 to all ranks."
|
"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:
|
if rank == 0:
|
||||||
master_ip = get_local_ip_auto()
|
master_ip = get_local_ip_auto()
|
||||||
ip_list = [master_ip]
|
ip_list = [master_ip]
|
||||||
else:
|
else:
|
||||||
ip_list = [None]
|
ip_list = [None]
|
||||||
|
|
||||||
torch.distributed.broadcast_object_list(ip_list, src=0)
|
torch.distributed.broadcast_object_list(ip_list, src=0)
|
||||||
master_ip = ip_list[0]
|
master_ip = ip_list[0]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
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(
|
tcp_store = TCPStore(
|
||||||
host_name=master_ip,
|
host_name=master_ip,
|
||||||
port=base_store_port,
|
port=base_store_port,
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
is_master=(rank == 0),
|
is_master=is_master,
|
||||||
)
|
)
|
||||||
set_global_tcp_store(tcp_store)
|
set_global_tcp_store(tcp_store)
|
||||||
logger.info(
|
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,
|
master_ip,
|
||||||
base_store_port,
|
base_store_port,
|
||||||
rank,
|
rank,
|
||||||
world_size,
|
is_master,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -1906,6 +1961,7 @@ def init_distributed_environment(
|
|||||||
timeout: Optional[int] = None,
|
timeout: Optional[int] = None,
|
||||||
moe_a2a_backend: Optional[str] = None,
|
moe_a2a_backend: Optional[str] = None,
|
||||||
recovered_rank: bool = False,
|
recovered_rank: bool = False,
|
||||||
|
max_world_size: Optional[int] = None,
|
||||||
):
|
):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"world_size=%d rank=%d local_rank=%d " "distributed_init_method=%s backend=%s",
|
"world_size=%d rank=%d local_rank=%d " "distributed_init_method=%s backend=%s",
|
||||||
@@ -1942,8 +1998,15 @@ def init_distributed_environment(
|
|||||||
if backend == "mooncake":
|
if backend == "mooncake":
|
||||||
from mooncake.ep import MooncakeBackendOptions
|
from mooncake.ep import MooncakeBackendOptions
|
||||||
|
|
||||||
# Setting "cuda" as device here is safe, as it is guarded under the mooncake case
|
use_max_ws = max_world_size and max_world_size > world_size
|
||||||
active_ranks = torch.ones(world_size, dtype=torch.int32, device="cuda")
|
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)
|
pg_options = MooncakeBackendOptions(active_ranks, recovered_rank)
|
||||||
else:
|
else:
|
||||||
pg_options = get_torch_distributed_pg_options()
|
pg_options = get_torch_distributed_pg_options()
|
||||||
@@ -1960,7 +2023,15 @@ def init_distributed_environment(
|
|||||||
|
|
||||||
# Create a global TCPStore for coordination (used by NIXL)
|
# Create a global TCPStore for coordination (used by NIXL)
|
||||||
if moe_a2a_backend == "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
|
# set the local rank
|
||||||
# local_rank is not available in torch ProcessGroup,
|
# local_rank is not available in torch ProcessGroup,
|
||||||
@@ -1996,6 +2067,8 @@ def initialize_model_parallel(
|
|||||||
duplicate_tp_group: bool = False,
|
duplicate_tp_group: bool = False,
|
||||||
enable_symm_mem: bool = False,
|
enable_symm_mem: bool = False,
|
||||||
recovered_rank: bool = False,
|
recovered_rank: bool = False,
|
||||||
|
rank_offset: int = 0,
|
||||||
|
max_world_size: Optional[int] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Initialize model parallel groups.
|
Initialize model parallel groups.
|
||||||
@@ -2048,9 +2121,15 @@ def initialize_model_parallel(
|
|||||||
"""
|
"""
|
||||||
# Get world size and rank. Ensure some consistencies.
|
# Get world size and rank. Ensure some consistencies.
|
||||||
assert torch.distributed.is_initialized()
|
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)
|
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:
|
if world_size != tensor_model_parallel_size * pipeline_model_parallel_size:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"world_size ({world_size}) is not equal to "
|
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(),
|
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
|
||||||
group_name="tp",
|
group_name="tp",
|
||||||
recovered_rank=recovered_rank,
|
recovered_rank=recovered_rank,
|
||||||
|
rank_offset=rank_offset,
|
||||||
|
max_world_size=max_world_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
if duplicate_tp_group:
|
if duplicate_tp_group:
|
||||||
@@ -2110,6 +2191,8 @@ def initialize_model_parallel(
|
|||||||
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
|
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
|
||||||
group_name="pdmux_prefill_tp",
|
group_name="pdmux_prefill_tp",
|
||||||
recovered_rank=recovered_rank,
|
recovered_rank=recovered_rank,
|
||||||
|
rank_offset=rank_offset,
|
||||||
|
max_world_size=max_world_size,
|
||||||
)
|
)
|
||||||
if _TP.pynccl_comm:
|
if _TP.pynccl_comm:
|
||||||
_TP.pynccl_comm.disabled = False
|
_TP.pynccl_comm.disabled = False
|
||||||
@@ -2172,6 +2255,8 @@ def initialize_model_parallel(
|
|||||||
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
|
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
|
||||||
group_name="attn_cp",
|
group_name="attn_cp",
|
||||||
recovered_rank=recovered_rank,
|
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
|
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(),
|
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
|
||||||
group_name="attention_tp",
|
group_name="attention_tp",
|
||||||
recovered_rank=recovered_rank,
|
recovered_rank=recovered_rank,
|
||||||
|
rank_offset=rank_offset,
|
||||||
|
max_world_size=max_world_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
moe_ep_size = expert_model_parallel_size
|
moe_ep_size = expert_model_parallel_size
|
||||||
@@ -2239,6 +2326,8 @@ def initialize_model_parallel(
|
|||||||
backend,
|
backend,
|
||||||
group_name="moe_dp",
|
group_name="moe_dp",
|
||||||
recovered_rank=recovered_rank,
|
recovered_rank=recovered_rank,
|
||||||
|
rank_offset=rank_offset,
|
||||||
|
max_world_size=max_world_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
global _MOE_EP
|
global _MOE_EP
|
||||||
@@ -2267,6 +2356,8 @@ def initialize_model_parallel(
|
|||||||
use_custom_allreduce=False,
|
use_custom_allreduce=False,
|
||||||
group_name="moe_ep",
|
group_name="moe_ep",
|
||||||
recovered_rank=recovered_rank,
|
recovered_rank=recovered_rank,
|
||||||
|
rank_offset=rank_offset,
|
||||||
|
max_world_size=max_world_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
global _MOE_TP
|
global _MOE_TP
|
||||||
@@ -2295,6 +2386,8 @@ def initialize_model_parallel(
|
|||||||
use_custom_allreduce=False,
|
use_custom_allreduce=False,
|
||||||
group_name="moe_tp",
|
group_name="moe_tp",
|
||||||
recovered_rank=recovered_rank,
|
recovered_rank=recovered_rank,
|
||||||
|
rank_offset=rank_offset,
|
||||||
|
max_world_size=max_world_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Build the pipeline model-parallel groups.
|
# Build the pipeline model-parallel groups.
|
||||||
@@ -2315,6 +2408,8 @@ def initialize_model_parallel(
|
|||||||
use_custom_allreduce=False,
|
use_custom_allreduce=False,
|
||||||
group_name="pp",
|
group_name="pp",
|
||||||
recovered_rank=recovered_rank,
|
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 logging
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Iterator, List, Optional
|
from typing import TYPE_CHECKING, Callable, Iterator, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.distributed import get_world_group, parallel_state
|
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.eplb.expert_location import broadcast_global_expert_location_metadata
|
||||||
from sglang.srt.managers.schedule_batch import ServerArgs
|
from sglang.srt.managers.schedule_batch import ServerArgs
|
||||||
from sglang.srt.utils import broadcast_pyobj, is_cpu, is_cuda
|
from sglang.srt.utils import broadcast_pyobj, is_cpu, is_cuda
|
||||||
@@ -17,12 +18,39 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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
|
@dataclass
|
||||||
class ElasticEPState:
|
class ElasticEPState:
|
||||||
active_ranks: Optional[torch.Tensor]
|
active_ranks: Optional[torch.Tensor]
|
||||||
last_active_ranks: Optional[torch.Tensor]
|
last_active_ranks: Optional[torch.Tensor]
|
||||||
active_ranks_cpu: 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:
|
def is_active_equal_last(self) -> bool:
|
||||||
return torch.equal(self.active_ranks, self.last_active_ranks)
|
return torch.equal(self.active_ranks, self.last_active_ranks)
|
||||||
@@ -37,13 +65,16 @@ class ElasticEPState:
|
|||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
if self.active_ranks is not None:
|
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.snapshot_active_to_last()
|
||||||
self.sync_active_to_cpu()
|
self.sync_active_to_cpu()
|
||||||
|
|
||||||
|
|
||||||
class ElasticEPStateManager:
|
class ElasticEPStateManager:
|
||||||
_instance: Optional[ElasticEPState] = None
|
_instance: Optional[ElasticEPState] = None
|
||||||
|
_on_scale: Optional[Callable[[int, int], None]] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def instance(cls) -> ElasticEPState:
|
def instance(cls) -> ElasticEPState:
|
||||||
@@ -55,16 +86,53 @@ class ElasticEPStateManager:
|
|||||||
return cls._instance
|
return cls._instance
|
||||||
|
|
||||||
if server_args.elastic_ep_backend is not None:
|
if server_args.elastic_ep_backend is not None:
|
||||||
cls._instance = cls._build_state(ep_size=None, device=None)
|
world_size = torch.distributed.get_world_size()
|
||||||
if server_args.elastic_ep_rejoin:
|
active_rank_capacity = server_args.max_ep_size or world_size
|
||||||
# Mask out peer ranks to perform cuda graph capture on its own
|
assert active_rank_capacity >= world_size, (
|
||||||
cls._instance.active_ranks.zero_()
|
f"--max-ep-size ({active_rank_capacity}) must be >= "
|
||||||
cls._instance.active_ranks[torch.distributed.get_rank()] = 1
|
f"world_size ({world_size})."
|
||||||
cls._instance.snapshot_active_to_last()
|
)
|
||||||
cls._instance.sync_active_to_cpu()
|
|
||||||
|
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
|
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
|
@staticmethod
|
||||||
def _select_device() -> torch.device:
|
def _select_device() -> torch.device:
|
||||||
if is_cuda():
|
if is_cuda():
|
||||||
@@ -94,39 +162,197 @@ class ElasticEPStateManager:
|
|||||||
|
|
||||||
return torch.ones(size, dtype=torch.int32, device=dev)
|
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
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
@classmethod
|
||||||
# Helpers for elastic EP recovery
|
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
|
_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]:
|
def _iter_live_parallel_groups() -> Iterator[parallel_state.GroupCoordinator]:
|
||||||
groups = []
|
groups = []
|
||||||
for group_ref in parallel_state._groups.values():
|
for group_ref in parallel_state._groups.values():
|
||||||
group = group_ref()
|
group = group_ref()
|
||||||
if group is not None:
|
if group is not None:
|
||||||
groups.append(group)
|
groups.append(group)
|
||||||
for group in sorted(groups, key=lambda x: x.unique_name):
|
yield from sorted(groups, key=lambda group: group.unique_name)
|
||||||
yield group
|
|
||||||
|
|
||||||
|
|
||||||
def _map_global_to_group_local_ranks(
|
def _map_global_to_group_local_ranks(
|
||||||
group_ranks: List[int], global_ranks: List[int]
|
group_ranks: List[int], global_ranks: List[int]
|
||||||
) -> 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]
|
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:
|
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)):
|
while not all(mooncake_ep.get_peer_state(backend, ranks)):
|
||||||
time.sleep(_PEER_STATE_POLL_INTERVAL_SEC)
|
time.sleep(_PEER_STATE_POLL_INTERVAL_SEC)
|
||||||
|
|
||||||
@@ -142,66 +368,71 @@ def _maybe_create_message_queue(group) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _refresh_ep_members() -> None:
|
def _try_recover_world(global_ranks: List[int]) -> bool:
|
||||||
from sglang.srt.layers.moe.token_dispatcher.mooncake import EPBuffer
|
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:
|
def try_recover_ranks(global_ranks: List[int]) -> bool:
|
||||||
from mooncake import ep as mooncake_ep
|
"""Recover ranks in WORLD and every launch-time parallel group."""
|
||||||
|
if not _try_recover_world(global_ranks):
|
||||||
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.
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# Recover the world backend first, then recover each derived process group
|
from mooncake import ep as mooncake_ep
|
||||||
# using ranks mapped into that group's local rank space.
|
|
||||||
mooncake_ep.recover_ranks(world_backend, global_ranks)
|
|
||||||
|
|
||||||
for group in _iter_live_parallel_groups():
|
for group in _iter_live_parallel_groups():
|
||||||
group_local_ranks = _map_global_to_group_local_ranks(group.ranks, global_ranks)
|
local_ranks = _map_global_to_group_local_ranks(group.ranks, global_ranks)
|
||||||
if not group_local_ranks:
|
if not local_ranks:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
device_backend = _get_process_group_backend(group.device_group, "cuda")
|
_wait_for_peer_state(mooncake_ep, group.device_group, local_ranks)
|
||||||
_wait_for_peer_state(mooncake_ep, device_backend, group_local_ranks)
|
mooncake_ep.recover_ranks(group.device_group, local_ranks)
|
||||||
mooncake_ep.recover_ranks(device_backend, group_local_ranks)
|
_wait_for_peer_state(mooncake_ep, group.cpu_group, local_ranks)
|
||||||
|
mooncake_ep.recover_ranks(group.cpu_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)
|
|
||||||
_maybe_create_message_queue(group)
|
_maybe_create_message_queue(group)
|
||||||
|
|
||||||
_refresh_ep_members()
|
_refresh_ep_members()
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def join_process_groups():
|
def _join_world_group() -> None:
|
||||||
from mooncake import ep as mooncake_ep
|
from mooncake import ep as mooncake_ep
|
||||||
|
|
||||||
def join_backend(label: str, backend) -> None:
|
mooncake_ep.join_group(torch.distributed.group.WORLD)
|
||||||
logger.info("Recovered rank joining Mooncake backend %s", label)
|
|
||||||
mooncake_ep.join_group(backend)
|
|
||||||
|
|
||||||
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():
|
for group in _iter_live_parallel_groups():
|
||||||
if group.world_size <= 1:
|
if group.world_size <= 1:
|
||||||
continue
|
continue
|
||||||
|
mooncake_ep.join_group(group.device_group)
|
||||||
join_backend(
|
mooncake_ep.join_group(group.cpu_group)
|
||||||
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"),
|
|
||||||
)
|
|
||||||
_maybe_create_message_queue(group)
|
_maybe_create_message_queue(group)
|
||||||
|
|
||||||
_refresh_ep_members()
|
_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 is None for RayEngine (uses Ray actors instead).
|
||||||
"""
|
"""
|
||||||
scheduler_procs = []
|
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
|
# Launch tensor parallel scheduler processes
|
||||||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||||
enable=server_args.enable_memory_saver
|
enable=server_args.enable_memory_saver
|
||||||
@@ -678,8 +681,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
def wait_for_ready():
|
def wait_for_ready():
|
||||||
infos = _wait_for_scheduler_ready(scheduler_pipe_readers, scheduler_procs)
|
infos = _wait_for_scheduler_ready(scheduler_pipe_readers, scheduler_procs)
|
||||||
scheduler_infos.extend(infos)
|
scheduler_infos.extend(infos)
|
||||||
# For dp_size > 1, collect child scheduler PIDs from the DP controller
|
if use_dp_controller:
|
||||||
if server_args.dp_size > 1:
|
|
||||||
for info in infos:
|
for info in infos:
|
||||||
if SCHEDULER_PIDS_ARG in info:
|
if SCHEDULER_PIDS_ARG in info:
|
||||||
all_child_pids.extend(info[SCHEDULER_PIDS_ARG])
|
all_child_pids.extend(info[SCHEDULER_PIDS_ARG])
|
||||||
@@ -833,8 +835,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
run_expert_backup_manager(server_args, port_args)
|
run_expert_backup_manager(server_args, port_args)
|
||||||
|
|
||||||
if server_args.node_rank >= 1:
|
if server_args.node_rank >= 1:
|
||||||
# In multi-node cases, non-zero rank nodes do not need to run tokenizer or detokenizer,
|
# Non-zero-rank nodes do not run tokenizer processes.
|
||||||
# so they can just wait here.
|
|
||||||
scheduler_init_result.wait_for_ready()
|
scheduler_init_result.wait_for_ready()
|
||||||
|
|
||||||
if os.getenv("SGLANG_BLOCK_NONZERO_RANK_CHILDREN") == "0":
|
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)
|
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:
|
def _anthropic_validation_message(raw_errors) -> str:
|
||||||
"""Render Pydantic-style errors for an Anthropic /v1/messages route.
|
"""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:
|
if server_args.checkpoint_engine_wait_weights_before_ready:
|
||||||
_wait_weights_ready()
|
_wait_weights_ready()
|
||||||
|
|
||||||
# Send a warmup request
|
# Joiner schedulers are served through the primary after adoption.
|
||||||
if not server_args.skip_server_warmup:
|
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):
|
if not execute_warmup_func(server_args):
|
||||||
return
|
return
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -52,6 +52,8 @@ class EPLBManager:
|
|||||||
self._server_args.eplb_rebalance_layers_per_chunk
|
self._server_args.eplb_rebalance_layers_per_chunk
|
||||||
)
|
)
|
||||||
self._rebalance_num_iterations = self._server_args.eplb_rebalance_num_iterations
|
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.
|
# Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented.
|
||||||
assert (
|
assert (
|
||||||
@@ -74,6 +76,11 @@ class EPLBManager:
|
|||||||
def reset_generator(self):
|
def reset_generator(self):
|
||||||
self._main_generator = self._entrypoint()
|
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
|
# can be more complex if needed
|
||||||
def _entrypoint(self):
|
def _entrypoint(self):
|
||||||
while True:
|
while True:
|
||||||
@@ -83,6 +90,15 @@ class EPLBManager:
|
|||||||
yield from self.rebalance()
|
yield from self.rebalance()
|
||||||
|
|
||||||
def rebalance(self):
|
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")
|
logger.info("[EPLBManager] rebalance start")
|
||||||
|
|
||||||
enable_timing = self._rebalance_layers_per_chunk is None
|
enable_timing = self._rebalance_layers_per_chunk is None
|
||||||
|
|||||||
@@ -331,7 +331,9 @@ class _SinglePassGatherer(ABC):
|
|||||||
return _SelectExpertsSinglePassGatherer(expert_location_metadata, rank)
|
return _SelectExpertsSinglePassGatherer(expert_location_metadata, rank)
|
||||||
elif server_args.deepep_mode == "low_latency":
|
elif server_args.deepep_mode == "low_latency":
|
||||||
return _DeepepLowLatencySinglePassGatherer(
|
return _DeepepLowLatencySinglePassGatherer(
|
||||||
expert_location_metadata, rank
|
expert_location_metadata,
|
||||||
|
rank,
|
||||||
|
elastic_ep_enabled=server_args.elastic_ep_backend is not None,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
@@ -574,13 +576,26 @@ class _DeepepNormalSinglePassGatherer(_LayerBasedCpuSinglePassGatherer):
|
|||||||
|
|
||||||
|
|
||||||
class _DeepepLowLatencySinglePassGatherer(_LayerBasedGpuSinglePassGatherer):
|
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)
|
super().__init__(*args, **kwargs, enable_global_physical_experts=False)
|
||||||
|
self._elastic_ep_enabled = elastic_ep_enabled
|
||||||
|
|
||||||
def on_deepep_dispatch_low_latency(
|
def on_deepep_dispatch_low_latency(
|
||||||
self, layer_idx: int, local_physical_count_of_layer: torch.Tensor
|
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
|
self._data[layer_idx, :] += local_physical_count_of_layer
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -32,6 +32,25 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
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
|
@dataclass
|
||||||
class ExpertLocationMetadata:
|
class ExpertLocationMetadata:
|
||||||
physical_to_logical_map: torch.Tensor # (layers, num_physical_experts)
|
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: torch.Tensor # (layers, num_logical_experts, X)
|
||||||
logical_to_all_physical_map_cpu: torch.Tensor # CPU copy for performance
|
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)
|
logical_to_all_physical_map_num_valid: torch.Tensor # (layers, num_logical_experts)
|
||||||
|
ep_size: int
|
||||||
# (layers, num_logical_experts)
|
# (layers, num_logical_experts)
|
||||||
logical_to_rank_dispatch_physical_map: Optional[torch.Tensor]
|
logical_to_rank_dispatch_physical_map: Optional[torch.Tensor]
|
||||||
|
|
||||||
@@ -62,11 +82,6 @@ class ExpertLocationMetadata:
|
|||||||
def num_logical_experts(self) -> int:
|
def num_logical_experts(self) -> int:
|
||||||
return self.logical_to_all_physical_map.shape[1]
|
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):
|
def __post_init__(self):
|
||||||
num_layers_0, num_physical_experts_0 = self.physical_to_logical_map.shape
|
num_layers_0, num_physical_experts_0 = self.physical_to_logical_map.shape
|
||||||
num_layers_1, num_logical_experts_0, num_physical_experts_1 = (
|
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_layers = model_config_for_expert_location.num_layers
|
||||||
num_logical_experts = model_config_for_expert_location.num_logical_experts
|
num_logical_experts = model_config_for_expert_location.num_logical_experts
|
||||||
|
|
||||||
|
base_num_physical_experts = common["base_num_physical_experts"]
|
||||||
physical_to_logical_map = (
|
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
|
% 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(
|
return ExpertLocationMetadata.init_by_mapping(
|
||||||
server_args,
|
server_args,
|
||||||
@@ -125,6 +146,15 @@ class ExpertLocationMetadata:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
model_config_for_expert_location = common["model_config_for_expert_location"]
|
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(
|
logical_to_all_physical_map = _compute_logical_to_all_physical_map(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
physical_to_logical_map=physical_to_logical_map,
|
physical_to_logical_map=physical_to_logical_map,
|
||||||
@@ -138,6 +168,7 @@ class ExpertLocationMetadata:
|
|||||||
ep_size=common["ep_size"],
|
ep_size=common["ep_size"],
|
||||||
physical_to_logical_map=physical_to_logical_map,
|
physical_to_logical_map=physical_to_logical_map,
|
||||||
logical_to_all_physical_map=logical_to_all_physical_map,
|
logical_to_all_physical_map=logical_to_all_physical_map,
|
||||||
|
moe_ep_rank=moe_ep_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -195,16 +226,33 @@ class ExpertLocationMetadata:
|
|||||||
if model_config_for_expert_location is None:
|
if model_config_for_expert_location is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
num_physical_experts = (
|
base_num_physical_experts = (
|
||||||
model_config_for_expert_location.num_logical_experts
|
model_config_for_expert_location.num_logical_experts
|
||||||
+ server_args.ep_num_redundant_experts
|
+ server_args.ep_num_redundant_experts
|
||||||
)
|
)
|
||||||
ep_size = server_args.ep_size
|
ep_size = server_args.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
|
assert num_physical_experts % ep_size == 0
|
||||||
num_local_physical_experts = num_physical_experts // ep_size
|
num_local_physical_experts = num_physical_experts // ep_size
|
||||||
|
|
||||||
return dict(
|
return dict(
|
||||||
model_config_for_expert_location=model_config_for_expert_location,
|
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_physical_experts=num_physical_experts,
|
||||||
num_local_physical_experts=num_local_physical_experts,
|
num_local_physical_experts=num_local_physical_experts,
|
||||||
ep_size=ep_size,
|
ep_size=ep_size,
|
||||||
@@ -216,6 +264,7 @@ class ExpertLocationMetadata:
|
|||||||
ep_size: int,
|
ep_size: int,
|
||||||
physical_to_logical_map: torch.Tensor,
|
physical_to_logical_map: torch.Tensor,
|
||||||
logical_to_all_physical_map: torch.Tensor,
|
logical_to_all_physical_map: torch.Tensor,
|
||||||
|
moe_ep_rank: Optional[int] = None,
|
||||||
):
|
):
|
||||||
_, num_physical_experts = physical_to_logical_map.shape
|
_, 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=logical_to_all_physical_map_padded,
|
||||||
logical_to_all_physical_map_cpu=logical_to_all_physical_map_padded.cpu(),
|
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,
|
logical_to_all_physical_map_num_valid=logical_to_all_physical_map_num_valid,
|
||||||
|
ep_size=ep_size,
|
||||||
logical_to_rank_dispatch_physical_map=(
|
logical_to_rank_dispatch_physical_map=(
|
||||||
compute_logical_to_rank_dispatch_physical_map(
|
compute_logical_to_rank_dispatch_physical_map(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
logical_to_all_physical_map=logical_to_all_physical_map,
|
logical_to_all_physical_map=logical_to_all_physical_map,
|
||||||
ep_size=ep_size,
|
ep_size=ep_size,
|
||||||
num_physical_experts=num_physical_experts,
|
num_physical_experts=num_physical_experts,
|
||||||
# TODO improve when we have real EP rank
|
ep_rank=(
|
||||||
ep_rank=torch.distributed.get_rank() % ep_size,
|
moe_ep_rank
|
||||||
|
if moe_ep_rank is not None
|
||||||
|
else torch.distributed.get_rank() % ep_size
|
||||||
|
),
|
||||||
)
|
)
|
||||||
if server_args.ep_dispatch_algorithm == "static"
|
if server_args.ep_dispatch_algorithm == "static"
|
||||||
else None
|
else None
|
||||||
@@ -425,58 +478,57 @@ def get_global_expert_location_metadata():
|
|||||||
return get_resources().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
|
from sglang.srt.runtime_context import get_resources
|
||||||
|
|
||||||
resources = get_resources()
|
resources = get_resources()
|
||||||
|
if not allow_overwrite:
|
||||||
assert resources.expert_location_metadata is None
|
assert resources.expert_location_metadata is None
|
||||||
resources.expert_location_metadata = value
|
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(
|
def broadcast_global_expert_location_metadata(
|
||||||
src_rank: int = 0, group: Optional[torch.distributed.ProcessGroup] = None
|
model_config: ModelConfig,
|
||||||
):
|
moe_ep_rank: int,
|
||||||
"""Broadcast the global ExpertLocationMetadata from src_rank to all ranks.
|
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
|
server_args = get_server_args()
|
||||||
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.
|
|
||||||
"""
|
|
||||||
metadata = get_global_expert_location_metadata()
|
metadata = get_global_expert_location_metadata()
|
||||||
assert metadata is not None
|
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.physical_to_logical_map = metadata.physical_to_logical_map.contiguous()
|
||||||
metadata.logical_to_all_physical_map = (
|
torch.distributed.broadcast(
|
||||||
metadata.logical_to_all_physical_map.contiguous()
|
metadata.physical_to_logical_map, src=src_rank, group=group
|
||||||
)
|
)
|
||||||
metadata.logical_to_all_physical_map_num_valid = (
|
metadata = ExpertLocationMetadata.init_by_mapping(
|
||||||
metadata.logical_to_all_physical_map_num_valid.contiguous()
|
server_args,
|
||||||
)
|
model_config,
|
||||||
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.physical_to_logical_map,
|
metadata.physical_to_logical_map,
|
||||||
metadata.logical_to_all_physical_map,
|
moe_ep_rank=moe_ep_rank,
|
||||||
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()
|
|
||||||
)
|
)
|
||||||
|
set_global_expert_location_metadata(metadata, allow_overwrite=True)
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
|
||||||
def _compute_logical_to_all_physical_map(
|
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
|
# Replace by the physical expert on local GPU or node if possible
|
||||||
if moe_ep_rank is not None:
|
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
|
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_node_physical_experts = (
|
||||||
num_local_gpu_physical_experts * num_gpus_per_node
|
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 layer_id in range(num_layers):
|
||||||
for logical_expert_id in range(num_logical_experts):
|
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()
|
logical_to_all_physical_map = logical_to_all_physical_map.cpu()
|
||||||
|
|
||||||
num_local_gpu_physical_experts = num_physical_experts // ep_size
|
num_local_gpu_physical_experts = num_physical_experts // ep_size
|
||||||
num_gpus_per_node = server_args.ep_size // server_args.nnodes
|
prefer_same_node = _prefer_same_node_experts(server_args)
|
||||||
num_local_node_physical_experts = num_local_gpu_physical_experts * num_gpus_per_node
|
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
|
num_layers, num_logical_experts, _ = logical_to_all_physical_map.shape
|
||||||
dtype = logical_to_all_physical_map.dtype
|
dtype = logical_to_all_physical_map.dtype
|
||||||
|
|
||||||
@@ -633,8 +697,8 @@ def _find_nearest_expert(
|
|||||||
candidate_physical_expert_ids: List[int],
|
candidate_physical_expert_ids: List[int],
|
||||||
num_local_gpu_physical_experts: int,
|
num_local_gpu_physical_experts: int,
|
||||||
moe_ep_rank: int,
|
moe_ep_rank: int,
|
||||||
num_gpus_per_node: int,
|
num_gpus_per_node: Optional[int],
|
||||||
num_local_node_physical_experts: int,
|
num_local_node_physical_experts: Optional[int],
|
||||||
) -> int:
|
) -> int:
|
||||||
# 1. If only one candidate, return it directly
|
# 1. If only one candidate, return it directly
|
||||||
if len(candidate_physical_expert_ids) == 1:
|
if len(candidate_physical_expert_ids) == 1:
|
||||||
@@ -652,7 +716,8 @@ def _find_nearest_expert(
|
|||||||
if len(same_gpu_physical_expert_ids) > 0:
|
if len(same_gpu_physical_expert_ids) > 0:
|
||||||
return same_gpu_physical_expert_ids[0]
|
return same_gpu_physical_expert_ids[0]
|
||||||
|
|
||||||
# 3. Otherwise, prefer same-node experts
|
# 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
|
node_rank = moe_ep_rank // num_gpus_per_node
|
||||||
same_node_physical_expert_ids = [
|
same_node_physical_expert_ids = [
|
||||||
physical_expert_id
|
physical_expert_id
|
||||||
@@ -662,7 +727,7 @@ def _find_nearest_expert(
|
|||||||
)
|
)
|
||||||
== node_rank
|
== node_rank
|
||||||
]
|
]
|
||||||
if len(same_node_physical_expert_ids) > 0:
|
if 0 < len(same_node_physical_expert_ids) < len(candidate_physical_expert_ids):
|
||||||
return same_node_physical_expert_ids[0]
|
return same_node_physical_expert_ids[0]
|
||||||
|
|
||||||
# 4. At last, leave it as -1 to indicate not found.
|
# 4. At last, leave it as -1 to indicate not found.
|
||||||
|
|||||||
@@ -5,9 +5,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import DpPaddingMode
|
||||||
DpPaddingMode,
|
|
||||||
)
|
|
||||||
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
||||||
is_in_breakable_cuda_graph,
|
is_in_breakable_cuda_graph,
|
||||||
)
|
)
|
||||||
@@ -164,6 +162,9 @@ def cal_padded_tokens(forward_batch: "ForwardBatch"):
|
|||||||
cp_align_size = get_cp_padding_align_size()
|
cp_align_size = get_cp_padding_align_size()
|
||||||
for i in range(sync_group_size):
|
for i in range(sync_group_size):
|
||||||
global_num_tokens[i] = ceil_align(global_num_tokens[i], cp_align_size)
|
global_num_tokens[i] = ceil_align(global_num_tokens[i], cp_align_size)
|
||||||
|
# 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(
|
dp_padding_mode = DpPaddingMode.get_dp_padding_mode(
|
||||||
forward_batch.is_extend_in_batch, global_num_tokens
|
forward_batch.is_extend_in_batch, global_num_tokens
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -39,6 +39,29 @@ if TYPE_CHECKING:
|
|||||||
_ATTN_DP_RANK: Optional[int] = None
|
_ATTN_DP_RANK: Optional[int] = None
|
||||||
_ATTN_DP_SIZE: 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()
|
_is_hip = is_hip()
|
||||||
_USE_ROCM700A_WA = _is_hip and get_bool_env_var("SGLANG_USE_ROCM700A")
|
_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
|
enable_dp_attention = server_args.enable_dp_attention
|
||||||
dp_size = server_args.dp_size
|
dp_size = server_args.dp_size
|
||||||
moe_dense_tp_size = server_args.moe_dense_tp_size
|
|
||||||
attn_cp_size = server_args.attn_cp_size
|
attn_cp_size = server_args.attn_cp_size
|
||||||
|
|
||||||
dp.enabled = enable_dp_attention
|
dp.enabled = enable_dp_attention
|
||||||
@@ -300,6 +322,11 @@ def initialize_dp_attention(
|
|||||||
)
|
)
|
||||||
_ATTN_DP_SIZE = dp_size if enable_dp_attention else 1
|
_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(
|
_DpGatheredBufferWrapper.set_metadata(
|
||||||
hidden_size=model_config.hidden_size,
|
hidden_size=model_config.hidden_size,
|
||||||
dtype=model_config.dtype,
|
dtype=model_config.dtype,
|
||||||
@@ -408,6 +435,13 @@ 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.
|
# Input IDs are in int 32. We should use inplace_all_reduce for local case because of custom all reduce.
|
||||||
|
if world_dp_gather_enabled():
|
||||||
|
torch.distributed.all_reduce(
|
||||||
|
global_tokens,
|
||||||
|
op=torch.distributed.ReduceOp.SUM,
|
||||||
|
group=torch.distributed.group.WORLD,
|
||||||
|
)
|
||||||
|
else:
|
||||||
NUM_GPUS_PER_NODE = 8
|
NUM_GPUS_PER_NODE = 8
|
||||||
if (
|
if (
|
||||||
not local_tokens.dtype.is_floating_point
|
not local_tokens.dtype.is_floating_point
|
||||||
@@ -427,7 +461,16 @@ def _dp_gather_via_all_gather(
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
is_partial: bool,
|
is_partial: bool,
|
||||||
):
|
):
|
||||||
|
use_world = world_dp_gather_enabled()
|
||||||
|
|
||||||
if get_attn_tensor_model_parallel_world_size() == 1:
|
if get_attn_tensor_model_parallel_world_size() == 1:
|
||||||
|
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)
|
get_tp_group().all_gather_into_tensor(global_tokens, local_tokens)
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -438,6 +481,13 @@ def _dp_gather_via_all_gather(
|
|||||||
get_attn_tensor_model_parallel_world_size()
|
get_attn_tensor_model_parallel_world_size()
|
||||||
)[get_attn_tensor_model_parallel_rank()]
|
)[get_attn_tensor_model_parallel_rank()]
|
||||||
get_attn_tp_group().reduce_scatter_tensor(scattered_local_tokens, local_tokens)
|
get_attn_tp_group().reduce_scatter_tensor(scattered_local_tokens, 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)
|
get_tp_group().all_gather_into_tensor(global_tokens, scattered_local_tokens)
|
||||||
|
|
||||||
|
|
||||||
@@ -461,6 +511,7 @@ def is_dp_gatherv_active() -> bool:
|
|||||||
dp_reduce_scatter_tensor) consistent."""
|
dp_reduce_scatter_tensor) consistent."""
|
||||||
return (
|
return (
|
||||||
_USE_DP_GATHERV
|
_USE_DP_GATHERV
|
||||||
|
and not world_dp_gather_enabled()
|
||||||
and get_attn_tensor_model_parallel_world_size() == 1
|
and get_attn_tensor_model_parallel_world_size() == 1
|
||||||
and get_tensor_model_parallel_world_size() == get_attention_dp_size()
|
and get_tensor_model_parallel_world_size() == get_attention_dp_size()
|
||||||
and not _DpGatheredBufferWrapper.is_dp_max_padding()
|
and not _DpGatheredBufferWrapper.is_dp_max_padding()
|
||||||
@@ -541,7 +592,10 @@ def _dp_gather(
|
|||||||
global_tokens, local_tokens, forward_batch, is_partial, _gatherv_sizes
|
global_tokens, local_tokens, forward_batch, is_partial, _gatherv_sizes
|
||||||
)
|
)
|
||||||
return
|
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(
|
_dp_gather_via_all_gather(
|
||||||
global_tokens, local_tokens, forward_batch, is_partial
|
global_tokens, local_tokens, forward_batch, is_partial
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -229,9 +229,19 @@ class FusedMoE(torch.nn.Module):
|
|||||||
else:
|
else:
|
||||||
num_shared_slots = num_fused_shared_experts
|
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_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.num_local_experts = self._num_local_routed + num_fused_shared_experts
|
||||||
self._has_fused_shared = num_fused_shared_experts > 0
|
self._has_fused_shared = num_fused_shared_experts > 0
|
||||||
self._pending_fp8_shared_weights: dict[tuple[int, str], torch.Tensor] = {}
|
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)
|
expert_data.copy_(loaded_weight)
|
||||||
|
|
||||||
def _map_global_expert_id_to_local_expert_id(self, expert_id: int) -> int:
|
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
|
end_idx = start_idx + self._num_local_routed
|
||||||
if start_idx <= expert_id < end_idx:
|
if start_idx <= expert_id < end_idx:
|
||||||
return expert_id - start_idx
|
return expert_id - start_idx
|
||||||
|
|||||||
@@ -56,10 +56,46 @@ class NixlEPBuffer:
|
|||||||
num_max_dispatch_tokens_per_rank=None,
|
num_max_dispatch_tokens_per_rank=None,
|
||||||
num_experts=None,
|
num_experts=None,
|
||||||
num_local_experts=None,
|
num_local_experts=None,
|
||||||
|
connected_ep_size=None,
|
||||||
|
scale_to=None,
|
||||||
|
dispatch_ep_size=None,
|
||||||
)
|
)
|
||||||
buffers["nixl_ep_state"] = state
|
buffers["nixl_ep_state"] = state
|
||||||
return 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
|
@classmethod
|
||||||
def get_nixl_buffer(
|
def get_nixl_buffer(
|
||||||
cls,
|
cls,
|
||||||
@@ -72,6 +108,12 @@ class NixlEPBuffer:
|
|||||||
):
|
):
|
||||||
state = cls._state()
|
state = cls._state()
|
||||||
if state.buffer is not None:
|
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
|
return state.buffer
|
||||||
|
|
||||||
state.hidden_size = hidden_size
|
state.hidden_size = hidden_size
|
||||||
@@ -79,23 +121,31 @@ class NixlEPBuffer:
|
|||||||
state.num_experts = num_experts
|
state.num_experts = num_experts
|
||||||
state.num_local_experts = num_local_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
|
num_rdma_bytes = 0
|
||||||
if deepep_mode.enable_normal():
|
if deepep_mode.enable_normal():
|
||||||
raise NotImplementedError("Normal mode is not supported for Nixl EP yet.")
|
raise NotImplementedError("Normal mode is not supported for Nixl EP yet.")
|
||||||
if deepep_mode.enable_low_latency():
|
if deepep_mode.enable_low_latency():
|
||||||
assert num_max_dispatch_tokens_per_rank != -1
|
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_rdma_bytes = Buffer.get_rdma_size_hint(
|
||||||
num_max_dispatch_tokens_per_rank,
|
num_max_dispatch_tokens_per_rank,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
group.size(),
|
nixl_max_ranks,
|
||||||
num_experts,
|
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()
|
tcp_store = get_global_tcp_store()
|
||||||
if tcp_store is None:
|
if tcp_store is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -104,23 +154,28 @@ class NixlEPBuffer:
|
|||||||
)
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Using NIXL EP (world_size={world_size}, rank={rank}, "
|
f"Using NIXL EP (world_size={world_size}, max_ep_size={max_ep_size}, "
|
||||||
f"num_experts={state.num_experts}, num_experts_per_rank={state.num_local_experts}) "
|
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(
|
state.buffer = Buffer(
|
||||||
rank=rank,
|
rank=global_rank,
|
||||||
tcp_store_group=tcp_store,
|
tcp_store_group=tcp_store,
|
||||||
)
|
)
|
||||||
|
|
||||||
state.buffer.update_memory_buffers(
|
state.buffer.update_memory_buffers(
|
||||||
num_ranks=world_size,
|
num_ranks=nixl_max_ranks,
|
||||||
num_experts_per_rank=state.num_local_experts,
|
num_experts_per_rank=state.num_local_experts,
|
||||||
num_rdma_bytes=num_rdma_bytes,
|
num_rdma_bytes=num_rdma_bytes,
|
||||||
)
|
)
|
||||||
all_ranks = list(range(world_size))
|
initial_ep_size = offset + world_size
|
||||||
state.buffer.connect_ranks(all_ranks)
|
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
|
return state.buffer
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -170,8 +225,12 @@ class _NixlEPDispatcherImplBase:
|
|||||||
self.active_ranks = (
|
self.active_ranks = (
|
||||||
elastic_state.active_ranks if elastic_state is not None else None
|
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 = (
|
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
|
if self.active_ranks is not None
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
@@ -232,14 +291,21 @@ class _NixlEPDispatcherImpl(_NixlEPDispatcherImplBase):
|
|||||||
buffer = self._get_buffer()
|
buffer = self._get_buffer()
|
||||||
topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids
|
topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids
|
||||||
topk_ids = topk_ids.to(torch.int64)
|
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 = (
|
expected_m = (
|
||||||
hidden_states.shape[0] * buffer.group_size * topk_ids.shape[1]
|
hidden_states.shape[0] * dispatch_ep_size * topk_ids.shape[1]
|
||||||
+ self.num_experts
|
+ num_dispatch_experts
|
||||||
) // self.num_experts
|
) // num_dispatch_experts
|
||||||
|
|
||||||
hidden_states, masked_m, event, hook = self._dispatch_core(
|
hidden_states, masked_m, event, hook = self._dispatch_core(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
topk_ids,
|
topk_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
hidden_states,
|
hidden_states,
|
||||||
topk_ids,
|
topk_ids,
|
||||||
@@ -289,12 +355,17 @@ class _NixlEPDispatcherImpl(_NixlEPDispatcherImplBase):
|
|||||||
use_fp8 = not envs.SGLANG_NIXL_EP_BF16_DISPATCH.get()
|
use_fp8 = not envs.SGLANG_NIXL_EP_BF16_DISPATCH.get()
|
||||||
|
|
||||||
buffer = self._get_buffer()
|
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 = (
|
packed_recv_hidden, self.packed_recv_count, self.handle, event, hook = (
|
||||||
buffer.dispatch(
|
buffer.dispatch(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
topk_idx,
|
topk_idx,
|
||||||
self.num_max_dispatch_tokens_per_rank,
|
self.num_max_dispatch_tokens_per_rank,
|
||||||
self.num_experts,
|
nixl_num_experts,
|
||||||
use_fp8=use_fp8,
|
use_fp8=use_fp8,
|
||||||
async_finish=not self.return_recv_hook,
|
async_finish=not self.return_recv_hook,
|
||||||
return_recv_hook=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:
|
if self._mask_buffer is not None:
|
||||||
buffer.query_mask_buffer(self._mask_buffer)
|
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
|
self.packed_recv_count = self.handle = None
|
||||||
return combined_hidden_states, event, hook
|
return combined_hidden_states, event, hook
|
||||||
|
|||||||
@@ -2,8 +2,11 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import copy
|
import copy
|
||||||
|
import logging
|
||||||
from typing import Callable, Generic, List, Optional, TypeVar
|
from typing import Callable, Generic, List, Optional, TypeVar
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
T = TypeVar("T")
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
@@ -30,6 +33,7 @@ class FanOutCommunicator(Generic[T]):
|
|||||||
self._mode = mode
|
self._mode = mode
|
||||||
self._result_event: Optional[asyncio.Event] = None
|
self._result_event: Optional[asyncio.Event] = None
|
||||||
self._result_values: Optional[List[T]] = None
|
self._result_values: Optional[List[T]] = None
|
||||||
|
self._result_fan_out: Optional[int] = None
|
||||||
self._queueing_lock = asyncio.Lock()
|
self._queueing_lock = asyncio.Lock()
|
||||||
|
|
||||||
assert mode in ["queueing", "watching"]
|
assert mode in ["queueing", "watching"]
|
||||||
@@ -45,10 +49,11 @@ class FanOutCommunicator(Generic[T]):
|
|||||||
|
|
||||||
self._result_event = asyncio.Event()
|
self._result_event = asyncio.Event()
|
||||||
self._result_values = []
|
self._result_values = []
|
||||||
|
self._result_fan_out = self._fan_out
|
||||||
await self._result_event.wait()
|
await self._result_event.wait()
|
||||||
result_values = self._result_values
|
result_values = self._result_values
|
||||||
self._result_event = self._result_values = None
|
self._result_event = self._result_values = None
|
||||||
|
self._result_fan_out = None
|
||||||
return result_values
|
return result_values
|
||||||
|
|
||||||
async def watching_call(self, obj):
|
async def watching_call(self, obj):
|
||||||
@@ -56,6 +61,7 @@ class FanOutCommunicator(Generic[T]):
|
|||||||
assert self._result_values is None
|
assert self._result_values is None
|
||||||
self._result_values = []
|
self._result_values = []
|
||||||
self._result_event = asyncio.Event()
|
self._result_event = asyncio.Event()
|
||||||
|
self._result_fan_out = self._fan_out
|
||||||
|
|
||||||
if obj is not None:
|
if obj is not None:
|
||||||
self._send(obj)
|
self._send(obj)
|
||||||
@@ -69,6 +75,7 @@ class FanOutCommunicator(Generic[T]):
|
|||||||
result_values = copy.deepcopy(values)
|
result_values = copy.deepcopy(values)
|
||||||
if self._result_event is event:
|
if self._result_event is event:
|
||||||
self._result_event = self._result_values = None
|
self._result_event = self._result_values = None
|
||||||
|
self._result_fan_out = None
|
||||||
return result_values
|
return result_values
|
||||||
|
|
||||||
async def __call__(self, obj):
|
async def __call__(self, obj):
|
||||||
@@ -77,9 +84,22 @@ class FanOutCommunicator(Generic[T]):
|
|||||||
else:
|
else:
|
||||||
return await self.watching_call(obj)
|
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):
|
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)
|
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()
|
self._result_event.set()
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ from sglang.srt.managers.io_struct import (
|
|||||||
BatchTokenizedEmbeddingReqInput,
|
BatchTokenizedEmbeddingReqInput,
|
||||||
BatchTokenizedGenerateReqInput,
|
BatchTokenizedGenerateReqInput,
|
||||||
BlockReqInput,
|
BlockReqInput,
|
||||||
|
ElasticScaleUpdateReq,
|
||||||
ProfileReq,
|
ProfileReq,
|
||||||
TokenizedEmbeddingReqInput,
|
TokenizedEmbeddingReqInput,
|
||||||
TokenizedGenerateReqInput,
|
TokenizedGenerateReqInput,
|
||||||
@@ -164,7 +165,17 @@ class DataParallelController:
|
|||||||
LoadBalanceMethod.TOTAL_TOKENS,
|
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.dp_budget = DPBudget(server_args.dp_size)
|
||||||
self.load_snapshot_reader = create_load_snapshot_reader(
|
self.load_snapshot_reader = create_load_snapshot_reader(
|
||||||
server_args,
|
server_args,
|
||||||
@@ -178,8 +189,10 @@ class DataParallelController:
|
|||||||
|
|
||||||
# Launch data parallel workers
|
# Launch data parallel workers
|
||||||
self.scheduler_procs = []
|
self.scheduler_procs = []
|
||||||
self.workers: List[zmq.Socket] = [None] * server_args.dp_size
|
self.workers: List[Optional[zmq.Socket]] = [None] * self.max_dp_size
|
||||||
self.status: List[bool] = [True] * server_args.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:
|
if server_args.enable_dp_attention:
|
||||||
self.launch_dp_attention_schedulers(server_args, port_args)
|
self.launch_dp_attention_schedulers(server_args, port_args)
|
||||||
@@ -208,16 +221,80 @@ class DataParallelController:
|
|||||||
|
|
||||||
def send_to_all_workers(self, obj):
|
def send_to_all_workers(self, obj):
|
||||||
for i, worker in enumerate(self.workers):
|
for i, worker in enumerate(self.workers):
|
||||||
if self.status[i]:
|
if worker is not None and self.status[i]:
|
||||||
sock_send(worker, obj)
|
sock_send(worker, obj)
|
||||||
|
|
||||||
def send_control_message(self, obj):
|
def send_control_message(self, obj):
|
||||||
# Send control messages to first worker of tp group
|
for i in self._active_workers[:: self.control_message_step]:
|
||||||
for worker in self.workers[:: self.control_message_step]:
|
worker = self.workers[i]
|
||||||
|
if worker is not None:
|
||||||
sock_send(worker, obj)
|
sock_send(worker, obj)
|
||||||
|
|
||||||
def update_active_ranks(self, ranks: ActiveRanksOutput):
|
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):
|
def refresh_load_budget(self):
|
||||||
# Throttle to at most once per 20ms. When a burst of requests
|
# Throttle to at most once per 20ms. When a burst of requests
|
||||||
@@ -272,6 +349,12 @@ class DataParallelController:
|
|||||||
(BlockReqInput, self.send_to_all_workers),
|
(BlockReqInput, self.send_to_all_workers),
|
||||||
(ProfileReq, self.send_to_all_workers),
|
(ProfileReq, self.send_to_all_workers),
|
||||||
(ActiveRanksOutput, self.update_active_ranks),
|
(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)
|
self._request_dispatcher.add_fallback_fn(self.send_control_message)
|
||||||
@@ -355,8 +438,8 @@ class DataParallelController:
|
|||||||
Returns:
|
Returns:
|
||||||
List of worker ports (same on all nodes after broadcast).
|
List of worker ports (same on all nodes after broadcast).
|
||||||
"""
|
"""
|
||||||
# Determine the endpoint for inter-node communication
|
is_joiner = server_args.is_ep_scale_joiner
|
||||||
if server_args.dist_init_addr is None:
|
if server_args.dist_init_addr is None or is_joiner:
|
||||||
na = NetworkAddress(
|
na = NetworkAddress(
|
||||||
server_args.host or "127.0.0.1",
|
server_args.host or "127.0.0.1",
|
||||||
server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||||
@@ -411,11 +494,10 @@ class DataParallelController:
|
|||||||
).start()
|
).start()
|
||||||
|
|
||||||
def _reply_ports_as_server(self, rep_socket: zmq.Socket, worker_ports: List[int]):
|
def _reply_ports_as_server(self, rep_socket: zmq.Socket, worker_ports: List[int]):
|
||||||
"""
|
"""Background thread: serve the pre-bound worker-port list to
|
||||||
Runs as a background thread to broadcast worker ports for recovered EP ranks
|
late-arriving elastic joiners. Publishes port numbers only; the primary
|
||||||
"""
|
keeps ownership of every socket."""
|
||||||
while True:
|
while True:
|
||||||
# Wait for client handshake
|
|
||||||
try:
|
try:
|
||||||
client_rank = sock_recv(rep_socket)
|
client_rank = sock_recv(rep_socket)
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -453,6 +535,12 @@ class DataParallelController:
|
|||||||
finally:
|
finally:
|
||||||
req_socket.close()
|
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(
|
def launch_dp_attention_schedulers(
|
||||||
self, server_args: ServerArgs, port_args: PortArgs
|
self, server_args: ServerArgs, port_args: PortArgs
|
||||||
):
|
):
|
||||||
@@ -461,25 +549,42 @@ class DataParallelController:
|
|||||||
else:
|
else:
|
||||||
bind_host = NetworkAddress.parse(server_args.dist_init_addr).host
|
bind_host = NetworkAddress.parse(server_args.dist_init_addr).host
|
||||||
|
|
||||||
# Pre-allocate worker ports on node 0 to avoid conflicts
|
|
||||||
worker_ports = []
|
worker_ports = []
|
||||||
if server_args.node_rank == 0:
|
if server_args.is_ep_scale_joiner:
|
||||||
for dp_rank in range(server_args.dp_size):
|
# 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(
|
worker_port, worker_socket = get_zmq_socket_on_host(
|
||||||
self.context, zmq.PUSH, host=bind_host
|
self.context, zmq.PUSH, host=bind_host
|
||||||
)
|
)
|
||||||
worker_ports.append(worker_port)
|
worker_ports.append(worker_port)
|
||||||
self.workers[dp_rank] = worker_socket
|
self.workers[slot] = worker_socket
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Assigned port %s to worker %s on host %s",
|
"Assigned port %s to worker slot %s on host %s",
|
||||||
worker_port,
|
worker_port,
|
||||||
dp_rank,
|
slot,
|
||||||
bind_host,
|
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(
|
self.launch_tensor_parallel_group(
|
||||||
server_args, port_args, 0, None, broadcasted_ports
|
server_args, port_args, 0, None, broadcasted_ports
|
||||||
)
|
)
|
||||||
@@ -510,6 +615,11 @@ class DataParallelController:
|
|||||||
|
|
||||||
nnodes_per_tp_group = nnodes_per_pp_rank
|
nnodes_per_tp_group = nnodes_per_pp_rank
|
||||||
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
||||||
|
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_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),
|
||||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
||||||
@@ -534,6 +644,16 @@ class DataParallelController:
|
|||||||
rank_port_args = PortArgs.init_new(
|
rank_port_args = PortArgs.init_new(
|
||||||
server_args, dp_rank, worker_ports
|
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,
|
# Data parallelism reuses the tensor parallelism group,
|
||||||
# so all dp ranks should use the same nccl port.
|
# so all dp ranks should use the same nccl port.
|
||||||
rank_port_args.nccl_port = port_args.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:
|
with self.env_lock, maybe_reindex_device_id(gpu_id) as gpu_id:
|
||||||
proc = mp.Process(
|
proc = mp.Process(
|
||||||
target=self.run_scheduler_process_func,
|
target=self.run_scheduler_process_func,
|
||||||
@@ -584,6 +710,9 @@ class DataParallelController:
|
|||||||
pp_rank,
|
pp_rank,
|
||||||
dp_rank,
|
dp_rank,
|
||||||
writer,
|
writer,
|
||||||
|
display_tp_rank,
|
||||||
|
display_dp_rank,
|
||||||
|
display_moe_ep_rank,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
with (
|
with (
|
||||||
@@ -604,8 +733,16 @@ class DataParallelController:
|
|||||||
|
|
||||||
def maybe_external_dp_rank_routing(self, req: Req):
|
def maybe_external_dp_rank_routing(self, req: Req):
|
||||||
if req.routed_dp_rank is not None:
|
if req.routed_dp_rank is not None:
|
||||||
logger.debug(f"Direct routing to DP rank {req.routed_dp_rank}")
|
rank = req.routed_dp_rank
|
||||||
sock_send(self.workers[req.routed_dp_rank], req)
|
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 True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -613,16 +750,21 @@ class DataParallelController:
|
|||||||
if self.maybe_external_dp_rank_routing(req):
|
if self.maybe_external_dp_rank_routing(req):
|
||||||
return
|
return
|
||||||
|
|
||||||
while True:
|
active = self._active_workers
|
||||||
if self.status[self.round_robin_counter]:
|
if not active:
|
||||||
logger.debug(f"Choose worker {self.round_robin_counter}")
|
raise RuntimeError("No active DP workers are available for routing.")
|
||||||
sock_send(self.workers[self.round_robin_counter], req)
|
attempts = 0
|
||||||
self.round_robin_counter = (self.round_robin_counter + 1) % len(
|
while attempts < len(active):
|
||||||
self.workers
|
slot = active[self.round_robin_counter % len(active)]
|
||||||
)
|
self.round_robin_counter = (self.round_robin_counter + 1) % len(active)
|
||||||
break
|
if self.status[slot]:
|
||||||
self.round_robin_counter = (self.round_robin_counter + 1) % len(
|
logger.debug(f"Choose worker {slot}")
|
||||||
self.workers
|
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):
|
def follow_bootstrap_room_scheduler(self, req: Req):
|
||||||
@@ -702,7 +844,8 @@ def run_data_parallel_controller_process(
|
|||||||
SCHEDULER_PIDS_ARG: scheduler_pids,
|
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()
|
controller.event_loop()
|
||||||
for proc in controller.scheduler_procs:
|
for proc in controller.scheduler_procs:
|
||||||
proc.join()
|
proc.join()
|
||||||
|
|||||||
@@ -1809,6 +1809,31 @@ class ActiveRanksOutput(BaseReq, kw_only=True):
|
|||||||
status: List[bool]
|
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):
|
class GetInternalStateReq(BaseReq, kw_only=True):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|||||||
@@ -127,6 +127,8 @@ from sglang.srt.managers.io_struct import (
|
|||||||
ResumeMemoryOccupationReqInput,
|
ResumeMemoryOccupationReqInput,
|
||||||
RpcReqInput,
|
RpcReqInput,
|
||||||
RpcReqOutput,
|
RpcReqOutput,
|
||||||
|
ScaleElasticEPReqInput,
|
||||||
|
ScaleElasticEPReqOutput,
|
||||||
SendWeightsToRemoteInstanceReqInput,
|
SendWeightsToRemoteInstanceReqInput,
|
||||||
SendWeightsToRemoteInstanceReqOutput,
|
SendWeightsToRemoteInstanceReqOutput,
|
||||||
SetInternalStateReq,
|
SetInternalStateReq,
|
||||||
@@ -973,6 +975,7 @@ class Scheduler(
|
|||||||
is_fully_idle=self.is_fully_idle,
|
is_fully_idle=self.is_fully_idle,
|
||||||
ipc_channels=self.ipc_channels,
|
ipc_channels=self.ipc_channels,
|
||||||
)
|
)
|
||||||
|
self._last_logged_elastic_radix_namespace: Optional[str] = None
|
||||||
self.session_controller = SessionController(self.tree_cache)
|
self.session_controller = SessionController(self.tree_cache)
|
||||||
self.forward_sleep_time = None
|
self.forward_sleep_time = None
|
||||||
self._engine_paused = False
|
self._engine_paused = False
|
||||||
@@ -1408,6 +1411,7 @@ class Scheduler(
|
|||||||
(PauseGenerationReqInput, self.pause_generation),
|
(PauseGenerationReqInput, self.pause_generation),
|
||||||
(ContinueGenerationReqInput, self.continue_generation),
|
(ContinueGenerationReqInput, self.continue_generation),
|
||||||
(ConfigureLoggingReq, self.configure_logging),
|
(ConfigureLoggingReq, self.configure_logging),
|
||||||
|
(ScaleElasticEPReqInput, self.handle_scale_elastic_ep),
|
||||||
(DumperControlReqInput, self.handle_dumper_control),
|
(DumperControlReqInput, self.handle_dumper_control),
|
||||||
(AddExternalCorpusReqInput, self.add_external_corpus),
|
(AddExternalCorpusReqInput, self.add_external_corpus),
|
||||||
(
|
(
|
||||||
@@ -1994,6 +1998,33 @@ class Scheduler(
|
|||||||
mm.mrope_positions = mrope_positions
|
mm.mrope_positions = mrope_positions
|
||||||
mm.mrope_position_delta = mrope_position_delta
|
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:
|
def _maybe_clear_mm_inputs(self, batch: ScheduleBatch) -> None:
|
||||||
for req in batch.reqs:
|
for req in batch.reqs:
|
||||||
if not req.finished() or not (mm_inputs := req.multimodal_inputs):
|
if not req.finished() or not (mm_inputs := req.multimodal_inputs):
|
||||||
@@ -2134,6 +2165,8 @@ class Scheduler(
|
|||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
self._maybe_namespace_elastic_radix_cache(req)
|
||||||
|
|
||||||
if self.spec_algorithm.is_dflash_family():
|
if self.spec_algorithm.is_dflash_family():
|
||||||
error_msg = validate_dflash_request(req, self.enable_overlap)
|
error_msg = validate_dflash_request(req, self.enable_overlap)
|
||||||
if error_msg is not None:
|
if error_msg is not None:
|
||||||
@@ -2471,6 +2504,7 @@ class Scheduler(
|
|||||||
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
|
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
|
||||||
)
|
)
|
||||||
req.tokenizer = self.tokenizer
|
req.tokenizer = self.tokenizer
|
||||||
|
self._maybe_namespace_elastic_radix_cache(req)
|
||||||
|
|
||||||
# Handle multimodal inputs
|
# Handle multimodal inputs
|
||||||
if recv_req.mm_inputs is not None:
|
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
|
self.enable_dp_attention and self.server_args.elastic_ep_backend is not None
|
||||||
):
|
):
|
||||||
return
|
return
|
||||||
# Get the tensors indicating rank activeness
|
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||||
tp_active_ranks = self.tp_group.active_ranks.detach().cpu().numpy()
|
|
||||||
tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy()
|
inst = ElasticEPStateManager.instance()
|
||||||
tp_active_ranks &= tp_active_ranks_cpu
|
if inst is not None and inst.active_ranks_cpu is not None:
|
||||||
dp_active_ranks = tp_active_ranks.reshape(self.ps.dp_size, -1).prod(axis=1)
|
|
||||||
self.ipc_channels.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
ActiveRanksOutput(status=dp_active_ranks.tolist())
|
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(
|
def _relay_forward_payload(
|
||||||
self, future_indices: torch.Tensor, batch_result: GenerationBatchResult
|
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
|
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 (
|
if (
|
||||||
not self.spec_algorithm.is_none()
|
not self.spec_algorithm.is_none()
|
||||||
and self.metrics_reporter.spec_total_num_forward_ct > 0
|
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.disagg_decode_prealloc_queue.enqueue_held_rebootstrap()
|
||||||
self._engine_paused = False
|
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(
|
def load_lora_adapter(
|
||||||
self, recv_req: LoadLoRAAdapterReqInput
|
self, recv_req: LoadLoRAAdapterReqInput
|
||||||
) -> LoadLoRAAdapterReqOutput:
|
) -> LoadLoRAAdapterReqOutput:
|
||||||
@@ -4334,11 +4465,13 @@ def configure_scheduler_process(
|
|||||||
moe_ep_rank: int,
|
moe_ep_rank: int,
|
||||||
pp_rank: int,
|
pp_rank: int,
|
||||||
dp_rank: Optional[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]:
|
) -> Optional[int]:
|
||||||
"""Configure scheduler worker: logging, process title, etc.
|
"""Configure scheduler worker logging and process title.
|
||||||
|
|
||||||
Returns:
|
display_* ranks are cosmetic; runtime ranks stay local.
|
||||||
dp_rank
|
|
||||||
"""
|
"""
|
||||||
kill_itself_when_parent_died()
|
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
|
# [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"])
|
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 = ""
|
prefix = ""
|
||||||
if dp_rank is not None:
|
if shown_dp is not None:
|
||||||
prefix += f" DP{dp_rank}"
|
prefix += f" DP{shown_dp}"
|
||||||
if server_args.pp_size > 1:
|
if server_args.pp_size > 1:
|
||||||
prefix += f" PP{pp_rank}"
|
prefix += f" PP{pp_rank}"
|
||||||
if server_args.attn_cp_size > 1:
|
if server_args.attn_cp_size > 1:
|
||||||
@@ -4357,9 +4496,9 @@ def configure_scheduler_process(
|
|||||||
if server_args.moe_dp_size > 1:
|
if server_args.moe_dp_size > 1:
|
||||||
prefix += f" MOE_DP{moe_dp_rank}"
|
prefix += f" MOE_DP{moe_dp_rank}"
|
||||||
if server_args.tp_size > 1:
|
if server_args.tp_size > 1:
|
||||||
prefix += f" TP{tp_rank}"
|
prefix += f" TP{shown_tp}"
|
||||||
if server_args.ep_size > 1:
|
if server_args.ep_size > 1:
|
||||||
prefix += f" EP{moe_ep_rank}"
|
prefix += f" EP{shown_moe_ep}"
|
||||||
|
|
||||||
# Config the process
|
# Config the process
|
||||||
setproctitle.setproctitle(f"sglang::scheduler{prefix.replace(' ', '_')}")
|
setproctitle.setproctitle(f"sglang::scheduler{prefix.replace(' ', '_')}")
|
||||||
@@ -4393,6 +4532,9 @@ def run_scheduler_process(
|
|||||||
pp_rank: int,
|
pp_rank: int,
|
||||||
dp_rank: Optional[int],
|
dp_rank: Optional[int],
|
||||||
pipe_writer,
|
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 so hooks can override Scheduler and its dependencies.
|
||||||
load_plugins()
|
load_plugins()
|
||||||
@@ -4405,6 +4547,9 @@ def run_scheduler_process(
|
|||||||
moe_ep_rank,
|
moe_ep_rank,
|
||||||
pp_rank,
|
pp_rank,
|
||||||
dp_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()
|
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 import get_tp_group
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.environ import envs
|
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.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler_components.recv_skipper import (
|
from sglang.srt.managers.scheduler_components.recv_skipper import (
|
||||||
SchedulerRecvSkipper,
|
SchedulerRecvSkipper,
|
||||||
@@ -36,6 +37,43 @@ if TYPE_CHECKING:
|
|||||||
_ENABLE_METRICS_DP_ATTENTION = envs.SGLANG_ENABLE_METRICS_DP_ATTENTION.get()
|
_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
|
@dataclass
|
||||||
class MLPSyncBatchInfo:
|
class MLPSyncBatchInfo:
|
||||||
dp_size: int
|
dp_size: int
|
||||||
@@ -88,27 +126,56 @@ class MLPSyncBatchInfo:
|
|||||||
dtype=dtype,
|
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)
|
local_info_tensor = self._get_local_tensor(device=device)
|
||||||
global_info_tensor = torch.empty(
|
fallback_tensor = self._get_fallback_tensor(device=device)
|
||||||
(self.dp_size, self.tp_size * self.cp_size, 7),
|
info_width = local_info_tensor.numel()
|
||||||
dtype=torch.int64,
|
# Inactive max_world_size slots must decode as IDLE.
|
||||||
device=device,
|
global_info_tensor = fallback_tensor.expand(
|
||||||
)
|
self.dp_size, self.tp_size * self.cp_size, info_width
|
||||||
|
).contiguous()
|
||||||
|
|
||||||
|
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(
|
torch.distributed.all_gather_into_tensor(
|
||||||
global_info_tensor.flatten(),
|
global_info_tensor.flatten(),
|
||||||
local_info_tensor,
|
local_info_tensor,
|
||||||
group=group,
|
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":
|
if device == "cpu":
|
||||||
tp_active_ranks = get_tp_group().active_ranks_cpu
|
tp_active_ranks = get_tp_group().active_ranks_cpu
|
||||||
else:
|
else:
|
||||||
tp_active_ranks = get_tp_group().active_ranks
|
tp_active_ranks = get_tp_group().active_ranks
|
||||||
|
if tp_active_ranks.shape[0] < num_ranks_in_tp_info:
|
||||||
# Set fallback values for inactive ranks
|
tp_active_ranks = torch.ones(
|
||||||
tp_info = global_info_tensor.view(self.dp_size * self.tp_size * self.cp_size, 7)
|
num_ranks_in_tp_info,
|
||||||
tp_info[tp_active_ranks == 0] = self._get_fallback_tensor(device=device)
|
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, :]
|
tp0_info = global_info_tensor[:, 0, :]
|
||||||
self.tp0_info = tp0_info
|
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
|
local_batch.is_extend_in_batch = is_extend_in_batch
|
||||||
|
|
||||||
tbo_preparer = TboDPAttentionPreparer()
|
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
|
disable_overlap_schedule
|
||||||
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
|
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
|
||||||
):
|
):
|
||||||
@@ -217,6 +291,13 @@ def prepare_mlp_sync_batch_raw(
|
|||||||
device = "cpu"
|
device = "cpu"
|
||||||
|
|
||||||
local_can_run_tbo, local_forward_mode = tbo_preparer.prepare_all_gather(local_batch)
|
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(
|
mlp_sync_info = MLPSyncBatchInfo(
|
||||||
dp_size=dp_size,
|
dp_size=dp_size,
|
||||||
@@ -232,7 +313,11 @@ def prepare_mlp_sync_batch_raw(
|
|||||||
)
|
)
|
||||||
|
|
||||||
if not skip_all_gather:
|
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 = (
|
mlp_sync_info.tbo_split_seq_index, mlp_sync_info.global_forward_mode = (
|
||||||
tbo_preparer.compute_output(
|
tbo_preparer.compute_output(
|
||||||
|
|||||||
@@ -167,7 +167,10 @@ class SchedulerRequestReceiver:
|
|||||||
# controller, so we broadcast within attn_tp_group + attn_cp_group
|
# controller, so we broadcast within attn_tp_group + attn_cp_group
|
||||||
# instead of the full tp_group. This avoids an expensive
|
# instead of the full tp_group. This avoids an expensive
|
||||||
# all-ranks gloo sync.
|
# 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 _local_ctrl:
|
||||||
if self.ps.attn_tp_size != 1:
|
if self.ps.attn_tp_size != 1:
|
||||||
control_reqs = broadcast_pyobj(
|
control_reqs = broadcast_pyobj(
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ from sglang.srt.managers.io_struct import (
|
|||||||
RemoveExternalCorpusReqOutput,
|
RemoveExternalCorpusReqOutput,
|
||||||
ResumeMemoryOccupationReqInput,
|
ResumeMemoryOccupationReqInput,
|
||||||
ResumeMemoryOccupationReqOutput,
|
ResumeMemoryOccupationReqOutput,
|
||||||
|
ScaleElasticEPReqOutput,
|
||||||
SendWeightsToRemoteInstanceReqInput,
|
SendWeightsToRemoteInstanceReqInput,
|
||||||
SendWeightsToRemoteInstanceReqOutput,
|
SendWeightsToRemoteInstanceReqOutput,
|
||||||
SetInternalStateReq,
|
SetInternalStateReq,
|
||||||
@@ -118,6 +119,7 @@ _COMMUNICATOR_SPECS = [
|
|||||||
("expert_distribution", ExpertDistributionReqOutput),
|
("expert_distribution", ExpertDistributionReqOutput),
|
||||||
("update_lora_adapter", LoRAUpdateOutput),
|
("update_lora_adapter", LoRAUpdateOutput),
|
||||||
("dumper_control", DumperControlReqOutput),
|
("dumper_control", DumperControlReqOutput),
|
||||||
|
("scale_elastic_ep", ScaleElasticEPReqOutput),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@@ -141,6 +143,23 @@ class TokenizerControlMixin:
|
|||||||
dispatch_pairs.append((resp_type, comm.handle_recv))
|
dispatch_pairs.append((resp_type, comm.handle_recv))
|
||||||
self._result_dispatcher += TypeBasedDispatcher(dispatch_pairs)
|
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(
|
async def add_external_corpus(
|
||||||
self: TokenizerManager, obj: AddExternalCorpusReqInput
|
self: TokenizerManager, obj: AddExternalCorpusReqInput
|
||||||
) -> AddExternalCorpusReqOutput:
|
) -> AddExternalCorpusReqOutput:
|
||||||
@@ -821,7 +840,9 @@ class TokenizerControlMixin:
|
|||||||
List of LoadSnapshot, one per scheduler (filtered by dp_rank if specified)
|
List of LoadSnapshot, one per scheduler (filtered by dp_rank if specified)
|
||||||
"""
|
"""
|
||||||
self.auto_create_handle_loop()
|
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 []
|
return []
|
||||||
|
|
||||||
reader = self.load_snapshot_reader
|
reader = self.load_snapshot_reader
|
||||||
|
|||||||
@@ -65,6 +65,7 @@ from sglang.srt.managers.io_struct import (
|
|||||||
BatchTokenizedGenerateReqInput,
|
BatchTokenizedGenerateReqInput,
|
||||||
ConfigureLoggingReq,
|
ConfigureLoggingReq,
|
||||||
ContinueGenerationReqInput,
|
ContinueGenerationReqInput,
|
||||||
|
ElasticScaleUpdateReq,
|
||||||
EmbeddingReqInput,
|
EmbeddingReqInput,
|
||||||
FreezeGCReq,
|
FreezeGCReq,
|
||||||
GenerateReqInput,
|
GenerateReqInput,
|
||||||
@@ -72,6 +73,8 @@ from sglang.srt.managers.io_struct import (
|
|||||||
LoadLoRAAdapterReqInput,
|
LoadLoRAAdapterReqInput,
|
||||||
OpenSessionReqOutput,
|
OpenSessionReqOutput,
|
||||||
PauseGenerationReqInput,
|
PauseGenerationReqInput,
|
||||||
|
ScaleElasticEPReqInput,
|
||||||
|
ScaleElasticEPReqOutput,
|
||||||
SessionParams,
|
SessionParams,
|
||||||
ShutdownReq,
|
ShutdownReq,
|
||||||
TokenizedEmbeddingReqInput,
|
TokenizedEmbeddingReqInput,
|
||||||
@@ -279,6 +282,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
):
|
):
|
||||||
# Parse args
|
# Parse args
|
||||||
self.server_args = server_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.enable_metrics = server_args.enable_metrics
|
||||||
self.incremental_streaming_output = server_args.incremental_streaming_output
|
self.incremental_streaming_output = server_args.incremental_streaming_output
|
||||||
self.enable_lora = server_args.enable_lora
|
self.enable_lora = server_args.enable_lora
|
||||||
@@ -491,6 +498,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self.model_update_result: Optional[Awaitable[UpdateWeightFromDiskReqOutput]] = (
|
self.model_update_result: Optional[Awaitable[UpdateWeightFromDiskReqOutput]] = (
|
||||||
None
|
None
|
||||||
)
|
)
|
||||||
|
self.model_update_expected_workers = self.elastic_worker_count
|
||||||
|
self.model_update_tmp: List[UpdateWeightFromDiskReqOutput] = []
|
||||||
self.is_pause = False
|
self.is_pause = False
|
||||||
self.is_pause_cond = asyncio.Condition()
|
self.is_pause_cond = asyncio.Condition()
|
||||||
|
|
||||||
@@ -604,6 +613,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
# Same skip-detokenizer forwarding case as above.
|
# Same skip-detokenizer forwarding case as above.
|
||||||
(ConfigureLoggingReq, lambda x: None),
|
(ConfigureLoggingReq, lambda x: None),
|
||||||
(ActiveRanksOutput, self.update_active_ranks),
|
(ActiveRanksOutput, self.update_active_ranks),
|
||||||
|
(ElasticScaleUpdateReq, self.forward_elastic_scale_update),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.init_communicators(self.server_args)
|
self.init_communicators(self.server_args)
|
||||||
@@ -623,7 +633,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self._set_default_priority(obj)
|
self._set_default_priority(obj)
|
||||||
|
|
||||||
if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None:
|
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:
|
if dp_size <= 1 and obj.routed_dp_rank == 0:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"routed_dp_rank={obj.routed_dp_rank} is ignored because dp_size={dp_size}"
|
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(
|
async def _wait_for_model_update_from_disk(
|
||||||
self, obj: UpdateWeightFromDiskReqInput
|
self, obj: UpdateWeightFromDiskReqInput
|
||||||
) -> Tuple[bool, str]:
|
) -> 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()
|
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
|
result = await self.model_update_result
|
||||||
if result.success:
|
if result.success:
|
||||||
self._update_model_path_info(obj.model_path, obj.load_format)
|
self._update_model_path_info(obj.model_path, obj.load_format)
|
||||||
return result.success, result.message, result.num_paused_requests
|
return result.success, result.message, result.num_paused_requests
|
||||||
else: # self.server_args.dp_size > 1
|
else:
|
||||||
self.model_update_tmp = []
|
|
||||||
result = await self.model_update_result
|
result = await self.model_update_result
|
||||||
|
|
||||||
all_success = all([r.success for r in 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):
|
def update_active_ranks(self, ranks: ActiveRanksOutput):
|
||||||
self._dispatch_to_scheduler(ranks)
|
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):
|
def _handle_open_session_req_output(self, recv_obj):
|
||||||
future = self.session_futures.get(recv_obj.session_id)
|
future = self.session_futures.get(recv_obj.session_id)
|
||||||
if future is None:
|
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)
|
future.set_result(recv_obj.session_id if recv_obj.success else None)
|
||||||
|
|
||||||
def _handle_update_weights_from_disk_req_output(self, recv_obj):
|
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)
|
self.model_update_result.set_result(recv_obj)
|
||||||
else: # self.server_args.dp_size > 1
|
else:
|
||||||
self.model_update_tmp.append(recv_obj)
|
self.model_update_tmp.append(recv_obj)
|
||||||
# set future if the all results are received
|
if len(self.model_update_tmp) == self.model_update_expected_workers:
|
||||||
if len(self.model_update_tmp) == self.server_args.dp_size:
|
|
||||||
self.model_update_result.set_result(self.model_update_tmp)
|
self.model_update_result.set_result(self.model_update_tmp)
|
||||||
|
|
||||||
async def _validate_and_resolve_lora(
|
async def _validate_and_resolve_lora(
|
||||||
|
|||||||
@@ -310,7 +310,11 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
self.world_group = get_world_group()
|
self.world_group = get_world_group()
|
||||||
|
|
||||||
# Sync random seed across TP workers
|
# 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(
|
self.random_seed = broadcast_pyobj(
|
||||||
[server_args.random_seed],
|
[server_args.random_seed],
|
||||||
self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
|
self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
|
world_dp_gather_enabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
||||||
ForwardBatchDeepSeekMHAMixin,
|
ForwardBatchDeepSeekMHAMixin,
|
||||||
@@ -75,6 +76,25 @@ _skip_attn_backend_init_warned = False
|
|||||||
_is_npu = is_npu()
|
_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):
|
class ForwardMode(IntEnum):
|
||||||
# Extend a sequence. The KV cache of the beginning part of the sequence is already computed (e.g., system prompt).
|
# 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.
|
# 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
|
assert self.global_num_tokens_for_logprob_cpu is not None
|
||||||
|
|
||||||
self._original_batch_size = self.batch_size
|
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)
|
sync_group_size = len(global_num_tokens)
|
||||||
attn_tp_size = get_parallel().attn_tp_size
|
attn_tp_size = get_parallel().attn_tp_size
|
||||||
|
|
||||||
@@ -1168,6 +1188,13 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
dp_padding_mode = DpPaddingMode.get_dp_padding_mode(
|
dp_padding_mode = DpPaddingMode.get_dp_padding_mode(
|
||||||
self.is_extend_in_batch, global_num_tokens
|
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
|
# 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
|
# 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
|
# 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
|
from typing import Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
from sglang.srt.configs.load_config import LoadConfig
|
from sglang.srt.configs.load_config import LoadConfig
|
||||||
from sglang.srt.configs.model_config import (
|
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 (
|
from sglang.srt.elastic_ep.elastic_ep import (
|
||||||
ElasticEPStateManager,
|
ElasticEPStateManager,
|
||||||
get_healthy_expert_location_src_rank,
|
get_healthy_expert_location_src_rank,
|
||||||
|
get_scale_cohort_target,
|
||||||
join_process_groups,
|
join_process_groups,
|
||||||
|
join_scale_process_group,
|
||||||
maybe_rebalance_after_rank_fault,
|
maybe_rebalance_after_rank_fault,
|
||||||
maybe_recover_ep_ranks,
|
maybe_recover_ep_ranks,
|
||||||
|
register_scale_cohort,
|
||||||
|
try_admit_scale_ranks,
|
||||||
)
|
)
|
||||||
from sglang.srt.elastic_ep.expert_backup_client import ExpertBackupClient
|
from sglang.srt.elastic_ep.expert_backup_client import ExpertBackupClient
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -55,6 +60,8 @@ from sglang.srt.eplb.expert_distribution import (
|
|||||||
set_global_expert_distribution_recorder,
|
set_global_expert_distribution_recorder,
|
||||||
)
|
)
|
||||||
from sglang.srt.eplb.expert_location import (
|
from sglang.srt.eplb.expert_location import (
|
||||||
|
ExpertLocationMetadata,
|
||||||
|
append_trivial_expert_slots,
|
||||||
broadcast_global_expert_location_metadata,
|
broadcast_global_expert_location_metadata,
|
||||||
compute_initial_expert_location_metadata,
|
compute_initial_expert_location_metadata,
|
||||||
format_expert_location_layout,
|
format_expert_location_layout,
|
||||||
@@ -276,6 +283,7 @@ class ModelRunner:
|
|||||||
self.attention_chunk_size = model_config.attention_chunk_size
|
self.attention_chunk_size = model_config.attention_chunk_size
|
||||||
self.enable_elastic_ep = server_args.elastic_ep_backend is not None
|
self.enable_elastic_ep = server_args.elastic_ep_backend is not None
|
||||||
self.forward_pass_id = 0
|
self.forward_pass_id = 0
|
||||||
|
self._pending_elastic_scale_update = None
|
||||||
self.init_new_workspace = False
|
self.init_new_workspace = False
|
||||||
self.draft_model_idx = draft_model_idx
|
self.draft_model_idx = draft_model_idx
|
||||||
self.enable_hisparse = server_args.enable_hisparse
|
self.enable_hisparse = server_args.enable_hisparse
|
||||||
@@ -349,17 +357,7 @@ class ModelRunner:
|
|||||||
self.initialize()
|
self.initialize()
|
||||||
self.check_quantized_moe_compatibility()
|
self.check_quantized_moe_compatibility()
|
||||||
|
|
||||||
if (
|
self._initialize_elastic_ep_joiner()
|
||||||
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()
|
|
||||||
|
|
||||||
if self.is_multimodal:
|
if self.is_multimodal:
|
||||||
sanity_check_mm_pad_shift_value(self.model_config.vocab_size)
|
sanity_check_mm_pad_shift_value(self.model_config.vocab_size)
|
||||||
@@ -378,6 +376,87 @@ class ModelRunner:
|
|||||||
self.init_weight_updater()
|
self.init_weight_updater()
|
||||||
self.init_weight_exporter()
|
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):
|
def init_msprobe(self):
|
||||||
self.msprobe_debugger = misc_utils.create_msprobe_debugger(self.server_args)
|
self.msprobe_debugger = misc_utils.create_msprobe_debugger(self.server_args)
|
||||||
|
|
||||||
@@ -526,11 +605,16 @@ class ModelRunner:
|
|||||||
def maybe_init_expert_location_metadata(self):
|
def maybe_init_expert_location_metadata(self):
|
||||||
if self.is_draft_worker:
|
if self.is_draft_worker:
|
||||||
return
|
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(
|
set_global_expert_location_metadata(
|
||||||
compute_initial_expert_location_metadata(
|
compute_initial_expert_location_metadata(
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
model_config=self.model_config,
|
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():
|
if self.ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get():
|
||||||
@@ -542,7 +626,7 @@ class ModelRunner:
|
|||||||
ExpertDistributionRecorder.init_new(
|
ExpertDistributionRecorder.init_new(
|
||||||
self.server_args,
|
self.server_args,
|
||||||
get_global_expert_location_metadata(),
|
get_global_expert_location_metadata(),
|
||||||
rank=self.ps.tp_rank,
|
rank=expert_rank,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -920,6 +1004,7 @@ class ModelRunner:
|
|||||||
dist_barrier_after_load(
|
dist_barrier_after_load(
|
||||||
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
|
is_ep_scale_joiner=self.server_args.is_ep_scale_joiner,
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_lora_manager(self):
|
def init_lora_manager(self):
|
||||||
@@ -1224,13 +1309,7 @@ class ModelRunner:
|
|||||||
self.msprobe_debugger.step()
|
self.msprobe_debugger.step()
|
||||||
|
|
||||||
if self.server_args.elastic_ep_backend is not None:
|
if self.server_args.elastic_ep_backend is not None:
|
||||||
recovered = maybe_recover_ep_ranks(
|
self.maybe_join_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 output
|
return output
|
||||||
|
|
||||||
@@ -1471,6 +1550,216 @@ class ModelRunner:
|
|||||||
action=action, allow_quant_error=allow_quant_error
|
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(
|
def _maybe_rebalance_after_rank_fault(
|
||||||
self,
|
self,
|
||||||
output: ModelRunnerOutput,
|
output: ModelRunnerOutput,
|
||||||
|
|||||||
@@ -248,9 +248,15 @@ 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":
|
if elastic_ep_backend == "mooncake":
|
||||||
# Mooncake does not support `monitored_barrier`
|
# Mooncake does not support `monitored_barrier`
|
||||||
|
if not is_ep_scale_joiner:
|
||||||
dist.barrier(group=get_tp_group().cpu_group)
|
dist.barrier(group=get_tp_group().cpu_group)
|
||||||
else:
|
else:
|
||||||
# Handle the case where some ranks do not finish loading.
|
# Handle the case where some ranks do not finish loading.
|
||||||
|
|||||||
@@ -345,6 +345,8 @@ class DpFlags(_FlagGroupBase):
|
|||||||
migrates them."""
|
migrates them."""
|
||||||
|
|
||||||
enabled: bool = False
|
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
|
# Hybrid-SSM models materialize idle ranks via the MAX_LEN fabricated-row
|
||||||
# conversion (set when hf_config has hybrid_override_pattern).
|
# conversion (set when hf_config has hybrid_override_pattern).
|
||||||
max_len_with_idle: bool = False
|
max_len_with_idle: bool = False
|
||||||
|
|||||||
@@ -1998,9 +1998,40 @@ class ServerArgs:
|
|||||||
bool,
|
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.",
|
"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
|
] = 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[
|
elastic_ep_rejoin: A[
|
||||||
bool,
|
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
|
] = False
|
||||||
disable_flashinfer_cutlass_moe_fp4_allgather: A[
|
disable_flashinfer_cutlass_moe_fp4_allgather: A[
|
||||||
bool,
|
bool,
|
||||||
@@ -5634,10 +5665,21 @@ class ServerArgs:
|
|||||||
):
|
):
|
||||||
self.ep_dispatch_algorithm = "static"
|
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
|
assert self._resolved().ep_size > 1
|
||||||
|
|
||||||
def _handle_elastic_ep(self):
|
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.elastic_ep_backend is not None:
|
||||||
if self.enable_eplb:
|
if self.enable_eplb:
|
||||||
if self.eplb_algorithm == "auto":
|
if self.eplb_algorithm == "auto":
|
||||||
@@ -5653,10 +5695,144 @@ class ServerArgs:
|
|||||||
self.mooncake_ib_device = self._validate_ib_devices(
|
self.mooncake_ib_device = self._validate_ib_devices(
|
||||||
self.mooncake_ib_device
|
self.mooncake_ib_device
|
||||||
)
|
)
|
||||||
if self.elastic_ep_rejoin:
|
if self.ep_join_mode is not None:
|
||||||
assert (
|
assert (
|
||||||
self.elastic_ep_backend is not None
|
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):
|
def _handle_expert_distribution_metrics(self):
|
||||||
if self.enable_expert_distribution_metrics and (
|
if self.enable_expert_distribution_metrics and (
|
||||||
@@ -6917,6 +7093,15 @@ class ServerArgs:
|
|||||||
def engine_info_bootstrap_url(self):
|
def engine_info_bootstrap_url(self):
|
||||||
return self.url(port=self.engine_info_bootstrap_port)
|
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):
|
def ssl_verify(self):
|
||||||
"""Return the value for the requests library's verify= parameter.
|
"""Return the value for the requests library's verify= parameter.
|
||||||
|
|
||||||
@@ -7113,6 +7298,7 @@ class ServerArgs:
|
|||||||
|
|
||||||
def check_server_args(self):
|
def check_server_args(self):
|
||||||
# Check parallel size constraints
|
# Check parallel size constraints
|
||||||
|
if self.ep_join_mode != "scale":
|
||||||
assert (
|
assert (
|
||||||
self.tp_size * self.pp_size
|
self.tp_size * self.pp_size
|
||||||
) % self.nnodes == 0, "tp_size must be divisible by number of nodes"
|
) % self.nnodes == 0, "tp_size must be divisible by number of nodes"
|
||||||
@@ -7870,7 +8056,11 @@ class PortArgs:
|
|||||||
# (no availability-based search). If incrementing would
|
# (no availability-based search). If incrementing would
|
||||||
# overflow the valid TCP range, decrement instead.
|
# overflow the valid TCP range, decrement instead.
|
||||||
NUM_DERIVED_PORTS = 5
|
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
|
port_base = dist_init_port - NUM_DERIVED_PORTS - 1
|
||||||
else:
|
else:
|
||||||
port_base = dist_init_port + 1
|
port_base = dist_init_port + 1
|
||||||
@@ -7886,8 +8076,10 @@ class PortArgs:
|
|||||||
assert worker_ports is not None
|
assert worker_ports is not None
|
||||||
scheduler_input_port = worker_ports[dp_rank]
|
scheduler_input_port = worker_ports[dp_rank]
|
||||||
|
|
||||||
|
is_joiner = server_args.is_ep_scale_joiner
|
||||||
try:
|
try:
|
||||||
if dp_rank is None:
|
if dp_rank is None:
|
||||||
|
if not is_joiner:
|
||||||
wait_port_available(dist_init_port, "dist_init_port")
|
wait_port_available(dist_init_port, "dist_init_port")
|
||||||
wait_port_available(port_base, "port_base")
|
wait_port_available(port_base, "port_base")
|
||||||
wait_port_available(detokenizer_port, "detokenizer_port")
|
wait_port_available(detokenizer_port, "detokenizer_port")
|
||||||
|
|||||||
@@ -2401,7 +2401,7 @@ def _get_fastapi_request_path(request) -> Tuple[str, bool]:
|
|||||||
for route in request.app.routes:
|
for route in request.app.routes:
|
||||||
match, child_scope = route.matches(request.scope)
|
match, child_scope = route.matches(request.scope)
|
||||||
if match == Match.FULL:
|
if match == Match.FULL:
|
||||||
return route.path, True
|
return getattr(route, "path", request.url.path), True
|
||||||
|
|
||||||
return request.url.path, False
|
return request.url.path, False
|
||||||
|
|
||||||
@@ -3475,6 +3475,13 @@ def require_mlp_tp_gather(server_args: ServerArgs):
|
|||||||
|
|
||||||
if server_args.enable_dp_attention:
|
if server_args.enable_dp_attention:
|
||||||
assert server_args.dp_size > 1, "dp_size must be greater than 1"
|
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 (
|
if (
|
||||||
server_args.moe_dense_tp_size is None
|
server_args.moe_dense_tp_size is None
|
||||||
): # TODO(ch-wan): some MoE models do not have dense layers
|
): # 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):
|
class _FakeTokenizerManager(TokenizerControlMixin):
|
||||||
def __init__(self, reader, dp_size: int):
|
def __init__(self, reader, dp_size: int):
|
||||||
self.load_snapshot_reader = reader
|
self.load_snapshot_reader = reader
|
||||||
|
self.elastic_worker_count = dp_size
|
||||||
self.server_args = SimpleNamespace(
|
self.server_args = SimpleNamespace(
|
||||||
dp_size=dp_size,
|
dp_size=dp_size,
|
||||||
enable_dp_attention=False,
|
enable_dp_attention=False,
|
||||||
|
|||||||
@@ -10,14 +10,16 @@ import unittest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.eplb.expert_location import (
|
from sglang.srt.eplb.expert_location import (
|
||||||
|
_compute_logical_to_all_physical_map,
|
||||||
|
append_trivial_expert_slots,
|
||||||
compute_logical_to_rank_dispatch_physical_map,
|
compute_logical_to_rank_dispatch_physical_map,
|
||||||
)
|
)
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
def _make_server_args(ep_size: int, nnodes: int):
|
def _make_server_args(ep_size: int, nnodes: int):
|
||||||
"""Minimal server_args stub — only ep_size and nnodes are used."""
|
"""Minimal server_args stub for expert placement tests."""
|
||||||
return types.SimpleNamespace(ep_size=ep_size, nnodes=nnodes)
|
return types.SimpleNamespace(ep_size=ep_size, nnodes=nnodes, ep_join_mode=None)
|
||||||
|
|
||||||
|
|
||||||
def _make_logical_to_all_physical_map(
|
def _make_logical_to_all_physical_map(
|
||||||
@@ -194,6 +196,24 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase):
|
|||||||
self.assertEqual(result.shape, (self.NUM_LAYERS, 1))
|
self.assertEqual(result.shape, (self.NUM_LAYERS, 1))
|
||||||
self.assertTrue(torch.all(result >= 0))
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.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.
|
tie-break on `total_requests` transitively covers that state.
|
||||||
|
|
||||||
Fragility: scheduler tests bypass `DataParallelController.__init__` via
|
Fragility: scheduler tests bypass `DataParallelController.__init__` via
|
||||||
`__new__` and inject only the attrs the schedulers read (`workers`,
|
`__new__` and inject only the attrs the schedulers read (`workers`, `status`,
|
||||||
`status`, `round_robin_counter`, `dp_budget`). Update `_make_controller`
|
`_active_workers`, `round_robin_counter`, `dp_budget`). Update `_make_controller`
|
||||||
if a scheduler starts reading another attr. `maybe_external_dp_rank_routing`
|
if a scheduler starts reading another attr. `maybe_external_dp_rank_routing`
|
||||||
is exercised as the real method, no mock.
|
is exercised as the real method, no mock.
|
||||||
"""
|
"""
|
||||||
@@ -48,6 +48,7 @@ def _make_controller(dp_size: int) -> DataParallelController:
|
|||||||
ctl = DataParallelController.__new__(DataParallelController)
|
ctl = DataParallelController.__new__(DataParallelController)
|
||||||
ctl.workers = [MagicMock(name=f"worker_{i}") for i in range(dp_size)]
|
ctl.workers = [MagicMock(name=f"worker_{i}") for i in range(dp_size)]
|
||||||
ctl.status = [True] * dp_size
|
ctl.status = [True] * dp_size
|
||||||
|
ctl._active_workers = list(range(dp_size))
|
||||||
ctl.round_robin_counter = 0
|
ctl.round_robin_counter = 0
|
||||||
ctl.dp_budget = DPBudget(dp_size=dp_size)
|
ctl.dp_budget = DPBudget(dp_size=dp_size)
|
||||||
return ctl
|
return ctl
|
||||||
|
|||||||
@@ -1812,11 +1812,16 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(
|
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},
|
{"enable_dp_attention": False, "enable_dp_lm_head": False},
|
||||||
)
|
)
|
||||||
self.assertEqual(
|
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:
|
with patch("sglang.srt.environ.envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE") as e:
|
||||||
|
|||||||
Reference in New Issue
Block a user