config: route parallel config-leaf reads through get_parallel() (#33170)
The parallel namespace joins the accessor migration: 106 config-leaf reads (enable_dp_lm_head, enable_dp_attention, pp_async_batch_depth, dp_size, ep_join_rank_offset, dwdp_size, ...) flip from get_server_args()/ self.server_args to get_parallel(), which serves config leaves from the published parallel bag via __getattr__. - ParallelContext.__getattr__ is restructured to stay dynamo-traceable (object.__getattribute__ graph-breaks): gate helpers such as enable_moe_dense_fully_dp() run inside compiled model forwards. A fullgraph regression test pins the pattern. - The five live-shadowed topology sizes (tp/pp/dcp/attn_cp/moe_dp_size) keep their server_args reads: the live @property wins on the accessor, and conditionally-initialized groups would fail loud at unconditional call sites. - Elastic-EP scale writers (ep_size/dp_size x4 in model_runner) reroute to get_context().override together with their remaining instance readers (expert_location gpus-per-node paths); the ServerArgs.override ratchet drops 39 -> 35. - The expert placement helpers (compute_logical_to_rank_dispatch_ physical_map, _compute_logical_to_all_physical_map, _prefer_same_node_experts) now read everything from the bags and drop their server_args parameter; their unit tests publish the config they need instead of stubbing it.
This commit is contained in:
@@ -758,7 +758,7 @@ class TboForwardBatchPreparer:
|
|||||||
|
|
||||||
# TODO improve, e.g. unify w/ `init_raw`
|
# TODO improve, e.g. unify w/ `init_raw`
|
||||||
if (
|
if (
|
||||||
get_server_args().moe_dense_tp_size == 1
|
get_parallel().moe_dense_tp_size == 1
|
||||||
and batch.global_dp_buffer_len is not None
|
and batch.global_dp_buffer_len is not None
|
||||||
):
|
):
|
||||||
sum_len = end_token_index - start_token_index
|
sum_len = end_token_index - start_token_index
|
||||||
|
|||||||
@@ -174,7 +174,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self.attn_dp_size = get_attention_dp_size()
|
self.attn_dp_size = get_attention_dp_size()
|
||||||
self.attn_dp_rank = get_attention_dp_rank()
|
self.attn_dp_rank = get_attention_dp_rank()
|
||||||
self.system_dp_size = (
|
self.system_dp_size = (
|
||||||
1 if server_args.enable_dp_attention else server_args.dp_size
|
1 if get_parallel().enable_dp_attention else get_parallel().dp_size
|
||||||
)
|
)
|
||||||
self.system_dp_rank = (
|
self.system_dp_rank = (
|
||||||
self.kv_args.system_dp_rank if self.kv_args.system_dp_rank else 0
|
self.kv_args.system_dp_rank if self.kv_args.system_dp_rank else 0
|
||||||
@@ -183,7 +183,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self.pp_rank = self.kv_args.pp_rank
|
self.pp_rank = self.kv_args.pp_rank
|
||||||
self.local_ip = get_local_ip_auto()
|
self.local_ip = get_local_ip_auto()
|
||||||
cp_sharded_prefill = self.attn_cp_size > 1 and (
|
cp_sharded_prefill = self.attn_cp_size > 1 and (
|
||||||
self.is_hybrid_mla_backend or server_args.enable_dsa_cache_layer_split
|
self.is_hybrid_mla_backend or get_parallel().enable_dsa_cache_layer_split
|
||||||
)
|
)
|
||||||
|
|
||||||
hybrid_decode_pulls_all_ranks = (
|
hybrid_decode_pulls_all_ranks = (
|
||||||
@@ -651,7 +651,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
`Connection refused`, and the leader's `prefill_port_table` ends
|
`Connection refused`, and the leader's `prefill_port_table` ends
|
||||||
up missing rows.
|
up missing rows.
|
||||||
"""
|
"""
|
||||||
if not self.dist_init_addr or self.server_args.nnodes == 1:
|
if not self.dist_init_addr or get_parallel().nnodes == 1:
|
||||||
return local_port
|
return local_port
|
||||||
|
|
||||||
if not (dist.is_available() and dist.is_initialized()):
|
if not (dist.is_available() and dist.is_initialized()):
|
||||||
@@ -703,10 +703,8 @@ class CommonKVManager(BaseKVManager):
|
|||||||
"rank_port": self.rank_port,
|
"rank_port": self.rank_port,
|
||||||
"page_size": self.kv_args.page_size,
|
"page_size": self.kv_args.page_size,
|
||||||
"kv_cache_dtype": self.kv_cache_dtype_str,
|
"kv_cache_dtype": self.kv_cache_dtype_str,
|
||||||
"load_balance_method": self.server_args.load_balance_method,
|
"load_balance_method": get_parallel().load_balance_method,
|
||||||
"enable_dsa_cache_layer_split": getattr(
|
"enable_dsa_cache_layer_split": get_parallel().enable_dsa_cache_layer_split,
|
||||||
self.server_args, "enable_dsa_cache_layer_split", False
|
|
||||||
),
|
|
||||||
# Self-register the HTTP API port so the decode can derive the PD
|
# Self-register the HTTP API port so the decode can derive the PD
|
||||||
# retract rebootstrap /generate URL from bootstrap info instead of a
|
# retract rebootstrap /generate URL from bootstrap info instead of a
|
||||||
# router-injected pd_rebootstrap_prefill_url.
|
# router-injected pd_rebootstrap_prefill_url.
|
||||||
@@ -1078,12 +1076,11 @@ class CommonKVSender(BaseKVSender):
|
|||||||
return
|
return
|
||||||
|
|
||||||
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping)
|
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping)
|
||||||
if self.kv_mgr.server_args.dp_size > 1 and not req_has_disagg_prefill_dp_rank:
|
if get_parallel().dp_size > 1 and not req_has_disagg_prefill_dp_rank:
|
||||||
if self.kv_mgr.server_args.load_balance_method != "follow_bootstrap_room":
|
if get_parallel().load_balance_method != "follow_bootstrap_room":
|
||||||
self._register_prefill_dp_rank()
|
self._register_prefill_dp_rank()
|
||||||
elif (
|
elif (
|
||||||
self.kv_mgr.attn_dp_rank
|
self.kv_mgr.attn_dp_rank != self.bootstrap_room % get_parallel().dp_size
|
||||||
!= self.bootstrap_room % self.kv_mgr.server_args.dp_size
|
|
||||||
):
|
):
|
||||||
# follow_bootstrap_room was overridden by external routed_dp_rank
|
# follow_bootstrap_room was overridden by external routed_dp_rank
|
||||||
if envs.SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK.get():
|
if envs.SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK.get():
|
||||||
@@ -1094,7 +1091,7 @@ class CommonKVSender(BaseKVSender):
|
|||||||
f"follow_bootstrap_room conflict: dispatched to dp_rank "
|
f"follow_bootstrap_room conflict: dispatched to dp_rank "
|
||||||
f"{self.kv_mgr.attn_dp_rank} but bootstrap_room "
|
f"{self.kv_mgr.attn_dp_rank} but bootstrap_room "
|
||||||
f"{self.bootstrap_room} implies dp_rank "
|
f"{self.bootstrap_room} implies dp_rank "
|
||||||
f"{self.bootstrap_room % self.kv_mgr.server_args.dp_size}. "
|
f"{self.bootstrap_room % get_parallel().dp_size}. "
|
||||||
f"Set SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK=1 "
|
f"Set SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK=1 "
|
||||||
f"to allow mixed routing.",
|
f"to allow mixed routing.",
|
||||||
)
|
)
|
||||||
@@ -1168,7 +1165,7 @@ class CommonKVSender(BaseKVSender):
|
|||||||
|
|
||||||
if (
|
if (
|
||||||
self.kv_mgr.enable_all_cp_ranks_for_transfer
|
self.kv_mgr.enable_all_cp_ranks_for_transfer
|
||||||
and not self.kv_mgr.server_args.enable_dsa_cache_layer_split
|
and not get_parallel().enable_dsa_cache_layer_split
|
||||||
):
|
):
|
||||||
kv_indices, index_slice = filter_kv_indices_for_cp_rank(
|
kv_indices, index_slice = filter_kv_indices_for_cp_rank(
|
||||||
self.kv_mgr,
|
self.kv_mgr,
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ from sglang.srt.observability.trace import (
|
|||||||
TraceReqContext,
|
TraceReqContext,
|
||||||
trace_set_thread_info,
|
trace_set_thread_info,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_schedule
|
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
@@ -1099,7 +1099,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
if (
|
if (
|
||||||
self.attn_cp_size > 1
|
self.attn_cp_size > 1
|
||||||
and self.attn_cp_rank != 0
|
and self.attn_cp_rank != 0
|
||||||
and not self.server_args.enable_dsa_cache_layer_split
|
and not get_parallel().enable_dsa_cache_layer_split
|
||||||
):
|
):
|
||||||
skip_state = True
|
skip_state = True
|
||||||
|
|
||||||
|
|||||||
@@ -466,18 +466,18 @@ class MultimemAllGatherer:
|
|||||||
# Lazy import avoids a module-load dependency on the distributed facade.
|
# Lazy import avoids a module-load dependency on the distributed facade.
|
||||||
from sglang.srt.distributed import get_tp_group
|
from sglang.srt.distributed import get_tp_group
|
||||||
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
tp_group = get_tp_group()
|
tp_group = get_tp_group()
|
||||||
# Only probe node topology when the deployment can actually span
|
# Only probe node topology when the deployment can actually span
|
||||||
# nodes. Check world_size first so a TP=1 gatherer short-circuits
|
# nodes. Check world_size first so a TP=1 gatherer short-circuits
|
||||||
# before reading server args (which may be unpublished on offline
|
# before reading the parallel config (which may be unpublished on
|
||||||
# paths). On a single node every TP rank is co-located, so skip the
|
# offline paths). On a single node every TP rank is co-located, so skip the
|
||||||
# in_the_same_node_as() all-reduce, which can segfault under some
|
# in_the_same_node_as() all-reduce, which can segfault under some
|
||||||
# EP/mooncake setups, and keep multimem enabled.
|
# EP/mooncake setups, and keep multimem enabled.
|
||||||
if (
|
if (
|
||||||
tp_group.world_size > 1
|
tp_group.world_size > 1
|
||||||
and get_server_args().nnodes > 1
|
and get_parallel().nnodes > 1
|
||||||
and not all(in_the_same_node_as(tp_group.cpu_group, source_rank=0))
|
and not all(in_the_same_node_as(tp_group.cpu_group, source_rank=0))
|
||||||
):
|
):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from sglang.srt.distributed import get_world_group, parallel_state
|
|||||||
from sglang.srt.distributed.utils import get_global_tcp_store
|
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.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import is_cpu, is_cuda
|
from sglang.srt.utils import is_cpu, is_cuda
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -308,13 +309,10 @@ def elastic_expanded_world_enabled() -> bool:
|
|||||||
|
|
||||||
Launch-time TP groups exclude ranks admitted during scale-up.
|
Launch-time TP groups exclude ranks admitted during scale-up.
|
||||||
"""
|
"""
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
inst = ElasticEPStateManager.instance()
|
inst = ElasticEPStateManager.instance()
|
||||||
if inst is None:
|
if inst is None:
|
||||||
return False
|
return False
|
||||||
sa = get_server_args()
|
if get_parallel().max_ep_size is None:
|
||||||
if sa.max_ep_size is None:
|
|
||||||
return False
|
return False
|
||||||
active_target_size = inst.effective_ep_size
|
active_target_size = inst.effective_ep_size
|
||||||
if inst.pending_ep_size is not None and inst.scale_phase in (
|
if inst.pending_ep_size is not None and inst.scale_phase in (
|
||||||
|
|||||||
@@ -32,10 +32,13 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def _prefer_same_node_experts(server_args: ServerArgs) -> bool:
|
def _prefer_same_node_experts() -> bool:
|
||||||
from sglang.srt.elastic_ep.elastic_ep import elastic_expanded_world_enabled
|
from sglang.srt.elastic_ep.elastic_ep import elastic_expanded_world_enabled
|
||||||
|
from sglang.srt.runtime_context import get_exec
|
||||||
|
|
||||||
return server_args.ep_join_mode != "scale" and not elastic_expanded_world_enabled()
|
return (
|
||||||
|
get_exec().moe.ep_join_mode != "scale" and not elastic_expanded_world_enabled()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _compute_elastic_expert_layout(
|
def _compute_elastic_expert_layout(
|
||||||
@@ -156,7 +159,6 @@ class ExpertLocationMetadata:
|
|||||||
)
|
)
|
||||||
assert physical_to_logical_map.shape[-1] == common["num_physical_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,
|
|
||||||
physical_to_logical_map=physical_to_logical_map,
|
physical_to_logical_map=physical_to_logical_map,
|
||||||
num_logical_experts=model_config_for_expert_location.num_logical_experts,
|
num_logical_experts=model_config_for_expert_location.num_logical_experts,
|
||||||
ep_size=common["ep_size"],
|
ep_size=common["ep_size"],
|
||||||
@@ -164,7 +166,6 @@ class ExpertLocationMetadata:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return ExpertLocationMetadata._init_raw(
|
return ExpertLocationMetadata._init_raw(
|
||||||
server_args=server_args,
|
|
||||||
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,
|
||||||
@@ -185,6 +186,8 @@ class ExpertLocationMetadata:
|
|||||||
logical_count = logical_count.unsqueeze(0)
|
logical_count = logical_count.unsqueeze(0)
|
||||||
logical_count = logical_count.to(server_args.device)
|
logical_count = logical_count.to(server_args.device)
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
common = ExpertLocationMetadata._init_common(server_args, model_config)
|
common = ExpertLocationMetadata._init_common(server_args, model_config)
|
||||||
|
|
||||||
if common is None:
|
if common is None:
|
||||||
@@ -193,7 +196,7 @@ class ExpertLocationMetadata:
|
|||||||
model_config_for_expert_location = common["model_config_for_expert_location"]
|
model_config_for_expert_location = common["model_config_for_expert_location"]
|
||||||
num_physical_experts = common["num_physical_experts"]
|
num_physical_experts = common["num_physical_experts"]
|
||||||
num_groups = model_config_for_expert_location.num_groups
|
num_groups = model_config_for_expert_location.num_groups
|
||||||
num_nodes = 1 if use_flat_topology else server_args.nnodes
|
num_nodes = 1 if use_flat_topology else get_parallel().nnodes
|
||||||
|
|
||||||
from sglang.srt.eplb import eplb_algorithms
|
from sglang.srt.eplb import eplb_algorithms
|
||||||
|
|
||||||
@@ -213,7 +216,6 @@ class ExpertLocationMetadata:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return ExpertLocationMetadata._init_raw(
|
return ExpertLocationMetadata._init_raw(
|
||||||
server_args=server_args,
|
|
||||||
ep_size=common["ep_size"],
|
ep_size=common["ep_size"],
|
||||||
physical_to_logical_map=physical_to_logical_map.to(server_args.device),
|
physical_to_logical_map=physical_to_logical_map.to(server_args.device),
|
||||||
logical_to_all_physical_map=logical_to_all_physical_map.to(
|
logical_to_all_physical_map=logical_to_all_physical_map.to(
|
||||||
@@ -223,6 +225,8 @@ class ExpertLocationMetadata:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _init_common(server_args: ServerArgs, model_config: ModelConfig):
|
def _init_common(server_args: ServerArgs, model_config: ModelConfig):
|
||||||
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
|
|
||||||
model_config_for_expert_location = (
|
model_config_for_expert_location = (
|
||||||
ModelConfigForExpertLocation.from_model_config(model_config)
|
ModelConfigForExpertLocation.from_model_config(model_config)
|
||||||
)
|
)
|
||||||
@@ -232,16 +236,17 @@ class ExpertLocationMetadata:
|
|||||||
|
|
||||||
base_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
|
+ get_exec().moe.ep_num_redundant_experts
|
||||||
)
|
)
|
||||||
ep_size = server_args.ep_size
|
# elastic-EP scale-up rewrites ep_size on the published config
|
||||||
|
ep_size = get_parallel().ep_size
|
||||||
num_physical_experts = base_num_physical_experts
|
num_physical_experts = base_num_physical_experts
|
||||||
initial_ep_size = server_args.elastic_ep_initial_size
|
initial_ep_size = get_parallel().elastic_ep_initial_size
|
||||||
if initial_ep_size is not None:
|
if initial_ep_size is not None:
|
||||||
if server_args.ep_join_mode == "scale":
|
if get_exec().moe.ep_join_mode == "scale":
|
||||||
ep_size = max(
|
ep_size = max(
|
||||||
ep_size,
|
ep_size,
|
||||||
server_args.ep_join_rank_offset + server_args.tp_size,
|
get_parallel().ep_join_rank_offset + server_args.tp_size,
|
||||||
)
|
)
|
||||||
num_physical_experts, num_local_physical_experts = (
|
num_physical_experts, num_local_physical_experts = (
|
||||||
_compute_elastic_expert_layout(
|
_compute_elastic_expert_layout(
|
||||||
@@ -264,12 +269,13 @@ class ExpertLocationMetadata:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _init_raw(
|
def _init_raw(
|
||||||
server_args: ServerArgs,
|
|
||||||
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,
|
moe_ep_rank: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
from sglang.srt.runtime_context import get_exec
|
||||||
|
|
||||||
_, num_physical_experts = physical_to_logical_map.shape
|
_, num_physical_experts = physical_to_logical_map.shape
|
||||||
|
|
||||||
logical_to_all_physical_map_padded = F.pad(
|
logical_to_all_physical_map_padded = F.pad(
|
||||||
@@ -291,7 +297,6 @@ class ExpertLocationMetadata:
|
|||||||
ep_size=ep_size,
|
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,
|
|
||||||
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,
|
||||||
@@ -301,7 +306,7 @@ class ExpertLocationMetadata:
|
|||||||
else torch.distributed.get_rank() % ep_size
|
else torch.distributed.get_rank() % ep_size
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
if server_args.ep_dispatch_algorithm == "static"
|
if get_exec().moe.ep_dispatch_algorithm == "static"
|
||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -536,12 +541,13 @@ def broadcast_global_expert_location_metadata(
|
|||||||
|
|
||||||
|
|
||||||
def _compute_logical_to_all_physical_map(
|
def _compute_logical_to_all_physical_map(
|
||||||
server_args: ServerArgs,
|
|
||||||
physical_to_logical_map: torch.Tensor,
|
physical_to_logical_map: torch.Tensor,
|
||||||
num_logical_experts: int,
|
num_logical_experts: int,
|
||||||
ep_size: int,
|
ep_size: int,
|
||||||
moe_ep_rank: int,
|
moe_ep_rank: int,
|
||||||
):
|
):
|
||||||
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
|
|
||||||
# This is rarely called, so we use for loops for maximum clarity
|
# This is rarely called, so we use for loops for maximum clarity
|
||||||
|
|
||||||
num_layers, num_physical_experts = physical_to_logical_map.shape
|
num_layers, num_physical_experts = physical_to_logical_map.shape
|
||||||
@@ -564,11 +570,13 @@ def _compute_logical_to_all_physical_map(
|
|||||||
# without an a2a backend, where all EP ranks must agree on the pick: this
|
# without an a2a backend, where all EP ranks must agree on the pick: this
|
||||||
# collapse is per-rank, and the full candidate list is what lets the dispatch
|
# collapse is per-rank, and the full candidate list is what lets the dispatch
|
||||||
# spread a hot expert over its replicas. See ExpertLocationDispatchInfo.
|
# spread a hot expert over its replicas. See ExpertLocationDispatchInfo.
|
||||||
if moe_ep_rank is not None and server_args.moe_a2a_backend != "none":
|
if moe_ep_rank is not None and get_exec().moe.moe_a2a_backend != "none":
|
||||||
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)
|
prefer_same_node = _prefer_same_node_experts()
|
||||||
num_gpus_per_node = (
|
num_gpus_per_node = (
|
||||||
server_args.ep_size // server_args.nnodes if prefer_same_node else None
|
get_parallel().ep_size // get_parallel().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
|
||||||
@@ -614,22 +622,23 @@ def _pad_nested_array(arr, pad_value):
|
|||||||
|
|
||||||
# TODO optimize performance (rewrite and/or run in separate process with overlap)
|
# TODO optimize performance (rewrite and/or run in separate process with overlap)
|
||||||
def compute_logical_to_rank_dispatch_physical_map(
|
def compute_logical_to_rank_dispatch_physical_map(
|
||||||
server_args: ServerArgs,
|
|
||||||
logical_to_all_physical_map: torch.Tensor,
|
logical_to_all_physical_map: torch.Tensor,
|
||||||
ep_size: int,
|
ep_size: int,
|
||||||
num_physical_experts: int,
|
num_physical_experts: int,
|
||||||
ep_rank: int,
|
ep_rank: int,
|
||||||
seed: int = 42,
|
seed: int = 42,
|
||||||
):
|
):
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
r = random.Random(seed)
|
r = random.Random(seed)
|
||||||
|
|
||||||
device = logical_to_all_physical_map.device
|
device = logical_to_all_physical_map.device
|
||||||
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
|
||||||
prefer_same_node = _prefer_same_node_experts(server_args)
|
prefer_same_node = _prefer_same_node_experts()
|
||||||
num_gpus_per_node = (
|
num_gpus_per_node = (
|
||||||
server_args.ep_size // server_args.nnodes if prefer_same_node else None
|
get_parallel().ep_size // get_parallel().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
|
||||||
|
|||||||
@@ -98,14 +98,14 @@ def is_dsa_enable_prefill_cp():
|
|||||||
def is_dsa_prefill_cp_in_seq_split():
|
def is_dsa_prefill_cp_in_seq_split():
|
||||||
return (
|
return (
|
||||||
is_dsa_enable_prefill_cp()
|
is_dsa_enable_prefill_cp()
|
||||||
and get_server_args().dsa_prefill_cp_mode == "in-seq-split"
|
and get_parallel().dsa_prefill_cp_mode == "in-seq-split"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def is_dsa_prefill_cp_round_robin_split():
|
def is_dsa_prefill_cp_round_robin_split():
|
||||||
return (
|
return (
|
||||||
is_dsa_enable_prefill_cp()
|
is_dsa_enable_prefill_cp()
|
||||||
and get_server_args().dsa_prefill_cp_mode == "round-robin-split"
|
and get_parallel().dsa_prefill_cp_mode == "round-robin-split"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -73,13 +73,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
check_cuda_graph_backend,
|
check_cuda_graph_backend,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_spec
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_spec,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -275,7 +269,7 @@ class AttnTpContext:
|
|||||||
def init_context(self, q_lora_rank, is_dsa):
|
def init_context(self, q_lora_rank, is_dsa):
|
||||||
self.is_dsa = is_dsa
|
self.is_dsa = is_dsa
|
||||||
self.allow_input_scattered = (
|
self.allow_input_scattered = (
|
||||||
get_server_args().enable_attn_tp_input_scattered
|
get_parallel().enable_attn_tp_input_scattered
|
||||||
and (_is_cuda or _is_npu)
|
and (_is_cuda or _is_npu)
|
||||||
and q_lora_rank is not None
|
and q_lora_rank is not None
|
||||||
and not is_dsa
|
and not is_dsa
|
||||||
@@ -286,7 +280,7 @@ class AttnTpContext:
|
|||||||
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE)
|
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE)
|
||||||
and get_spec().speculative_algorithm != "EAGLE3"
|
and get_spec().speculative_algorithm != "EAGLE3"
|
||||||
)
|
)
|
||||||
if get_server_args().enable_attn_tp_input_scattered:
|
if get_parallel().enable_attn_tp_input_scattered:
|
||||||
if not self.allow_input_scattered:
|
if not self.allow_input_scattered:
|
||||||
logging.info(
|
logging.info(
|
||||||
"attn_tp_input_scattered is not enabled while other conditions are not met"
|
"attn_tp_input_scattered is not enabled while other conditions are not met"
|
||||||
@@ -444,11 +438,11 @@ class LayerScatterModes:
|
|||||||
|
|
||||||
|
|
||||||
def enable_moe_dense_fully_dp():
|
def enable_moe_dense_fully_dp():
|
||||||
return get_server_args().moe_dense_tp_size == 1
|
return get_parallel().moe_dense_tp_size == 1
|
||||||
|
|
||||||
|
|
||||||
def enable_dwdp():
|
def enable_dwdp():
|
||||||
return get_server_args().dwdp_size > 1
|
return get_parallel().dwdp_size > 1
|
||||||
|
|
||||||
|
|
||||||
class LayerCommunicator:
|
class LayerCommunicator:
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ import torch
|
|||||||
from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
|
from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
|
||||||
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -51,7 +51,7 @@ class CpDecodeAttnTpContext:
|
|||||||
"""Slices replicated attention weights across CP ranks during decode."""
|
"""Slices replicated attention weights across CP ranks during decode."""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
enable_attn_tp = get_server_args().enable_cp_decode_attn_tp
|
enable_attn_tp = get_parallel().enable_cp_decode_attn_tp
|
||||||
|
|
||||||
if enable_attn_tp and get_parallel().attn_cp_size > 1:
|
if enable_attn_tp and get_parallel().attn_cp_size > 1:
|
||||||
self.decode_tp_rank = get_parallel().attn_cp_rank
|
self.decode_tp_rank = get_parallel().attn_cp_rank
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardBatch,
|
ForwardBatch,
|
||||||
ForwardMode,
|
ForwardMode,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_npu,
|
is_npu,
|
||||||
@@ -349,7 +349,7 @@ class LogitsProcessor(nn.Module):
|
|||||||
self.config = config
|
self.config = config
|
||||||
self.vocab_size = config.vocab_size
|
self.vocab_size = config.vocab_size
|
||||||
self.logit_scale = logit_scale
|
self.logit_scale = logit_scale
|
||||||
self.use_attn_tp_group = get_server_args().enable_dp_lm_head
|
self.use_attn_tp_group = get_parallel().enable_dp_lm_head
|
||||||
self.use_fp32_lm_head = get_exec().features.enable_fp32_lm_head
|
self.use_fp32_lm_head = get_exec().features.enable_fp32_lm_head
|
||||||
if self.use_attn_tp_group:
|
if self.use_attn_tp_group:
|
||||||
self.attn_tp_size = get_parallel().attn_tp_size
|
self.attn_tp_size = get_parallel().attn_tp_size
|
||||||
|
|||||||
@@ -260,12 +260,11 @@ class FusedMoE(torch.nn.Module):
|
|||||||
num_shared_slots = num_fused_shared_experts
|
num_shared_slots = num_fused_shared_experts
|
||||||
|
|
||||||
self._num_global_routed = num_experts - num_shared_slots
|
self._num_global_routed = num_experts - num_shared_slots
|
||||||
server_args = get_server_args()
|
|
||||||
if get_exec().moe.ep_join_mode == "scale":
|
if get_exec().moe.ep_join_mode == "scale":
|
||||||
storage_ep_size = server_args.elastic_ep_initial_size
|
storage_ep_size = get_parallel().elastic_ep_initial_size
|
||||||
assert storage_ep_size is not None
|
assert storage_ep_size is not None
|
||||||
self._expert_storage_rank = (
|
self._expert_storage_rank = (
|
||||||
server_args.ep_join_rank_offset + self.moe_ep_rank
|
get_parallel().ep_join_rank_offset + self.moe_ep_rank
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
storage_ep_size = self.moe_ep_size
|
storage_ep_size = self.moe_ep_size
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from sglang.srt.layers.moe.moe_runner.base import (
|
|||||||
register_fused_func,
|
register_fused_func,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils.common import log_info_on_rank0, print_warning_once
|
from sglang.srt.utils.common import log_info_on_rank0, print_warning_once
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -276,7 +277,9 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None:
|
|||||||
else:
|
else:
|
||||||
# Standard allgather path: the MoE sees up to dp_size local forwards
|
# Standard allgather path: the MoE sees up to dp_size local forwards
|
||||||
# gathered together, so scale the per-rank forward bound by dp_size.
|
# gathered together, so scale the per-rank forward bound by dp_size.
|
||||||
max_num_tokens = server_args.dp_size * server_args.cutedsl_moe_max_num_tokens()
|
max_num_tokens = (
|
||||||
|
get_parallel().dp_size * server_args.cutedsl_moe_max_num_tokens()
|
||||||
|
)
|
||||||
top_k = layer.top_k if layer.top_k is not None else layer.moe_runner_config.top_k
|
top_k = layer.top_k if layer.top_k is not None else layer.moe_runner_config.top_k
|
||||||
# inference_mode(False) ensures the wrapper's pre-allocated CUDA-graph
|
# inference_mode(False) ensures the wrapper's pre-allocated CUDA-graph
|
||||||
# buffers are normal tensors. This call typically happens inside
|
# buffers are normal tensors. This call typically happens inside
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.layers.moe.token_dispatcher.deepep import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
from sglang.srt.layers.moe.topk import TopKOutput
|
||||||
from sglang.srt.layers.moe.utils import DeepEPMode
|
from sglang.srt.layers.moe.utils import DeepEPMode
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from nixl_ep import Buffer
|
from nixl_ep import Buffer
|
||||||
@@ -127,9 +128,7 @@ class NixlEPBuffer:
|
|||||||
offset = ElasticEPStateManager.get_ep_join_rank_offset()
|
offset = ElasticEPStateManager.get_ep_join_rank_offset()
|
||||||
global_rank = rank + offset
|
global_rank = rank + offset
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_server_args
|
max_ep_size = get_parallel().max_ep_size or world_size
|
||||||
|
|
||||||
max_ep_size = get_server_args().max_ep_size or world_size
|
|
||||||
nixl_max_ranks = max_ep_size
|
nixl_max_ranks = max_ep_size
|
||||||
|
|
||||||
num_rdma_bytes = 0
|
num_rdma_bytes = 0
|
||||||
@@ -226,9 +225,8 @@ class _NixlEPDispatcherImplBase:
|
|||||||
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)
|
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
|
_max_ep = get_parallel().max_ep_size or self._active_world_size
|
||||||
self._mask_buffer = (
|
self._mask_buffer = (
|
||||||
torch.zeros(_max_ep, dtype=torch.int32, device="cuda")
|
torch.zeros(_max_ep, dtype=torch.int32, device="cuda")
|
||||||
if self.active_ranks is not None
|
if self.active_ranks is not None
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ from sglang.srt.layers.moe.utils import (
|
|||||||
DispatcherOutputDtype,
|
DispatcherOutputDtype,
|
||||||
get_deepep_output_dtype,
|
get_deepep_output_dtype,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
# Block size used by pplx-kernels for FP8 block-wise scales, matching the
|
# Block size used by pplx-kernels for FP8 block-wise scales, matching the
|
||||||
# DeepSeek / DeepGEMM block quantization convention.
|
# DeepSeek / DeepGEMM block quantization convention.
|
||||||
@@ -155,7 +155,7 @@ class PplxAllToAllManager:
|
|||||||
# pplx forces ep_size == world_size
|
# pplx forces ep_size == world_size
|
||||||
# with pp_size == 1 (enforced in _ensure_nvshmem), so the EP group spans
|
# with pp_size == 1 (enforced in _ensure_nvshmem), so the EP group spans
|
||||||
# a single node iff the whole job runs on one node.
|
# a single node iff the whole job runs on one node.
|
||||||
is_internode = get_server_args().nnodes > 1
|
is_internode = get_parallel().nnodes > 1
|
||||||
|
|
||||||
if is_internode:
|
if is_internode:
|
||||||
cls._all_to_all = AllToAll.internode(
|
cls._all_to_all = AllToAll.internode(
|
||||||
|
|||||||
@@ -506,7 +506,7 @@ def should_skip_post_experts_all_reduce(*, is_tp_path: bool) -> bool:
|
|||||||
"""
|
"""
|
||||||
if should_skip_mlp_all_reduce():
|
if should_skip_mlp_all_reduce():
|
||||||
return True
|
return True
|
||||||
if get_server_args().dwdp_size > 1:
|
if get_parallel().dwdp_size > 1:
|
||||||
return True
|
return True
|
||||||
if should_use_dp_reduce_scatterv():
|
if should_use_dp_reduce_scatterv():
|
||||||
return True
|
return True
|
||||||
|
|||||||
@@ -58,19 +58,19 @@ class ContextParallelMetadata:
|
|||||||
|
|
||||||
|
|
||||||
def is_prefill_context_parallel_enabled():
|
def is_prefill_context_parallel_enabled():
|
||||||
return get_server_args().enable_prefill_context_parallel
|
return get_parallel().enable_prefill_context_parallel
|
||||||
|
|
||||||
|
|
||||||
def is_prefill_cp_in_seq_split():
|
def is_prefill_cp_in_seq_split():
|
||||||
return (
|
return (
|
||||||
is_prefill_context_parallel_enabled()
|
is_prefill_context_parallel_enabled()
|
||||||
and get_server_args().prefill_cp_mode == "in-seq-split"
|
and get_parallel().prefill_cp_mode == "in-seq-split"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def is_mla_prefill_cp_enabled() -> bool:
|
def is_mla_prefill_cp_enabled() -> bool:
|
||||||
sa = get_server_args()
|
sa = get_server_args()
|
||||||
return sa.enable_prefill_context_parallel and sa.use_mla_backend()
|
return get_parallel().enable_prefill_context_parallel and sa.use_mla_backend()
|
||||||
|
|
||||||
|
|
||||||
def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None):
|
def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None):
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_mm,
|
get_mm,
|
||||||
get_model,
|
get_model,
|
||||||
get_observability,
|
get_observability,
|
||||||
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
get_serving,
|
get_serving,
|
||||||
get_spec,
|
get_spec,
|
||||||
@@ -1270,7 +1271,7 @@ class Scheduler(
|
|||||||
gloo_group=self.attn_tp_cpu_group,
|
gloo_group=self.attn_tp_cpu_group,
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
tp_size=self.ps.tp_size,
|
tp_size=self.ps.tp_size,
|
||||||
dp_size=self.server_args.dp_size,
|
dp_size=get_parallel().dp_size,
|
||||||
gpu_id=self.ps.gpu_id,
|
gpu_id=self.ps.gpu_id,
|
||||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
@@ -4462,7 +4463,7 @@ class Scheduler(
|
|||||||
|
|
||||||
old_ep_size = ElasticEPStateManager.get_effective_ep_size()
|
old_ep_size = ElasticEPStateManager.get_effective_ep_size()
|
||||||
new_ep_size = recv_req.new_ep_size
|
new_ep_size = recv_req.new_ep_size
|
||||||
max_ep_size = self.server_args.max_ep_size or old_ep_size
|
max_ep_size = get_parallel().max_ep_size or old_ep_size
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"[Elastic EP][scale] request received: new_ep_size=%d "
|
"[Elastic EP][scale] request received: new_ep_size=%d "
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.observability.metrics_collector import DPCooperationInfo
|
from sglang.srt.observability.metrics_collector import DPCooperationInfo
|
||||||
from sglang.srt.runtime_context import get_schedule
|
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||||
@@ -402,7 +402,7 @@ class SchedulerDPAttnAdapter:
|
|||||||
return prepare_mlp_sync_batch_raw(
|
return prepare_mlp_sync_batch_raw(
|
||||||
local_batch,
|
local_batch,
|
||||||
model_runner=self.model_runner,
|
model_runner=self.model_runner,
|
||||||
dp_size=self.server_args.dp_size,
|
dp_size=get_parallel().dp_size,
|
||||||
attn_tp_size=self.ps.attn_tp_size,
|
attn_tp_size=self.ps.attn_tp_size,
|
||||||
attn_cp_size=self.ps.attn_cp_size,
|
attn_cp_size=self.ps.attn_cp_size,
|
||||||
tp_group=self.tp_group,
|
tp_group=self.tp_group,
|
||||||
@@ -411,7 +411,7 @@ class SchedulerDPAttnAdapter:
|
|||||||
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
|
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
|
||||||
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||||
offload_tags=self.offload_tags,
|
offload_tags=self.offload_tags,
|
||||||
dwdp=self.server_args.dwdp_size > 1,
|
dwdp=get_parallel().dwdp_size > 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
def maybe_prepare_mlp_sync_batch(
|
def maybe_prepare_mlp_sync_batch(
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ from sglang.srt.managers.mm_utils import (
|
|||||||
has_shm_features,
|
has_shm_features,
|
||||||
unwrap_shm_features,
|
unwrap_shm_features,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_disagg
|
from sglang.srt.runtime_context import get_disagg, get_parallel
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
broadcast_pyobj,
|
broadcast_pyobj,
|
||||||
point_to_point_pyobj,
|
point_to_point_pyobj,
|
||||||
@@ -151,7 +151,7 @@ class SchedulerRequestReceiver:
|
|||||||
return recv_reqs
|
return recv_reqs
|
||||||
|
|
||||||
def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
|
def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
|
||||||
if self.server_args.enable_dp_attention:
|
if get_parallel().enable_dp_attention:
|
||||||
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
|
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
|
||||||
work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
|
work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
|
||||||
else:
|
else:
|
||||||
@@ -180,7 +180,7 @@ class SchedulerRequestReceiver:
|
|||||||
# 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 = (
|
_local_ctrl = (
|
||||||
self.server_args.enable_dp_attention_local_control_broadcast
|
get_parallel().enable_dp_attention_local_control_broadcast
|
||||||
or self.server_args.is_ep_scale_joiner
|
or self.server_args.is_ep_scale_joiner
|
||||||
)
|
)
|
||||||
if _local_ctrl:
|
if _local_ctrl:
|
||||||
@@ -258,7 +258,7 @@ class SchedulerRequestReceiver:
|
|||||||
# peer ranks may still be unpickling ShmPointerMMData
|
# peer ranks may still be unpickling ShmPointerMMData
|
||||||
# (-> shm_open). Synchronize the same CPU groups that carried
|
# (-> shm_open). Synchronize the same CPU groups that carried
|
||||||
# SHM-backed work requests before materialize() unlinks them.
|
# SHM-backed work requests before materialize() unlinks them.
|
||||||
if self.server_args.enable_dp_attention:
|
if get_parallel().enable_dp_attention:
|
||||||
if self.ps.attn_tp_size > 1:
|
if self.ps.attn_tp_size > 1:
|
||||||
barrier(group=self.attn_tp_cpu_group)
|
barrier(group=self.attn_tp_cpu_group)
|
||||||
if self.ps.attn_cp_size > 1:
|
if self.ps.attn_cp_size > 1:
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
PPProxyTensors,
|
PPProxyTensors,
|
||||||
)
|
)
|
||||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
from sglang.srt.observability.req_time_stats import set_time_batch
|
||||||
from sglang.srt.runtime_context import get_disagg
|
from sglang.srt.runtime_context import get_disagg, get_parallel
|
||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
|
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
|
||||||
from sglang.srt.utils.common import get_device_module, is_xpu
|
from sglang.srt.utils.common import get_device_module, is_xpu
|
||||||
@@ -123,7 +123,7 @@ class SchedulerPPMixin:
|
|||||||
next_pp_outputs = None
|
next_pp_outputs = None
|
||||||
next_batch_result = None
|
next_batch_result = None
|
||||||
d2h_event = None
|
d2h_event = None
|
||||||
if self.server_args.pp_async_batch_depth > 0:
|
if get_parallel().pp_async_batch_depth > 0:
|
||||||
next_pp_outputs, next_batch_result, d2h_event = (
|
next_pp_outputs, next_batch_result, d2h_event = (
|
||||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||||
next_first_rank_mb_id,
|
next_first_rank_mb_id,
|
||||||
@@ -139,7 +139,7 @@ class SchedulerPPMixin:
|
|||||||
self.mb_metadata,
|
self.mb_metadata,
|
||||||
self.last_rank_comm_queue,
|
self.last_rank_comm_queue,
|
||||||
)
|
)
|
||||||
if self.server_args.pp_async_batch_depth == 0:
|
if get_parallel().pp_async_batch_depth == 0:
|
||||||
next_pp_outputs, next_batch_result, d2h_event = (
|
next_pp_outputs, next_batch_result, d2h_event = (
|
||||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||||
next_first_rank_mb_id,
|
next_first_rank_mb_id,
|
||||||
@@ -269,7 +269,7 @@ class SchedulerPPMixin:
|
|||||||
server_is_idle = False
|
server_is_idle = False
|
||||||
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
||||||
|
|
||||||
if self.server_args.pp_async_batch_depth > 0:
|
if get_parallel().pp_async_batch_depth > 0:
|
||||||
next_pp_outputs, next_batch_result, d2h_event = (
|
next_pp_outputs, next_batch_result, d2h_event = (
|
||||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||||
next_first_rank_mb_id,
|
next_first_rank_mb_id,
|
||||||
@@ -285,7 +285,7 @@ class SchedulerPPMixin:
|
|||||||
self.mb_metadata,
|
self.mb_metadata,
|
||||||
self.last_rank_comm_queue,
|
self.last_rank_comm_queue,
|
||||||
)
|
)
|
||||||
if self.server_args.pp_async_batch_depth == 0:
|
if get_parallel().pp_async_batch_depth == 0:
|
||||||
next_pp_outputs, next_batch_result, d2h_event = (
|
next_pp_outputs, next_batch_result, d2h_event = (
|
||||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||||
next_first_rank_mb_id,
|
next_first_rank_mb_id,
|
||||||
@@ -428,7 +428,7 @@ class SchedulerPPMixin:
|
|||||||
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
||||||
|
|
||||||
# early send output if possible
|
# early send output if possible
|
||||||
if self.server_args.pp_async_batch_depth > 0:
|
if get_parallel().pp_async_batch_depth > 0:
|
||||||
next_pp_outputs, next_batch_result, d2h_event = (
|
next_pp_outputs, next_batch_result, d2h_event = (
|
||||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||||
next_first_rank_mb_id,
|
next_first_rank_mb_id,
|
||||||
@@ -446,7 +446,7 @@ class SchedulerPPMixin:
|
|||||||
self.last_rank_comm_queue,
|
self.last_rank_comm_queue,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.server_args.pp_async_batch_depth == 0:
|
if get_parallel().pp_async_batch_depth == 0:
|
||||||
next_pp_outputs, next_batch_result, d2h_event = (
|
next_pp_outputs, next_batch_result, d2h_event = (
|
||||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||||
next_first_rank_mb_id,
|
next_first_rank_mb_id,
|
||||||
@@ -557,10 +557,10 @@ class SchedulerPPMixin:
|
|||||||
self.on_idle()
|
self.on_idle()
|
||||||
|
|
||||||
def init_pp_loop_state(self: Scheduler):
|
def init_pp_loop_state(self: Scheduler):
|
||||||
self.pp_loop_size: int = self.ps.pp_size + self.server_args.pp_async_batch_depth
|
self.pp_loop_size: int = self.ps.pp_size + get_parallel().pp_async_batch_depth
|
||||||
# In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation.
|
# In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation.
|
||||||
self.require_attn_tp_allgather = (
|
self.require_attn_tp_allgather = (
|
||||||
not self.server_args.enable_dsa_prefill_context_parallel
|
not get_parallel().enable_dsa_prefill_context_parallel
|
||||||
)
|
)
|
||||||
self.mbs = [None] * self.pp_loop_size
|
self.mbs = [None] * self.pp_loop_size
|
||||||
self.last_mbs = [None] * self.pp_loop_size
|
self.last_mbs = [None] * self.pp_loop_size
|
||||||
|
|||||||
@@ -585,7 +585,7 @@ class CPUGraphRunner:
|
|||||||
model_runner.server_args.enable_profile_cuda_graph
|
model_runner.server_args.enable_profile_cuda_graph
|
||||||
)
|
)
|
||||||
self.tp_size = model_runner.server_args.tp_size
|
self.tp_size = model_runner.server_args.tp_size
|
||||||
self.dp_size = model_runner.server_args.dp_size
|
self.dp_size = get_parallel().dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = model_runner.server_args.pp_size
|
||||||
|
|
||||||
self.capture_forward_mode = ForwardMode.DECODE
|
self.capture_forward_mode = ForwardMode.DECODE
|
||||||
|
|||||||
@@ -407,19 +407,17 @@ class ModelRunner:
|
|||||||
):
|
):
|
||||||
return
|
return
|
||||||
|
|
||||||
join_effective_ep_size = self.server_args.ep_join_rank_offset + self.ps.tp_size
|
join_effective_ep_size = get_parallel().ep_join_rank_offset + self.ps.tp_size
|
||||||
dist.barrier(group=self.tp_group.cpu_group)
|
dist.barrier(group=self.tp_group.cpu_group)
|
||||||
if self.ps.tp_rank == 0:
|
if self.ps.tp_rank == 0:
|
||||||
register_scale_cohort(
|
register_scale_cohort(
|
||||||
self.server_args.ep_join_rank_offset,
|
get_parallel().ep_join_rank_offset,
|
||||||
join_effective_ep_size,
|
join_effective_ep_size,
|
||||||
)
|
)
|
||||||
join_scale_process_group()
|
join_scale_process_group()
|
||||||
self.server_args.override(
|
get_context().override("elastic_ep.scale_join", ep_size=join_effective_ep_size)
|
||||||
"elastic_ep.scale_join", ep_size=join_effective_ep_size
|
|
||||||
)
|
|
||||||
|
|
||||||
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
global_ep_rank = self.ps.tp_rank + get_parallel().ep_join_rank_offset
|
||||||
broadcast_global_expert_location_metadata(
|
broadcast_global_expert_location_metadata(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
moe_ep_rank=global_ep_rank,
|
moe_ep_rank=global_ep_rank,
|
||||||
@@ -443,9 +441,7 @@ class ModelRunner:
|
|||||||
new_dp_size=join_effective_ep_size,
|
new_dp_size=join_effective_ep_size,
|
||||||
new_dp_rank=global_ep_rank,
|
new_dp_rank=global_ep_rank,
|
||||||
)
|
)
|
||||||
self.server_args.override(
|
get_context().override("elastic_ep.scale_join", dp_size=join_effective_ep_size)
|
||||||
"elastic_ep.scale_join", dp_size=join_effective_ep_size
|
|
||||||
)
|
|
||||||
if self.eplb_manager is not None:
|
if self.eplb_manager is not None:
|
||||||
self.eplb_manager.disable_rebalance(
|
self.eplb_manager.disable_rebalance(
|
||||||
"EPLB rebalance is disabled while elastic EP scale-up "
|
"EPLB rebalance is disabled while elastic EP scale-up "
|
||||||
@@ -622,7 +618,7 @@ class ModelRunner:
|
|||||||
if self.is_draft_worker:
|
if self.is_draft_worker:
|
||||||
return
|
return
|
||||||
expert_rank = self.ps.moe_ep_rank + (
|
expert_rank = self.ps.moe_ep_rank + (
|
||||||
self.server_args.ep_join_rank_offset
|
get_parallel().ep_join_rank_offset
|
||||||
if self.server_args.is_ep_scale_joiner
|
if self.server_args.is_ep_scale_joiner
|
||||||
else 0
|
else 0
|
||||||
)
|
)
|
||||||
@@ -802,7 +798,7 @@ class ModelRunner:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
tp_group=(
|
tp_group=(
|
||||||
self.attention_tp_group.cpu_group
|
self.attention_tp_group.cpu_group
|
||||||
if self.server_args.enable_dp_attention
|
if get_parallel().enable_dp_attention
|
||||||
else self.tp_group.cpu_group
|
else self.tp_group.cpu_group
|
||||||
),
|
),
|
||||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
||||||
@@ -833,7 +829,7 @@ class ModelRunner:
|
|||||||
def post_capture_elastic_ep_recover(self):
|
def post_capture_elastic_ep_recover(self):
|
||||||
join_process_groups()
|
join_process_groups()
|
||||||
|
|
||||||
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
global_ep_rank = self.ps.tp_rank + get_parallel().ep_join_rank_offset
|
||||||
broadcast_global_expert_location_metadata(
|
broadcast_global_expert_location_metadata(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
moe_ep_rank=global_ep_rank,
|
moe_ep_rank=global_ep_rank,
|
||||||
@@ -870,7 +866,7 @@ class ModelRunner:
|
|||||||
self.prefill_attention_backend_str = backends.prefill_attention_backend_str
|
self.prefill_attention_backend_str = backends.prefill_attention_backend_str
|
||||||
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
||||||
|
|
||||||
if self.server_args.dcp_size > 1 and self.server_args.dcp_replicate_q_proj:
|
if self.server_args.dcp_size > 1 and get_parallel().dcp_replicate_q_proj:
|
||||||
self._prepare_replicated_q_proj()
|
self._prepare_replicated_q_proj()
|
||||||
|
|
||||||
def _prepare_replicated_q_proj(self) -> None:
|
def _prepare_replicated_q_proj(self) -> None:
|
||||||
@@ -1107,7 +1103,7 @@ class ModelRunner:
|
|||||||
def maybe_init_dwdp(self):
|
def maybe_init_dwdp(self):
|
||||||
if self.is_draft_worker:
|
if self.is_draft_worker:
|
||||||
return
|
return
|
||||||
if self.server_args.dwdp_size <= 1:
|
if get_parallel().dwdp_size <= 1:
|
||||||
return
|
return
|
||||||
from sglang.srt.layers.moe.dwdp import DwdpManager
|
from sglang.srt.layers.moe.dwdp import DwdpManager
|
||||||
|
|
||||||
@@ -1658,9 +1654,9 @@ class ModelRunner:
|
|||||||
if added <= 0:
|
if added <= 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
initial_ep_size = self.server_args.elastic_ep_initial_size
|
initial_ep_size = get_parallel().elastic_ep_initial_size
|
||||||
assert initial_ep_size is not None
|
assert initial_ep_size is not None
|
||||||
self.server_args.override("elastic_ep.scale", ep_size=effective_size)
|
get_context().override("elastic_ep.scale", ep_size=effective_size)
|
||||||
|
|
||||||
expanded_p2l = append_trivial_expert_slots(
|
expanded_p2l = append_trivial_expert_slots(
|
||||||
metadata.physical_to_logical_map,
|
metadata.physical_to_logical_map,
|
||||||
@@ -1677,7 +1673,7 @@ class ModelRunner:
|
|||||||
set_global_expert_location_metadata(new_metadata, allow_overwrite=True)
|
set_global_expert_location_metadata(new_metadata, allow_overwrite=True)
|
||||||
|
|
||||||
def _elastic_global_rank(self) -> int:
|
def _elastic_global_rank(self) -> int:
|
||||||
return self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
return self.ps.tp_rank + get_parallel().ep_join_rank_offset
|
||||||
|
|
||||||
def _rearm_eplb_after_elastic_scale(self) -> None:
|
def _rearm_eplb_after_elastic_scale(self) -> None:
|
||||||
if self.eplb_manager is None:
|
if self.eplb_manager is None:
|
||||||
@@ -1775,7 +1771,7 @@ class ModelRunner:
|
|||||||
new_dp_size=target_size,
|
new_dp_size=target_size,
|
||||||
new_dp_rank=self._elastic_global_rank(),
|
new_dp_rank=self._elastic_global_rank(),
|
||||||
)
|
)
|
||||||
self.server_args.override("elastic_ep.scale", dp_size=target_size)
|
get_context().override("elastic_ep.scale", dp_size=target_size)
|
||||||
|
|
||||||
ElasticEPStateManager.mark_syncing_new_world()
|
ElasticEPStateManager.mark_syncing_new_world()
|
||||||
self._elastic_scale_ready_barrier(
|
self._elastic_scale_ready_barrier(
|
||||||
|
|||||||
+3
-3
@@ -11,7 +11,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
|||||||
RemoteInstanceWeightLoaderBackend,
|
RemoteInstanceWeightLoaderBackend,
|
||||||
register_memory_region,
|
register_memory_region,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_model
|
from sglang.srt.runtime_context import get_model, get_parallel
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||||
|
|
||||||
@@ -76,11 +76,11 @@ class RemoteInstanceWeightTransporter:
|
|||||||
"""
|
"""
|
||||||
import requests as http_requests
|
import requests as http_requests
|
||||||
|
|
||||||
if self.server_args.dist_init_addr:
|
if get_parallel().dist_init_addr:
|
||||||
# Multi-node: bootstrap server is on the head node (node_rank==0).
|
# Multi-node: bootstrap server is on the head node (node_rank==0).
|
||||||
# Derive host from dist_init_addr (shared across all nodes).
|
# Derive host from dist_init_addr (shared across all nodes).
|
||||||
bootstrap_host = (
|
bootstrap_host = (
|
||||||
NetworkAddress.parse(self.server_args.dist_init_addr).resolved().host
|
NetworkAddress.parse(get_parallel().dist_init_addr).resolved().host
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
bootstrap_host = "127.0.0.1"
|
bootstrap_host = "127.0.0.1"
|
||||||
|
|||||||
@@ -197,7 +197,8 @@ class BaseRunner(ABC):
|
|||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.server_args.tp_size
|
self.tp_size = model_runner.server_args.tp_size
|
||||||
self.dp_size = model_runner.server_args.dp_size
|
# elastic-EP scale-up rewrites dp_size on the published config
|
||||||
|
self.dp_size = get_parallel().dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = model_runner.server_args.pp_size
|
||||||
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
||||||
self.enable_return_hidden_states = (
|
self.enable_return_hidden_states = (
|
||||||
@@ -313,7 +314,7 @@ class BaseRunner(ABC):
|
|||||||
hidden_size=mr.model_config.hidden_size,
|
hidden_size=mr.model_config.hidden_size,
|
||||||
vocab_size=mr.model_config.vocab_size,
|
vocab_size=mr.model_config.vocab_size,
|
||||||
dtype=mr.model_config.dtype,
|
dtype=mr.model_config.dtype,
|
||||||
dp_size=mr.server_args.dp_size,
|
dp_size=get_parallel().dp_size,
|
||||||
pp_size=mr.server_args.pp_size,
|
pp_size=mr.server_args.pp_size,
|
||||||
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
||||||
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
|
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
|
||||||
@@ -483,7 +484,7 @@ class BaseRunner(ABC):
|
|||||||
assert require_mlp_tp_gather_ or require_attn_tp_gather_
|
assert require_mlp_tp_gather_ or require_attn_tp_gather_
|
||||||
|
|
||||||
if require_mlp_tp_gather_:
|
if require_mlp_tp_gather_:
|
||||||
global_num_tokens_cpu = [num_tokens] * mr.server_args.dp_size
|
global_num_tokens_cpu = [num_tokens] * get_parallel().dp_size
|
||||||
elif require_attn_tp_gather_:
|
elif require_attn_tp_gather_:
|
||||||
global_num_tokens_cpu = [num_tokens]
|
global_num_tokens_cpu = [num_tokens]
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -335,7 +335,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
self.moe_fusions = self.model_runner.moe_fusions
|
self.moe_fusions = self.model_runner.moe_fusions
|
||||||
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
|
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
|
||||||
|
|
||||||
self.dp_size = model_runner.server_args.dp_size
|
self.dp_size = get_parallel().dp_size
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||||
|
|
||||||
|
|||||||
@@ -47,7 +47,12 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
|||||||
get_remote_instance_transfer_engine_info_per_rank,
|
get_remote_instance_transfer_engine_info_per_rank,
|
||||||
register_memory_region,
|
register_memory_region,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_model, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_model,
|
||||||
|
get_parallel,
|
||||||
|
get_server_args,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_available_gpu_memory
|
from sglang.srt.utils import get_available_gpu_memory
|
||||||
|
|
||||||
# Try to import accelerate (optional dependency)
|
# Try to import accelerate (optional dependency)
|
||||||
@@ -1747,9 +1752,9 @@ class PreshardedModelLoader(DefaultModelLoader):
|
|||||||
"dp": _safe(lambda: parallel.moe_dp_size),
|
"dp": _safe(lambda: parallel.moe_dp_size),
|
||||||
"ep": _safe(lambda: parallel.moe_ep_size),
|
"ep": _safe(lambda: parallel.moe_ep_size),
|
||||||
"pp": _safe(lambda: parallel.pp_size),
|
"pp": _safe(lambda: parallel.pp_size),
|
||||||
"moe_dense_tp_size": server_args.moe_dense_tp_size,
|
"moe_dense_tp_size": parallel.moe_dense_tp_size,
|
||||||
"moe_dp_size": server_args.moe_dp_size,
|
"moe_dp_size": server_args.moe_dp_size,
|
||||||
"enable_dp_lm_head": server_args.enable_dp_lm_head,
|
"enable_dp_lm_head": parallel.enable_dp_lm_head,
|
||||||
"enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head,
|
"enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head,
|
||||||
"quantization": model_config.quantization,
|
"quantization": model_config.quantization,
|
||||||
"model_dtype": str(model_config.dtype),
|
"model_dtype": str(model_config.dtype),
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
kv_cache_scales_loader,
|
kv_cache_scales_loader,
|
||||||
maybe_remap_kv_scale_name,
|
maybe_remap_kv_scale_name,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix, make_layers
|
from sglang.srt.utils import add_prefix, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -442,7 +442,7 @@ class ApertusForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
||||||
|
|||||||
@@ -46,7 +46,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
kv_cache_scales_loader,
|
kv_cache_scales_loader,
|
||||||
maybe_remap_kv_scale_name,
|
maybe_remap_kv_scale_name,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix, make_layers
|
from sglang.srt.utils import add_prefix, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -405,7 +405,7 @@ class ArceeForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
||||||
|
|||||||
@@ -77,13 +77,7 @@ from sglang.srt.models.utils import (
|
|||||||
create_fused_set_kv_buffer_arg,
|
create_fused_set_kv_buffer_arg,
|
||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
|
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
|
||||||
|
|
||||||
LoraConfig = None
|
LoraConfig = None
|
||||||
@@ -823,7 +817,7 @@ class BailingMoEForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -58,13 +58,7 @@ from sglang.srt.model_executor.runner import get_is_capture_mode
|
|||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
|
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_device, get_forward, get_parallel, get_stream
|
||||||
get_device,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
BumpAllocator,
|
BumpAllocator,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
@@ -1090,7 +1084,7 @@ class BailingMoELinearForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
params_dtype=torch.float32,
|
params_dtype=torch.float32,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ from sglang.srt.models.bailing_moe_linear import (
|
|||||||
BailingMoeV2_5ForCausalLM,
|
BailingMoeV2_5ForCausalLM,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import BumpAllocator, add_prefix
|
from sglang.srt.utils import BumpAllocator, add_prefix
|
||||||
|
|
||||||
LoraConfig = None
|
LoraConfig = None
|
||||||
@@ -208,7 +208,7 @@ class BailingMoeForCausalLMNextN(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("model.shared_head.head", prefix),
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid":
|
if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid":
|
||||||
|
|||||||
@@ -359,7 +359,7 @@ class DeepseekMLAForwardMixin:
|
|||||||
# --dcp-replicate-q-proj: project full-head Q locally from pre-gathered
|
# --dcp-replicate-q-proj: project full-head Q locally from pre-gathered
|
||||||
# weights and skip the per-layer Q all-gather (bf16 decode absorb only).
|
# weights and skip the per-layer Q all-gather (bf16 decode absorb only).
|
||||||
q_replicate_active = (
|
q_replicate_active = (
|
||||||
get_server_args().dcp_replicate_q_proj
|
get_parallel().dcp_replicate_q_proj
|
||||||
and _is_dcp_mla_decode_phase(forward_batch)
|
and _is_dcp_mla_decode_phase(forward_batch)
|
||||||
and not self.use_deep_gemm_bmm
|
and not self.use_deep_gemm_bmm
|
||||||
and self.w_kc_qrep is not None
|
and self.w_kc_qrep is not None
|
||||||
@@ -1029,7 +1029,7 @@ class DeepseekMLAForwardMixin:
|
|||||||
self.num_local_heads * get_parallel().attn_dcp_size,
|
self.num_local_heads * get_parallel().attn_dcp_size,
|
||||||
self.kv_lora_rank,
|
self.kv_lora_rank,
|
||||||
)
|
)
|
||||||
dcp_comm_backend = get_server_args().dcp_comm_backend
|
dcp_comm_backend = get_parallel().dcp_comm_backend
|
||||||
if dcp_comm_backend in ("a2a", "fi_a2a"):
|
if dcp_comm_backend in ("a2a", "fi_a2a"):
|
||||||
# A2A exchange of head partials + LSE, then local Triton combine.
|
# A2A exchange of head partials + LSE, then local Triton combine.
|
||||||
# MLA decode LSE is base-2 (FlashInfer-MLA/FlashMLA) -> base_on_e=False.
|
# MLA decode LSE is base-2 (FlashInfer-MLA/FlashMLA) -> base_on_e=False.
|
||||||
|
|||||||
@@ -59,12 +59,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
|||||||
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
|
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
|
||||||
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
|
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_model, get_parallel, get_spec
|
||||||
get_model,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_spec,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
||||||
|
|
||||||
|
|
||||||
@@ -388,7 +383,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("model.shared_head.head", prefix),
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -2938,7 +2938,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# ranks other than the last rank will have a placeholder layer
|
# ranks other than the last rank will have a placeholder layer
|
||||||
|
|||||||
@@ -138,13 +138,7 @@ from sglang.srt.models.deepseek_v2 import (
|
|||||||
_is_npu,
|
_is_npu,
|
||||||
_is_xpu,
|
_is_xpu,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_device, get_exec, get_forward, get_parallel
|
||||||
get_device,
|
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
)
|
|
||||||
|
|
||||||
if not _is_hip:
|
if not _is_hip:
|
||||||
from sglang.srt.layers.utils.cp_utils import (
|
from sglang.srt.layers.utils.cp_utils import (
|
||||||
@@ -2501,7 +2495,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.lm_head = PPMissingLayer()
|
self.lm_head = PPMissingLayer()
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||||
from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM
|
from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -233,7 +233,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("model.shared_head.head", prefix),
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
default_weight_loader,
|
default_weight_loader,
|
||||||
maybe_remap_kv_scale_name,
|
maybe_remap_kv_scale_name,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix, make_layers
|
from sglang.srt.utils import add_prefix, make_layers
|
||||||
from sglang.utils import get_exception_traceback, logger
|
from sglang.utils import get_exception_traceback, logger
|
||||||
|
|
||||||
@@ -439,7 +439,7 @@ class Exaone4ForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|||||||
@@ -62,12 +62,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -648,7 +643,7 @@ class ExaoneMoEForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
# For EAGLE3 support
|
# For EAGLE3 support
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
|||||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel
|
from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -63,7 +63,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -33,12 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_forward, get_parallel, get_stream
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -477,7 +472,7 @@ class FalconH1ForCausalLM(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
org_num_embeddings=config.vocab_size,
|
org_num_embeddings=config.vocab_size,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.lm_head = self.lm_head.float()
|
self.lm_head = self.lm_head.float()
|
||||||
self.lm_head_multiplier = config.lm_head_multiplier
|
self.lm_head_multiplier = config.lm_head_multiplier
|
||||||
|
|||||||
@@ -83,13 +83,7 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader
|
|||||||
from sglang.srt.models.deepseek_nextn import DeepseekV3ForCausalLMNextN
|
from sglang.srt.models.deepseek_nextn import DeepseekV3ForCausalLMNextN
|
||||||
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
|
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
|
||||||
from sglang.srt.models.utils import WeightsMapper, apply_qk_norm
|
from sglang.srt.models.utils import WeightsMapper, apply_qk_norm
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -1171,7 +1165,7 @@ class Glm4MoeForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -74,13 +74,7 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter
|
from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter
|
||||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
|
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
BumpAllocator,
|
BumpAllocator,
|
||||||
LazyValue,
|
LazyValue,
|
||||||
@@ -911,7 +905,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import (
|
|||||||
Glm4MoeLiteDecoderLayer,
|
Glm4MoeLiteDecoderLayer,
|
||||||
Glm4MoeLiteForCausalLM,
|
Glm4MoeLiteForCausalLM,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
|
from sglang.srt.runtime_context import get_exec, get_parallel, get_spec
|
||||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
|
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -151,7 +151,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("model.shared_head.head", prefix),
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
|
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
|
from sglang.srt.runtime_context import get_exec, get_parallel, get_spec
|
||||||
from sglang.srt.utils import add_prefix, is_npu
|
from sglang.srt.utils import add_prefix, is_npu
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -137,7 +137,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("model.shared_head.head", prefix),
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
|||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models.glm4_moe import Glm4MoeModel
|
from sglang.srt.models.glm4_moe import Glm4MoeModel
|
||||||
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel
|
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel
|
||||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel
|
||||||
from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
|
from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
|
||||||
from sglang.srt.utils.hf_transformers_utils import get_processor
|
from sglang.srt.utils.hf_transformers_utils import get_processor
|
||||||
|
|
||||||
@@ -69,7 +69,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# ranks other than the last rank will have a placeholder layer
|
# ranks other than the last rank will have a placeholder layer
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.models.glm4 import Glm4DecoderLayer
|
from sglang.srt.models.glm4 import Glm4DecoderLayer
|
||||||
from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration
|
from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -134,7 +134,7 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("model.shared_head.head", prefix),
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -259,7 +259,7 @@ class GptOssSparseMoeBlock(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
forward_batch: Optional[ForwardBatch] = None,
|
forward_batch: Optional[ForwardBatch] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if get_server_args().dwdp_size > 1:
|
if get_parallel().dwdp_size > 1:
|
||||||
return self.forward_dwdp(hidden_states)
|
return self.forward_dwdp(hidden_states)
|
||||||
|
|
||||||
if not get_moe_a2a_backend().is_deepep():
|
if not get_moe_a2a_backend().is_deepep():
|
||||||
@@ -787,7 +787,7 @@ class GptOssForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
# quant_config=quant_config,
|
# quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
self.capture_aux_hidden_states = False
|
self.capture_aux_hidden_states = False
|
||||||
|
|||||||
@@ -53,12 +53,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models.utils import apply_qk_norm
|
from sglang.srt.models.utils import apply_qk_norm
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import LazyValue, add_prefix, make_layers
|
from sglang.srt.utils import LazyValue, add_prefix, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -645,7 +640,7 @@ class LagunaForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.lm_head = PPMissingLayer()
|
self.lm_head = PPMissingLayer()
|
||||||
|
|||||||
@@ -76,13 +76,7 @@ from sglang.srt.models.utils import (
|
|||||||
create_fused_set_kv_buffer_arg,
|
create_fused_set_kv_buffer_arg,
|
||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
LazyValue,
|
LazyValue,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
@@ -830,7 +824,7 @@ class LLaDA2MoeModelLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config, return_full_logits=True)
|
self.logits_processor = LogitsProcessor(config, return_full_logits=True)
|
||||||
|
|
||||||
|
|||||||
@@ -53,7 +53,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
maybe_remap_kv_scale_name,
|
maybe_remap_kv_scale_name,
|
||||||
)
|
)
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix, is_cuda, is_npu, is_xpu, make_layers
|
from sglang.srt.utils import add_prefix, is_cuda, is_npu, is_xpu, make_layers
|
||||||
from sglang.utils import get_exception_traceback
|
from sglang.utils import get_exception_traceback
|
||||||
|
|
||||||
@@ -530,7 +530,7 @@ class LlamaForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
||||||
|
|||||||
@@ -88,7 +88,7 @@ from sglang.srt.model_loader.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
|
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
|
from sglang.srt.runtime_context import get_parallel, get_stream
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
BumpAllocator,
|
BumpAllocator,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
@@ -721,7 +721,7 @@ class LongcatFlashForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
self.capture_aux_hidden_states = False
|
self.capture_aux_hidden_states = False
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ from sglang.srt.models.utils import (
|
|||||||
create_fused_set_kv_buffer_arg,
|
create_fused_set_kv_buffer_arg,
|
||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
from sglang.srt.utils import add_prefix, is_cuda
|
from sglang.srt.utils import add_prefix, is_cuda
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
@@ -520,7 +520,7 @@ class MellumForCausalLM(Qwen3MoeForCausalLM):
|
|||||||
cfg.hidden_size,
|
cfg.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(cfg)
|
self.logits_processor = LogitsProcessor(cfg)
|
||||||
self.capture_aux_hidden_states = False
|
self.capture_aux_hidden_states = False
|
||||||
|
|||||||
@@ -78,12 +78,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
|
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
|
||||||
from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig
|
from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
LazyValue,
|
LazyValue,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
@@ -1192,7 +1187,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.lm_head = PPMissingLayer()
|
self.lm_head = PPMissingLayer()
|
||||||
|
|||||||
@@ -44,7 +44,7 @@ from sglang.srt.models.mimo_v2 import (
|
|||||||
MiMoV2MLP,
|
MiMoV2MLP,
|
||||||
load_mimo_v2_qkv_proj_weight,
|
load_mimo_v2_qkv_proj_weight,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
MiMoV2Config = None
|
MiMoV2Config = None
|
||||||
@@ -259,7 +259,7 @@ class MiMoV2MTP(MiMoV2ForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -80,7 +80,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP
|
from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
get_device_sm,
|
get_device_sm,
|
||||||
@@ -1453,7 +1453,7 @@ class MiniMaxM3SparseForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|||||||
@@ -120,7 +120,7 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module):
|
|||||||
text_config.hidden_size,
|
text_config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("language_model.lm_head", prefix),
|
prefix=add_prefix("language_model.lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.lm_head = PPMissingLayer()
|
self.lm_head = PPMissingLayer()
|
||||||
|
|||||||
@@ -89,12 +89,7 @@ from sglang.srt.models.nemotron_h_utils import (
|
|||||||
pad_to_original_num_tokens,
|
pad_to_original_num_tokens,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
get_current_device_stream_fast,
|
get_current_device_stream_fast,
|
||||||
@@ -944,7 +939,7 @@ class NemotronHForCausalLM(nn.Module):
|
|||||||
else lora_config.lora_vocab_padding_size
|
else lora_config.lora_vocab_padding_size
|
||||||
),
|
),
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ from sglang.srt.models.nemotron_h import (
|
|||||||
NemotronHMoEDecoderLayer,
|
NemotronHMoEDecoderLayer,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.nemotron_h_utils import is_attn_layer
|
from sglang.srt.models.nemotron_h_utils import is_attn_layer
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
|
|
||||||
@@ -338,7 +338,7 @@ class NemotronHForCausalLMMTP(NemotronHForCausalLM):
|
|||||||
self.config.hidden_size,
|
self.config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|||||||
@@ -92,12 +92,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -1028,7 +1023,7 @@ class Qwen2MoeForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
# For EAGLE3 support
|
# For EAGLE3 support
|
||||||
|
|||||||
@@ -33,12 +33,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP
|
from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP
|
||||||
from sglang.srt.models.qwen2 import Qwen2Model
|
from sglang.srt.models.qwen2 import Qwen2Model
|
||||||
from sglang.srt.models.utils import apply_qk_norm
|
from sglang.srt.models.utils import apply_qk_norm
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu
|
from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu
|
||||||
|
|
||||||
Qwen3Config = None
|
Qwen3Config = None
|
||||||
@@ -497,7 +492,7 @@ class Qwen3ForCausalLM(nn.Module):
|
|||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTe
|
|||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models import qwen3_5
|
from sglang.srt.models import qwen3_5
|
||||||
from sglang.srt.models.qwen2_moe import Qwen2MoeSparseMoeBlock
|
from sglang.srt.models.qwen2_moe import Qwen2MoeSparseMoeBlock
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import LazyValue, add_prefix
|
from sglang.srt.utils import LazyValue, add_prefix
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -73,7 +73,7 @@ class Qwen3_5ForCausalLM(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
org_num_embeddings=config.vocab_size,
|
org_num_embeddings=config.vocab_size,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.lm_head = PPMissingLayer()
|
self.lm_head = PPMissingLayer()
|
||||||
|
|||||||
@@ -72,13 +72,7 @@ from sglang.srt.models.utils import (
|
|||||||
create_fused_set_kv_buffer_arg,
|
create_fused_set_kv_buffer_arg,
|
||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
LazyValue,
|
LazyValue,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
@@ -966,7 +960,7 @@ class Qwen3MoeForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
self.capture_aux_hidden_states = False
|
self.capture_aux_hidden_states = False
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
|||||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM, Qwen3MoeModel
|
from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM, Qwen3MoeModel
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -63,7 +63,7 @@ class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -49,12 +49,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
sharded_weight_loader,
|
sharded_weight_loader,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
|
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_forward, get_parallel, get_stream
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
LazyValue,
|
LazyValue,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
@@ -1032,7 +1027,7 @@ class Qwen3NextForCausalLM(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
org_num_embeddings=config.vocab_size,
|
org_num_embeddings=config.vocab_size,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
# For EAGLE3 support
|
# For EAGLE3 support
|
||||||
|
|||||||
@@ -32,12 +32,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
|||||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel
|
from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_model, get_parallel, get_spec
|
||||||
get_model,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_spec,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import add_prefix, is_npu
|
from sglang.srt.utils import add_prefix, is_npu
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -85,7 +80,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("model.shared_head.head", prefix),
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
# Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so
|
# Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so
|
||||||
|
|||||||
@@ -73,7 +73,7 @@ from sglang.srt.multimodal.mm_utils import (
|
|||||||
run_dp_sharded_mrope_vision_model,
|
run_dp_sharded_mrope_vision_model,
|
||||||
)
|
)
|
||||||
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -1280,7 +1280,7 @@ class Qwen3VLForConditionalGeneration(nn.Module):
|
|||||||
self.config.vocab_size,
|
self.config.vocab_size,
|
||||||
self.config.hidden_size,
|
self.config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1231,7 +1231,7 @@ class SarvamMLAForCausalLM(nn.Module):
|
|||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
|||||||
@@ -41,13 +41,7 @@ from sglang.srt.models.utils import (
|
|||||||
create_fused_set_kv_buffer_arg,
|
create_fused_set_kv_buffer_arg,
|
||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -475,7 +469,7 @@ class SDARForCausalLM(nn.Module):
|
|||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -57,13 +57,7 @@ from sglang.srt.models.utils import (
|
|||||||
create_fused_set_kv_buffer_arg,
|
create_fused_set_kv_buffer_arg,
|
||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -562,7 +556,7 @@ class SDARMoeForCausalLM(nn.Module):
|
|||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -46,13 +46,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream
|
||||||
get_exec,
|
|
||||||
get_forward,
|
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
|
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
|
||||||
|
|
||||||
Step3p5Config = None
|
Step3p5Config = None
|
||||||
@@ -822,7 +816,7 @@ class Step3p5ForCausalLM(nn.Module):
|
|||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -125,12 +125,19 @@ class ParallelContext:
|
|||||||
|
|
||||||
def __getattr__(self, name):
|
def __getattr__(self, name):
|
||||||
# Reached only for names that are neither a live @property nor a slot:
|
# Reached only for names that are neither a live @property nor a slot:
|
||||||
# serve parallel config leaves from the published bag.
|
# serve parallel config leaves from the published bag. The body must
|
||||||
try:
|
# stay dynamo-traceable — config-leaf reads such as
|
||||||
config = object.__getattribute__(self, "_config")
|
# ``get_parallel().moe_dense_tp_size`` run inside compiled model
|
||||||
except AttributeError:
|
# forwards, and ``object.__getattribute__`` graph-breaks.
|
||||||
config = None
|
if name.startswith("_"):
|
||||||
if config is not None and name in config:
|
# No config leaf is underscored; this also breaks the recursion
|
||||||
|
# when the ``_config`` slot itself is still unset (pickle/copy
|
||||||
|
# protocols probe attributes before __init__ runs).
|
||||||
|
raise AttributeError(name)
|
||||||
|
config = self._config
|
||||||
|
# ``_fields`` is a plain ``__dict__`` entry on the bag; ``in`` on the
|
||||||
|
# dict avoids ``_ConfigBag.__contains__`` (not traceable).
|
||||||
|
if config is not None and name in config._fields:
|
||||||
return getattr(config, name)
|
return getattr(config, name)
|
||||||
detail = (
|
detail = (
|
||||||
"not a published parallel config leaf"
|
"not a published parallel config leaf"
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
compute_position,
|
compute_position,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
from sglang.srt.runtime_context import get_exec, get_parallel, get_spec
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
||||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||||
@@ -389,7 +389,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
self, batch: ScheduleBatch, on_publish
|
self, batch: ScheduleBatch, on_publish
|
||||||
) -> GenerationBatchResult:
|
) -> GenerationBatchResult:
|
||||||
if batch.forward_mode.is_idle():
|
if batch.forward_mode.is_idle():
|
||||||
if self.server_args.enable_dp_attention:
|
if get_parallel().enable_dp_attention:
|
||||||
self.target_worker.forward_batch_generation(
|
self.target_worker.forward_batch_generation(
|
||||||
batch, capture_hidden_mode=CaptureHiddenMode.FULL
|
batch, capture_hidden_mode=CaptureHiddenMode.FULL
|
||||||
)
|
)
|
||||||
@@ -457,7 +457,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
def _dp_verify_tier_num_tokens(self, batch: ScheduleBatch) -> Optional[int]:
|
def _dp_verify_tier_num_tokens(self, batch: ScheduleBatch) -> Optional[int]:
|
||||||
if not (
|
if not (
|
||||||
self._draft_is_moe
|
self._draft_is_moe
|
||||||
and self.server_args.enable_dp_attention
|
and get_parallel().enable_dp_attention
|
||||||
and batch.global_num_tokens is not None
|
and batch.global_num_tokens is not None
|
||||||
and self._verify_planner.is_compact_mode
|
and self._verify_planner.is_compact_mode
|
||||||
):
|
):
|
||||||
@@ -501,7 +501,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
if batch.forward_mode.is_idle():
|
if batch.forward_mode.is_idle():
|
||||||
self._observers.note_idle_decode_step()
|
self._observers.note_idle_decode_step()
|
||||||
if self.server_args.enable_dp_attention:
|
if get_parallel().enable_dp_attention:
|
||||||
if self._draft_is_moe:
|
if self._draft_is_moe:
|
||||||
self._proposer.run_idle_participation(batch)
|
self._proposer.run_idle_participation(batch)
|
||||||
self._verify_executor.run_idle_participation(
|
self._verify_executor.run_idle_participation(
|
||||||
@@ -563,7 +563,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
global_num_reqs = (
|
global_num_reqs = (
|
||||||
max(batch.global_num_tokens)
|
max(batch.global_num_tokens)
|
||||||
if self._draft_is_moe
|
if self._draft_is_moe
|
||||||
and self.server_args.enable_dp_attention
|
and get_parallel().enable_dp_attention
|
||||||
and batch.global_num_tokens is not None
|
and batch.global_num_tokens is not None
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
@@ -731,14 +731,14 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
# Chain layout only: step index = commit_lens - 1. A tree (topk > 1)
|
# Chain layout only: step index = commit_lens - 1. A tree (topk > 1)
|
||||||
# layout would need the accept-index mapping the shared spec_utils
|
# layout would need the accept-index mapping the shared spec_utils
|
||||||
# commit helper does.
|
# commit helper does.
|
||||||
assert self.server_args.speculative_eagle_topk in (None, 1)
|
assert get_spec().speculative_eagle_topk in (None, 1)
|
||||||
attn_backend = self.target_worker.model_runner.attn_backend
|
attn_backend = self.target_worker.model_runner.attn_backend
|
||||||
|
|
||||||
last_correct_step_indices = commit_lens.to(torch.int64) - 1
|
last_correct_step_indices = commit_lens.to(torch.int64) - 1
|
||||||
mamba_steps_to_track = None
|
mamba_steps_to_track = None
|
||||||
|
|
||||||
if batch.mamba_track_indices is not None:
|
if batch.mamba_track_indices is not None:
|
||||||
mamba_track_interval = self.server_args.mamba_track_interval
|
mamba_track_interval = get_exec().mamba.mamba_track_interval
|
||||||
to_track_mask = (
|
to_track_mask = (
|
||||||
seq_lens_pre_verify // mamba_track_interval
|
seq_lens_pre_verify // mamba_track_interval
|
||||||
!= seq_lens_post_verify // mamba_track_interval
|
!= seq_lens_post_verify // mamba_track_interval
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
|||||||
from sglang.srt.model_executor.runner_backend_utils import (
|
from sglang.srt.model_executor.runner_backend_utils import (
|
||||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_spec
|
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||||
from sglang.srt.speculative.multi_layer_eagle_utils import (
|
from sglang.srt.speculative.multi_layer_eagle_utils import (
|
||||||
@@ -150,7 +150,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.ps.tp_size
|
||||||
self.dp_size = model_runner.server_args.dp_size
|
self.dp_size = get_parallel().dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = model_runner.server_args.pp_size
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||||
|
|||||||
@@ -69,15 +69,14 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
|
|||||||
topk_size = model_config.hf_text_config.num_experts_per_tok
|
topk_size = model_config.hf_text_config.num_experts_per_tok
|
||||||
num_layers = model_config.hf_text_config.num_hidden_layers
|
num_layers = model_config.hf_text_config.num_hidden_layers
|
||||||
|
|
||||||
server_args = get_server_args()
|
|
||||||
# Scale by dp_size so the buffer covers the full DP-concatenated batch.
|
# Scale by dp_size so the buffer covers the full DP-concatenated batch.
|
||||||
# _get_local_slice indexes into [attention_dp_rank * cuda_graph_batch, ...)
|
# _get_local_slice indexes into [attention_dp_rank * cuda_graph_batch, ...)
|
||||||
# and otherwise overflows on dp_rank > 0 when max_running_requests >
|
# and otherwise overflows on dp_rank > 0 when max_running_requests >
|
||||||
# chunked_prefill_size.
|
# chunked_prefill_size.
|
||||||
# FIXME: spec decoding's num_verify_tokens is still not accounted for.
|
# FIXME: spec decoding's num_verify_tokens is still not accounted for.
|
||||||
max_batch_size = max(
|
max_batch_size = max(
|
||||||
get_schedule().chunked_prefill_size * server_args.dp_size,
|
get_schedule().chunked_prefill_size * get_parallel().dp_size,
|
||||||
max_running_requests * server_args.dp_size,
|
max_running_requests * get_parallel().dp_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
|
|||||||
@@ -3541,10 +3541,12 @@ def require_mlp_tp_gather(server_args: ServerArgs):
|
|||||||
Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups.
|
Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups.
|
||||||
"""
|
"""
|
||||||
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
||||||
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
|
|
||||||
if server_args.enable_dp_attention:
|
# elastic-EP scale-up rewrites dp_size on the published config
|
||||||
assert server_args.dp_size > 1, "dp_size must be greater than 1"
|
if get_parallel().enable_dp_attention:
|
||||||
if server_args.elastic_ep_backend is not None:
|
assert get_parallel().dp_size > 1, "dp_size must be greater than 1"
|
||||||
|
if get_exec().moe.elastic_ep_backend is not None:
|
||||||
from sglang.srt.elastic_ep.elastic_ep import (
|
from sglang.srt.elastic_ep.elastic_ep import (
|
||||||
elastic_expanded_world_enabled,
|
elastic_expanded_world_enabled,
|
||||||
)
|
)
|
||||||
@@ -3552,10 +3554,10 @@ def require_mlp_tp_gather(server_args: ServerArgs):
|
|||||||
if elastic_expanded_world_enabled():
|
if elastic_expanded_world_enabled():
|
||||||
return True
|
return True
|
||||||
if (
|
if (
|
||||||
server_args.moe_dense_tp_size is None
|
get_parallel().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
|
||||||
return True
|
return True
|
||||||
elif not server_args.enable_dp_lm_head:
|
elif not get_parallel().enable_dp_lm_head:
|
||||||
return True
|
return True
|
||||||
elif get_moe_a2a_backend().is_none():
|
elif get_moe_a2a_backend().is_none():
|
||||||
return True
|
return True
|
||||||
@@ -3571,8 +3573,8 @@ def require_mlp_tp_gather(server_args: ServerArgs):
|
|||||||
return True
|
return True
|
||||||
else:
|
else:
|
||||||
return (
|
return (
|
||||||
server_args.moe_dense_tp_size
|
get_parallel().moe_dense_tp_size
|
||||||
> server_args.tp_size // server_args.dp_size
|
> server_args.tp_size // get_parallel().dp_size
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return False
|
return False
|
||||||
@@ -3586,14 +3588,19 @@ def require_attn_tp_gather(server_args: ServerArgs):
|
|||||||
# and do not consume the upstream gathered_buffer. Without this, the
|
# and do not consume the upstream gathered_buffer. Without this, the
|
||||||
# cuda graph runner pads num_tokens to attn_tp_size, which can cause
|
# cuda graph runner pads num_tokens to attn_tp_size, which can cause
|
||||||
# autotuners to pick suboptimal kernel variants at small batches.
|
# autotuners to pick suboptimal kernel variants at small batches.
|
||||||
if server_args.disable_attn_tp_gather:
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
|
if get_parallel().disable_attn_tp_gather:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
||||||
|
|
||||||
if not get_moe_a2a_backend().is_none() or server_args.moe_dense_tp_size is not None:
|
if (
|
||||||
if server_args.enable_dp_attention:
|
not get_moe_a2a_backend().is_none()
|
||||||
return server_args.dp_size < server_args.tp_size
|
or get_parallel().moe_dense_tp_size is not None
|
||||||
|
):
|
||||||
|
if get_parallel().enable_dp_attention:
|
||||||
|
return get_parallel().dp_size < server_args.tp_size
|
||||||
else:
|
else:
|
||||||
return True
|
return True
|
||||||
else:
|
else:
|
||||||
@@ -3605,7 +3612,9 @@ def require_gathered_buffer(server_args: ServerArgs):
|
|||||||
|
|
||||||
|
|
||||||
def require_mlp_sync(server_args: ServerArgs):
|
def require_mlp_sync(server_args: ServerArgs):
|
||||||
return server_args.enable_dp_attention or require_gathered_buffer(server_args)
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
|
return get_parallel().enable_dp_attention or require_gathered_buffer(server_args)
|
||||||
|
|
||||||
|
|
||||||
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int:
|
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int:
|
||||||
|
|||||||
@@ -484,9 +484,7 @@ class IpcModelLoader(BaseModelLoader):
|
|||||||
|
|
||||||
ep_size = ps.moe_ep_size
|
ep_size = ps.moe_ep_size
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_server_args
|
dp_size = get_parallel().dp_size
|
||||||
|
|
||||||
dp_size = get_server_args().dp_size
|
|
||||||
|
|
||||||
quant_method, quant_config = self._resolve_engine_quant(model_config)
|
quant_method, quant_config = self._resolve_engine_quant(model_config)
|
||||||
|
|
||||||
|
|||||||
@@ -389,7 +389,6 @@ class TestAiterAllreduceFusionGate(CustomTestCase):
|
|||||||
tp_size=8,
|
tp_size=8,
|
||||||
):
|
):
|
||||||
"""Run the gate with the aiter branch isolated (flashinfer forced off)."""
|
"""Run the gate with the aiter branch isolated (flashinfer forced off)."""
|
||||||
server_args = types.SimpleNamespace(enable_aiter_allreduce_fusion=aiter_enabled)
|
|
||||||
a2a_backend = types.SimpleNamespace(is_none=lambda: a2a_is_none)
|
a2a_backend = types.SimpleNamespace(is_none=lambda: a2a_is_none)
|
||||||
|
|
||||||
with ExitStack() as stack:
|
with ExitStack() as stack:
|
||||||
@@ -417,10 +416,14 @@ class TestAiterAllreduceFusionGate(CustomTestCase):
|
|||||||
lambda: types.SimpleNamespace(tp_size=tp_world_size),
|
lambda: types.SimpleNamespace(tp_size=tp_world_size),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
# the gate reads get_exec().comm.enable_aiter_allreduce_fusion
|
||||||
|
from sglang.srt.runtime_context import get_context, get_flags
|
||||||
|
|
||||||
stack.enter_context(
|
stack.enter_context(
|
||||||
mock.patch.object(comm, "get_server_args", lambda: server_args)
|
get_context().override_server_args(
|
||||||
|
enable_aiter_allreduce_fusion=aiter_enabled
|
||||||
|
)
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags
|
|
||||||
|
|
||||||
stack.enter_context(get_flags().dp.override(enabled=dp_attention))
|
stack.enter_context(get_flags().dp.override(enabled=dp_attention))
|
||||||
stack.enter_context(
|
stack.enter_context(
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ from sglang.test.ci.ci_register import register_cpu_ci
|
|||||||
|
|
||||||
register_cpu_ci(est_time=7, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=7, suite="base-a-test-cpu")
|
||||||
|
|
||||||
import types
|
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -14,20 +13,26 @@ from sglang.srt.eplb.expert_location import (
|
|||||||
append_trivial_expert_slots,
|
append_trivial_expert_slots,
|
||||||
compute_logical_to_rank_dispatch_physical_map,
|
compute_logical_to_rank_dispatch_physical_map,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
def _make_server_args(ep_size: int, nnodes: int, moe_a2a_backend: str = "deepep"):
|
def _published(
|
||||||
"""Minimal server_args stub for expert placement tests.
|
ep_size: int,
|
||||||
|
nnodes: int,
|
||||||
|
moe_a2a_backend: str = "deepep",
|
||||||
|
ep_join_mode=None,
|
||||||
|
):
|
||||||
|
"""Scoped publish of the config the placement functions read.
|
||||||
|
|
||||||
`moe_a2a_backend` defaults to an a2a backend because these tests cover the
|
`moe_a2a_backend` defaults to an a2a backend because these tests cover the
|
||||||
rank-local collapse, which is skipped when there is no a2a backend.
|
rank-local collapse, which is skipped when there is no a2a backend.
|
||||||
"""
|
"""
|
||||||
return types.SimpleNamespace(
|
return get_context().override_server_args(
|
||||||
ep_size=ep_size,
|
ep_size=ep_size,
|
||||||
nnodes=nnodes,
|
nnodes=nnodes,
|
||||||
ep_join_mode=None,
|
|
||||||
moe_a2a_backend=moe_a2a_backend,
|
moe_a2a_backend=moe_a2a_backend,
|
||||||
|
ep_join_mode=ep_join_mode,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -74,7 +79,6 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase):
|
|||||||
NUM_LAYERS = 2
|
NUM_LAYERS = 2
|
||||||
|
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.server_args = _make_server_args(self.EP_SIZE, self.NNODES)
|
|
||||||
self.logical_to_all_physical = _make_logical_to_all_physical_map(
|
self.logical_to_all_physical = _make_logical_to_all_physical_map(
|
||||||
num_layers=self.NUM_LAYERS,
|
num_layers=self.NUM_LAYERS,
|
||||||
num_logical_experts=self.NUM_LOGICAL,
|
num_logical_experts=self.NUM_LOGICAL,
|
||||||
@@ -83,14 +87,14 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _call(self, ep_rank, seed=42):
|
def _call(self, ep_rank, seed=42):
|
||||||
return compute_logical_to_rank_dispatch_physical_map(
|
with _published(self.EP_SIZE, self.NNODES):
|
||||||
server_args=self.server_args,
|
return compute_logical_to_rank_dispatch_physical_map(
|
||||||
logical_to_all_physical_map=self.logical_to_all_physical.clone(),
|
logical_to_all_physical_map=self.logical_to_all_physical.clone(),
|
||||||
ep_size=self.EP_SIZE,
|
ep_size=self.EP_SIZE,
|
||||||
num_physical_experts=self.NUM_PHYSICAL,
|
num_physical_experts=self.NUM_PHYSICAL,
|
||||||
ep_rank=ep_rank,
|
ep_rank=ep_rank,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
)
|
)
|
||||||
|
|
||||||
# ------------------------------------------------------------------ shape & range
|
# ------------------------------------------------------------------ shape & range
|
||||||
|
|
||||||
@@ -164,26 +168,25 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase):
|
|||||||
num_physical_experts=self.NUM_PHYSICAL,
|
num_physical_experts=self.NUM_PHYSICAL,
|
||||||
replicas_per_logical=2,
|
replicas_per_logical=2,
|
||||||
)
|
)
|
||||||
result = compute_logical_to_rank_dispatch_physical_map(
|
with _published(self.EP_SIZE, self.NNODES):
|
||||||
server_args=self.server_args,
|
result = compute_logical_to_rank_dispatch_physical_map(
|
||||||
logical_to_all_physical_map=logical_to_all_physical,
|
logical_to_all_physical_map=logical_to_all_physical,
|
||||||
ep_size=self.EP_SIZE,
|
ep_size=self.EP_SIZE,
|
||||||
num_physical_experts=self.NUM_PHYSICAL,
|
num_physical_experts=self.NUM_PHYSICAL,
|
||||||
ep_rank=0,
|
ep_rank=0,
|
||||||
)
|
)
|
||||||
self.assertEqual(result.shape, (1, self.NUM_LOGICAL))
|
self.assertEqual(result.shape, (1, self.NUM_LOGICAL))
|
||||||
self.assertTrue(torch.all(result >= 0))
|
self.assertTrue(torch.all(result >= 0))
|
||||||
|
|
||||||
def test_single_node(self):
|
def test_single_node(self):
|
||||||
"""With nnodes=1, all GPUs are on the same node."""
|
"""With nnodes=1, all GPUs are on the same node."""
|
||||||
server_args = _make_server_args(ep_size=4, nnodes=1)
|
with _published(ep_size=4, nnodes=1):
|
||||||
result = compute_logical_to_rank_dispatch_physical_map(
|
result = compute_logical_to_rank_dispatch_physical_map(
|
||||||
server_args=server_args,
|
logical_to_all_physical_map=self.logical_to_all_physical.clone(),
|
||||||
logical_to_all_physical_map=self.logical_to_all_physical.clone(),
|
ep_size=self.EP_SIZE,
|
||||||
ep_size=self.EP_SIZE,
|
num_physical_experts=self.NUM_PHYSICAL,
|
||||||
num_physical_experts=self.NUM_PHYSICAL,
|
ep_rank=0,
|
||||||
ep_rank=0,
|
)
|
||||||
)
|
|
||||||
self.assertEqual(result.shape, (self.NUM_LAYERS, self.NUM_LOGICAL))
|
self.assertEqual(result.shape, (self.NUM_LAYERS, self.NUM_LOGICAL))
|
||||||
self.assertTrue(torch.all(result >= 0))
|
self.assertTrue(torch.all(result >= 0))
|
||||||
self.assertTrue(torch.all(result < self.NUM_PHYSICAL))
|
self.assertTrue(torch.all(result < self.NUM_PHYSICAL))
|
||||||
@@ -195,13 +198,13 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase):
|
|||||||
torch.arange(self.NUM_PHYSICAL, dtype=torch.int64).unsqueeze(0).unsqueeze(0)
|
torch.arange(self.NUM_PHYSICAL, dtype=torch.int64).unsqueeze(0).unsqueeze(0)
|
||||||
)
|
)
|
||||||
mapping = mapping.expand(self.NUM_LAYERS, 1, self.NUM_PHYSICAL).clone()
|
mapping = mapping.expand(self.NUM_LAYERS, 1, self.NUM_PHYSICAL).clone()
|
||||||
result = compute_logical_to_rank_dispatch_physical_map(
|
with _published(self.EP_SIZE, self.NNODES):
|
||||||
server_args=self.server_args,
|
result = compute_logical_to_rank_dispatch_physical_map(
|
||||||
logical_to_all_physical_map=mapping,
|
logical_to_all_physical_map=mapping,
|
||||||
ep_size=self.EP_SIZE,
|
ep_size=self.EP_SIZE,
|
||||||
num_physical_experts=self.NUM_PHYSICAL,
|
num_physical_experts=self.NUM_PHYSICAL,
|
||||||
ep_rank=0,
|
ep_rank=0,
|
||||||
)
|
)
|
||||||
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))
|
||||||
|
|
||||||
@@ -210,16 +213,13 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase):
|
|||||||
physical_to_logical = append_trivial_expert_slots(
|
physical_to_logical = append_trivial_expert_slots(
|
||||||
physical_to_logical, count=16, num_logical_experts=64
|
physical_to_logical, count=16, num_logical_experts=64
|
||||||
)
|
)
|
||||||
server_args = _make_server_args(ep_size=5, nnodes=1)
|
with _published(ep_size=5, nnodes=1, ep_join_mode="scale"):
|
||||||
server_args.ep_join_mode = "scale"
|
logical_to_physical = _compute_logical_to_all_physical_map(
|
||||||
|
physical_to_logical_map=physical_to_logical,
|
||||||
logical_to_physical = _compute_logical_to_all_physical_map(
|
num_logical_experts=64,
|
||||||
server_args=server_args,
|
ep_size=5,
|
||||||
physical_to_logical_map=physical_to_logical,
|
moe_ep_rank=4,
|
||||||
num_logical_experts=64,
|
)
|
||||||
ep_size=5,
|
|
||||||
moe_ep_rank=4,
|
|
||||||
)
|
|
||||||
|
|
||||||
self.assertEqual(logical_to_physical[0, :16, 0].tolist(), list(range(64, 80)))
|
self.assertEqual(logical_to_physical[0, :16, 0].tolist(), list(range(64, 80)))
|
||||||
|
|
||||||
|
|||||||
@@ -787,9 +787,7 @@ class TestShardConfig(unittest.TestCase):
|
|||||||
# moe_dense_tp_size / LM-head flags out of the cache key before.
|
# moe_dense_tp_size / LM-head flags out of the cache key before.
|
||||||
loader = object.__new__(PreshardedModelLoader)
|
loader = object.__new__(PreshardedModelLoader)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
moe_dense_tp_size=1,
|
|
||||||
moe_dp_size=2,
|
moe_dp_size=2,
|
||||||
enable_dp_lm_head=True,
|
|
||||||
enable_fp32_lm_head=True,
|
enable_fp32_lm_head=True,
|
||||||
ep_num_redundant_experts=4,
|
ep_num_redundant_experts=4,
|
||||||
enable_eplb=True,
|
enable_eplb=True,
|
||||||
@@ -812,7 +810,14 @@ class TestShardConfig(unittest.TestCase):
|
|||||||
"init_expert_location",
|
"init_expert_location",
|
||||||
"structural_signature",
|
"structural_signature",
|
||||||
}
|
}
|
||||||
parallel = SimpleNamespace(tp_size=8, moe_dp_size=2, moe_ep_size=4, pp_size=1)
|
parallel = SimpleNamespace(
|
||||||
|
tp_size=8,
|
||||||
|
moe_dp_size=2,
|
||||||
|
moe_ep_size=4,
|
||||||
|
pp_size=1,
|
||||||
|
moe_dense_tp_size=1,
|
||||||
|
enable_dp_lm_head=True,
|
||||||
|
)
|
||||||
with mock.patch(
|
with mock.patch(
|
||||||
"sglang.srt.model_loader.loader.get_server_args",
|
"sglang.srt.model_loader.loader.get_server_args",
|
||||||
return_value=server_args,
|
return_value=server_args,
|
||||||
|
|||||||
@@ -760,6 +760,33 @@ class TestForwardFlags(_IsolatedServerArgs):
|
|||||||
self.assertEqual(probe(torch.zeros(())).item(), 28)
|
self.assertEqual(probe(torch.zeros(())).item(), 28)
|
||||||
self.assertEqual(probe(torch.zeros(())).item(), 0)
|
self.assertEqual(probe(torch.zeros(())).item(), 0)
|
||||||
|
|
||||||
|
def test_parallel_config_leaves_trace_under_torch_compile(self):
|
||||||
|
# Regression: parallel config leaves resolve through
|
||||||
|
# ``ParallelContext.__getattr__`` (the bag fallback), and gate helpers
|
||||||
|
# such as ``enable_moe_dense_fully_dp()`` read them inside compiled
|
||||||
|
# model forwards — the fallback body must stay dynamo-traceable
|
||||||
|
# (``object.__getattribute__`` graph-breaks). fullgraph=True turns any
|
||||||
|
# graph break back into a failure.
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
|
reset_context()
|
||||||
|
with get_context().override_server_args(moe_dense_tp_size=1, dwdp_size=4):
|
||||||
|
|
||||||
|
@torch.compile(fullgraph=True, backend="eager", dynamic=False)
|
||||||
|
def probe(x):
|
||||||
|
par = get_parallel()
|
||||||
|
if par.enable_prefill_context_parallel:
|
||||||
|
x = x + 1
|
||||||
|
if par.moe_dense_tp_size == 1:
|
||||||
|
x = x + 2
|
||||||
|
if par.dwdp_size > 1:
|
||||||
|
x = x + 4
|
||||||
|
return x
|
||||||
|
|
||||||
|
self.assertEqual(probe(torch.zeros(())).item(), 6)
|
||||||
|
|
||||||
def test_graph_visible_flags_are_process_visible_across_threads(self):
|
def test_graph_visible_flags_are_process_visible_across_threads(self):
|
||||||
# Documented divergence from the contextvar-backed flags: plain slots
|
# Documented divergence from the contextvar-backed flags: plain slots
|
||||||
# are process-global (the storage form these flags had before the
|
# are process-global (the storage form these flags had before the
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ _EXCLUDED = (
|
|||||||
"multimodal_gen",
|
"multimodal_gen",
|
||||||
)
|
)
|
||||||
|
|
||||||
_BASELINE = 38
|
_BASELINE = 34
|
||||||
|
|
||||||
|
|
||||||
class TestServerArgsWriterRatchet(CustomTestCase):
|
class TestServerArgsWriterRatchet(CustomTestCase):
|
||||||
|
|||||||
Reference in New Issue
Block a user