diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 772d9af0e..08a05b3ae 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -758,7 +758,7 @@ class TboForwardBatchPreparer: # TODO improve, e.g. unify w/ `init_raw` 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 ): sum_len = end_token_index - start_token_index diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 43b9bbbbb..8b29a73ce 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -174,7 +174,7 @@ class CommonKVManager(BaseKVManager): self.attn_dp_size = get_attention_dp_size() self.attn_dp_rank = get_attention_dp_rank() 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.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.local_ip = get_local_ip_auto() 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 = ( @@ -651,7 +651,7 @@ class CommonKVManager(BaseKVManager): `Connection refused`, and the leader's `prefill_port_table` ends 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 if not (dist.is_available() and dist.is_initialized()): @@ -703,10 +703,8 @@ class CommonKVManager(BaseKVManager): "rank_port": self.rank_port, "page_size": self.kv_args.page_size, "kv_cache_dtype": self.kv_cache_dtype_str, - "load_balance_method": self.server_args.load_balance_method, - "enable_dsa_cache_layer_split": getattr( - self.server_args, "enable_dsa_cache_layer_split", False - ), + "load_balance_method": get_parallel().load_balance_method, + "enable_dsa_cache_layer_split": get_parallel().enable_dsa_cache_layer_split, # Self-register the HTTP API port so the decode can derive the PD # retract rebootstrap /generate URL from bootstrap info instead of a # router-injected pd_rebootstrap_prefill_url. @@ -1078,12 +1076,11 @@ class CommonKVSender(BaseKVSender): return 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 self.kv_mgr.server_args.load_balance_method != "follow_bootstrap_room": + if get_parallel().dp_size > 1 and not req_has_disagg_prefill_dp_rank: + if get_parallel().load_balance_method != "follow_bootstrap_room": self._register_prefill_dp_rank() elif ( - self.kv_mgr.attn_dp_rank - != self.bootstrap_room % self.kv_mgr.server_args.dp_size + self.kv_mgr.attn_dp_rank != self.bootstrap_room % get_parallel().dp_size ): # follow_bootstrap_room was overridden by external routed_dp_rank 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"{self.kv_mgr.attn_dp_rank} but bootstrap_room " 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"to allow mixed routing.", ) @@ -1168,7 +1165,7 @@ class CommonKVSender(BaseKVSender): if ( 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( self.kv_mgr, diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 1e23c54ef..87cb6b45e 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -60,7 +60,7 @@ from sglang.srt.observability.trace import ( TraceReqContext, 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.utils.network import NetworkAddress @@ -1099,7 +1099,7 @@ class MooncakeKVManager(CommonKVManager): if ( self.attn_cp_size > 1 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 diff --git a/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py b/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py index 758e9c965..8d80f2ab8 100644 --- a/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py +++ b/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py @@ -466,18 +466,18 @@ class MultimemAllGatherer: # Lazy import avoids a module-load dependency on the distributed facade. from sglang.srt.distributed import get_tp_group 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() # Only probe node topology when the deployment can actually span # nodes. Check world_size first so a TP=1 gatherer short-circuits - # before reading server args (which may be unpublished on offline - # paths). On a single node every TP rank is co-located, so skip the + # before reading the parallel config (which may be unpublished on + # 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 # EP/mooncake setups, and keep multimem enabled. if ( 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)) ): logger.warning( diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 40e277af2..61037ded8 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -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.eplb.expert_location import broadcast_global_expert_location_metadata 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 if TYPE_CHECKING: @@ -308,13 +309,10 @@ def elastic_expanded_world_enabled() -> bool: Launch-time TP groups exclude ranks admitted during scale-up. """ - from sglang.srt.runtime_context import get_server_args - inst = ElasticEPStateManager.instance() if inst is None: return False - sa = get_server_args() - if sa.max_ep_size is None: + if get_parallel().max_ep_size is None: return False active_target_size = inst.effective_ep_size if inst.pending_ep_size is not None and inst.scale_phase in ( diff --git a/python/sglang/srt/eplb/expert_location.py b/python/sglang/srt/eplb/expert_location.py index 0d633f643..f2435e180 100644 --- a/python/sglang/srt/eplb/expert_location.py +++ b/python/sglang/srt/eplb/expert_location.py @@ -32,10 +32,13 @@ if TYPE_CHECKING: 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.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( @@ -156,7 +159,6 @@ class ExpertLocationMetadata: ) assert physical_to_logical_map.shape[-1] == common["num_physical_experts"] logical_to_all_physical_map = _compute_logical_to_all_physical_map( - server_args=server_args, physical_to_logical_map=physical_to_logical_map, num_logical_experts=model_config_for_expert_location.num_logical_experts, ep_size=common["ep_size"], @@ -164,7 +166,6 @@ class ExpertLocationMetadata: ) return ExpertLocationMetadata._init_raw( - server_args=server_args, ep_size=common["ep_size"], physical_to_logical_map=physical_to_logical_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.to(server_args.device) + from sglang.srt.runtime_context import get_parallel + common = ExpertLocationMetadata._init_common(server_args, model_config) if common is None: @@ -193,7 +196,7 @@ class ExpertLocationMetadata: model_config_for_expert_location = common["model_config_for_expert_location"] num_physical_experts = common["num_physical_experts"] 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 @@ -213,7 +216,6 @@ class ExpertLocationMetadata: ) return ExpertLocationMetadata._init_raw( - server_args=server_args, ep_size=common["ep_size"], physical_to_logical_map=physical_to_logical_map.to(server_args.device), logical_to_all_physical_map=logical_to_all_physical_map.to( @@ -223,6 +225,8 @@ class ExpertLocationMetadata: @staticmethod def _init_common(server_args: ServerArgs, model_config: ModelConfig): + from sglang.srt.runtime_context import get_exec, get_parallel + model_config_for_expert_location = ( ModelConfigForExpertLocation.from_model_config(model_config) ) @@ -232,16 +236,17 @@ class ExpertLocationMetadata: base_num_physical_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 - 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 server_args.ep_join_mode == "scale": + if get_exec().moe.ep_join_mode == "scale": ep_size = max( 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 = ( _compute_elastic_expert_layout( @@ -264,12 +269,13 @@ class ExpertLocationMetadata: @staticmethod def _init_raw( - server_args: ServerArgs, ep_size: int, physical_to_logical_map: torch.Tensor, logical_to_all_physical_map: torch.Tensor, moe_ep_rank: Optional[int] = None, ): + from sglang.srt.runtime_context import get_exec + _, num_physical_experts = physical_to_logical_map.shape logical_to_all_physical_map_padded = F.pad( @@ -291,7 +297,6 @@ class ExpertLocationMetadata: ep_size=ep_size, logical_to_rank_dispatch_physical_map=( compute_logical_to_rank_dispatch_physical_map( - server_args=server_args, logical_to_all_physical_map=logical_to_all_physical_map, ep_size=ep_size, num_physical_experts=num_physical_experts, @@ -301,7 +306,7 @@ class ExpertLocationMetadata: else torch.distributed.get_rank() % ep_size ), ) - if server_args.ep_dispatch_algorithm == "static" + if get_exec().moe.ep_dispatch_algorithm == "static" else None ), ) @@ -536,12 +541,13 @@ def broadcast_global_expert_location_metadata( def _compute_logical_to_all_physical_map( - server_args: ServerArgs, physical_to_logical_map: torch.Tensor, num_logical_experts: int, ep_size: 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 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 # collapse is per-rank, and the full candidate list is what lets the dispatch # 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 - prefer_same_node = _prefer_same_node_experts(server_args) + prefer_same_node = _prefer_same_node_experts() 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_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) def compute_logical_to_rank_dispatch_physical_map( - server_args: ServerArgs, logical_to_all_physical_map: torch.Tensor, ep_size: int, num_physical_experts: int, ep_rank: int, seed: int = 42, ): + from sglang.srt.runtime_context import get_parallel + r = random.Random(seed) device = logical_to_all_physical_map.device logical_to_all_physical_map = logical_to_all_physical_map.cpu() 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 = ( - 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_gpu_physical_experts * num_gpus_per_node diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index 674ad95f2..5caf46cc9 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -98,14 +98,14 @@ def is_dsa_enable_prefill_cp(): def is_dsa_prefill_cp_in_seq_split(): return ( 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(): return ( 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" ) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index dd81f9489..bb99bdf1f 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -73,13 +73,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( check_cuda_graph_backend, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_spec, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_spec from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( get_bool_env_var, @@ -275,7 +269,7 @@ class AttnTpContext: def init_context(self, q_lora_rank, is_dsa): self.is_dsa = is_dsa 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 q_lora_rank is not None and not is_dsa @@ -286,7 +280,7 @@ class AttnTpContext: and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) 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: logging.info( "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(): - return get_server_args().moe_dense_tp_size == 1 + return get_parallel().moe_dense_tp_size == 1 def enable_dwdp(): - return get_server_args().dwdp_size > 1 + return get_parallel().dwdp_size > 1 class LayerCommunicator: diff --git a/python/sglang/srt/layers/cp/cp_decode_attn_tp.py b/python/sglang/srt/layers/cp/cp_decode_attn_tp.py index b074e635c..0abfcf764 100644 --- a/python/sglang/srt/layers/cp/cp_decode_attn_tp.py +++ b/python/sglang/srt/layers/cp/cp_decode_attn_tp.py @@ -15,7 +15,7 @@ import torch 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.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__) @@ -51,7 +51,7 @@ class CpDecodeAttnTpContext: """Slices replicated attention weights across CP ranks during decode.""" 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: self.decode_tp_rank = get_parallel().attn_cp_rank diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 499840fe2..480ed81a5 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -51,7 +51,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, 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 ( is_cpu, is_npu, @@ -349,7 +349,7 @@ class LogitsProcessor(nn.Module): self.config = config self.vocab_size = config.vocab_size 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 if self.use_attn_tp_group: self.attn_tp_size = get_parallel().attn_tp_size diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 03d0ccc2d..006f0db17 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -260,12 +260,11 @@ class FusedMoE(torch.nn.Module): num_shared_slots = num_fused_shared_experts self._num_global_routed = num_experts - num_shared_slots - server_args = get_server_args() 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 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: storage_ep_size = self.moe_ep_size diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py index 02366f886..c4462e6c6 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py @@ -12,6 +12,7 @@ from sglang.srt.layers.moe.moe_runner.base import ( register_fused_func, ) 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 if TYPE_CHECKING: @@ -276,7 +277,9 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: else: # Standard allgather path: the MoE sees up to dp_size local forwards # 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 # inference_mode(False) ensures the wrapper's pre-allocated CUDA-graph # buffers are normal tensors. This call typically happens inside diff --git a/python/sglang/srt/layers/moe/token_dispatcher/nixl.py b/python/sglang/srt/layers/moe/token_dispatcher/nixl.py index 090eb05fa..93a72433e 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/nixl.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/nixl.py @@ -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.utils import DeepEPMode +from sglang.srt.runtime_context import get_parallel try: from nixl_ep import Buffer @@ -127,9 +128,7 @@ class NixlEPBuffer: offset = ElasticEPStateManager.get_ep_join_rank_offset() global_rank = rank + offset - from sglang.srt.runtime_context import get_server_args - - max_ep_size = get_server_args().max_ep_size or world_size + max_ep_size = get_parallel().max_ep_size or world_size nixl_max_ranks = max_ep_size num_rdma_bytes = 0 @@ -226,9 +225,8 @@ class _NixlEPDispatcherImplBase: elastic_state.active_ranks if elastic_state is not None else None ) self._active_world_size = dist.get_world_size(group) - from sglang.srt.runtime_context import get_server_args - _max_ep = get_server_args().max_ep_size or self._active_world_size + _max_ep = get_parallel().max_ep_size or self._active_world_size self._mask_buffer = ( torch.zeros(_max_ep, dtype=torch.int32, device="cuda") if self.active_ranks is not None diff --git a/python/sglang/srt/layers/moe/token_dispatcher/pplx.py b/python/sglang/srt/layers/moe/token_dispatcher/pplx.py index ecb0d8194..6337bb944 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/pplx.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/pplx.py @@ -22,7 +22,7 @@ from sglang.srt.layers.moe.utils import ( DispatcherOutputDtype, 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 # DeepSeek / DeepGEMM block quantization convention. @@ -155,7 +155,7 @@ class PplxAllToAllManager: # pplx forces ep_size == world_size # with pp_size == 1 (enforced in _ensure_nvshmem), so the EP group spans # 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: cls._all_to_all = AllToAll.internode( diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 824beff64..9ecaec671 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -506,7 +506,7 @@ def should_skip_post_experts_all_reduce(*, is_tp_path: bool) -> bool: """ if should_skip_mlp_all_reduce(): return True - if get_server_args().dwdp_size > 1: + if get_parallel().dwdp_size > 1: return True if should_use_dp_reduce_scatterv(): return True diff --git a/python/sglang/srt/layers/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index 0f289f602..553c51708 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -58,19 +58,19 @@ class ContextParallelMetadata: 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(): return ( 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: 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): diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index c4ee02cab..9d6046a16 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -36,6 +36,7 @@ from sglang.srt.runtime_context import ( get_mm, get_model, get_observability, + get_parallel, get_schedule, get_serving, get_spec, @@ -1270,7 +1271,7 @@ class Scheduler( gloo_group=self.attn_tp_cpu_group, tp_rank=self.ps.tp_rank, 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, bootstrap_port=get_disagg().disaggregation_bootstrap_port, max_total_num_tokens=self.max_total_num_tokens, @@ -4462,7 +4463,7 @@ class Scheduler( old_ep_size = ElasticEPStateManager.get_effective_ep_size() new_ep_size = recv_req.new_ep_size - max_ep_size = self.server_args.max_ep_size or old_ep_size + max_ep_size = get_parallel().max_ep_size or old_ep_size logger.debug( "[Elastic EP][scale] request received: new_ep_size=%d " diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index 5b2dfafeb..af3182611 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -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.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.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils.common import require_mlp_tp_gather @@ -402,7 +402,7 @@ class SchedulerDPAttnAdapter: return prepare_mlp_sync_batch_raw( local_batch, 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_cp_size=self.ps.attn_cp_size, tp_group=self.tp_group, @@ -411,7 +411,7 @@ class SchedulerDPAttnAdapter: require_mlp_tp_gather=require_mlp_tp_gather(self.server_args), disable_overlap_schedule=get_schedule().disable_overlap_schedule, offload_tags=self.offload_tags, - dwdp=self.server_args.dwdp_size > 1, + dwdp=get_parallel().dwdp_size > 1, ) def maybe_prepare_mlp_sync_batch( diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index bbda80ccd..e35bff05e 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -27,7 +27,7 @@ from sglang.srt.managers.mm_utils import ( has_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 ( broadcast_pyobj, point_to_point_pyobj, @@ -151,7 +151,7 @@ class SchedulerRequestReceiver: return recv_reqs 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: work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs) else: @@ -180,7 +180,7 @@ class SchedulerRequestReceiver: # instead of the full tp_group. This avoids an expensive # all-ranks gloo sync. _local_ctrl = ( - self.server_args.enable_dp_attention_local_control_broadcast + get_parallel().enable_dp_attention_local_control_broadcast or self.server_args.is_ep_scale_joiner ) if _local_ctrl: @@ -258,7 +258,7 @@ class SchedulerRequestReceiver: # peer ranks may still be unpickling ShmPointerMMData # (-> shm_open). Synchronize the same CPU groups that carried # 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: barrier(group=self.attn_tp_cpu_group) if self.ps.attn_cp_size > 1: diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 8bb69f489..715e20ebf 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -36,7 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import ( PPProxyTensors, ) 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.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj from sglang.srt.utils.common import get_device_module, is_xpu @@ -123,7 +123,7 @@ class SchedulerPPMixin: next_pp_outputs = None next_batch_result = 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 = ( self._pp_commit_send_output_work_and_preprocess_output_tensors( next_first_rank_mb_id, @@ -139,7 +139,7 @@ class SchedulerPPMixin: self.mb_metadata, 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 = ( self._pp_commit_send_output_work_and_preprocess_output_tensors( next_first_rank_mb_id, @@ -269,7 +269,7 @@ class SchedulerPPMixin: server_is_idle = False 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 = ( self._pp_commit_send_output_work_and_preprocess_output_tensors( next_first_rank_mb_id, @@ -285,7 +285,7 @@ class SchedulerPPMixin: self.mb_metadata, 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 = ( self._pp_commit_send_output_work_and_preprocess_output_tensors( next_first_rank_mb_id, @@ -428,7 +428,7 @@ class SchedulerPPMixin: pp_proxy_tensors = self._pp_recv_proxy_tensors() # 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 = ( self._pp_commit_send_output_work_and_preprocess_output_tensors( next_first_rank_mb_id, @@ -446,7 +446,7 @@ class SchedulerPPMixin: 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 = ( self._pp_commit_send_output_work_and_preprocess_output_tensors( next_first_rank_mb_id, @@ -557,10 +557,10 @@ class SchedulerPPMixin: self.on_idle() 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. 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.last_mbs = [None] * self.pp_loop_size diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index b55397550..2c77f8a2b 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -585,7 +585,7 @@ class CPUGraphRunner: model_runner.server_args.enable_profile_cuda_graph ) 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.capture_forward_mode = ForwardMode.DECODE diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 388845d25..880c81b6d 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -407,19 +407,17 @@ class ModelRunner: ): 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) if self.ps.tp_rank == 0: register_scale_cohort( - self.server_args.ep_join_rank_offset, + get_parallel().ep_join_rank_offset, join_effective_ep_size, ) join_scale_process_group() - self.server_args.override( - "elastic_ep.scale_join", ep_size=join_effective_ep_size - ) + get_context().override("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( model_config=self.model_config, moe_ep_rank=global_ep_rank, @@ -443,9 +441,7 @@ class ModelRunner: new_dp_size=join_effective_ep_size, new_dp_rank=global_ep_rank, ) - self.server_args.override( - "elastic_ep.scale_join", dp_size=join_effective_ep_size - ) + get_context().override("elastic_ep.scale_join", dp_size=join_effective_ep_size) if self.eplb_manager is not None: self.eplb_manager.disable_rebalance( "EPLB rebalance is disabled while elastic EP scale-up " @@ -622,7 +618,7 @@ class ModelRunner: if self.is_draft_worker: return 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 else 0 ) @@ -802,7 +798,7 @@ class ModelRunner: device=self.device, tp_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 ), host_to_device_ratio=hisparse_cfg.host_to_device_ratio, @@ -833,7 +829,7 @@ class ModelRunner: def post_capture_elastic_ep_recover(self): 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( model_config=self.model_config, moe_ep_rank=global_ep_rank, @@ -870,7 +866,7 @@ class ModelRunner: self.prefill_attention_backend_str = backends.prefill_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() def _prepare_replicated_q_proj(self) -> None: @@ -1107,7 +1103,7 @@ class ModelRunner: def maybe_init_dwdp(self): if self.is_draft_worker: return - if self.server_args.dwdp_size <= 1: + if get_parallel().dwdp_size <= 1: return from sglang.srt.layers.moe.dwdp import DwdpManager @@ -1658,9 +1654,9 @@ class ModelRunner: if added <= 0: 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 - 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( metadata.physical_to_logical_map, @@ -1677,7 +1673,7 @@ class ModelRunner: set_global_expert_location_metadata(new_metadata, allow_overwrite=True) def _elastic_global_rank(self) -> int: - return self.ps.tp_rank + self.server_args.ep_join_rank_offset + return self.ps.tp_rank + get_parallel().ep_join_rank_offset def _rearm_eplb_after_elastic_scale(self) -> None: if self.eplb_manager is None: @@ -1775,7 +1771,7 @@ class ModelRunner: new_dp_size=target_size, 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() self._elastic_scale_ready_barrier( diff --git a/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py b/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py index 1f4bb50ff..5fbcab266 100644 --- a/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py +++ b/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py @@ -11,7 +11,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( RemoteInstanceWeightLoaderBackend, 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.utils.network import NetworkAddress, get_local_ip_auto @@ -76,11 +76,11 @@ class RemoteInstanceWeightTransporter: """ 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). # Derive host from dist_init_addr (shared across all nodes). bootstrap_host = ( - NetworkAddress.parse(self.server_args.dist_init_addr).resolved().host + NetworkAddress.parse(get_parallel().dist_init_addr).resolved().host ) else: bootstrap_host = "127.0.0.1" diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 952731036..31dd6c5f3 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -197,7 +197,8 @@ class BaseRunner(ABC): self.device = model_runner.device self.device_module = torch.get_device_module(self.device) 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.enable_pdmux = model_runner.server_args.enable_pdmux self.enable_return_hidden_states = ( @@ -313,7 +314,7 @@ class BaseRunner(ABC): hidden_size=mr.model_config.hidden_size, vocab_size=mr.model_config.vocab_size, 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, is_encoder_decoder=mr.model_config.is_encoder_decoder, 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_ 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_: global_num_tokens_cpu = [num_tokens] else: diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 5250dc1dd..4c396c802 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -335,7 +335,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): self.moe_fusions = self.model_runner.moe_fusions 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_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 73f899add..c2f4e83dd 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -47,7 +47,12 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( get_remote_instance_transfer_engine_info_per_rank, 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 # Try to import accelerate (optional dependency) @@ -1747,9 +1752,9 @@ class PreshardedModelLoader(DefaultModelLoader): "dp": _safe(lambda: parallel.moe_dp_size), "ep": _safe(lambda: parallel.moe_ep_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, - "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, "quantization": model_config.quantization, "model_dtype": str(model_config.dtype), diff --git a/python/sglang/srt/models/apertus.py b/python/sglang/srt/models/apertus.py index b9f79000f..a016bfdde 100644 --- a/python/sglang/srt/models/apertus.py +++ b/python/sglang/srt/models/apertus.py @@ -52,7 +52,7 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, 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 logger = logging.getLogger(__name__) @@ -442,7 +442,7 @@ class ApertusForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) diff --git a/python/sglang/srt/models/arcee.py b/python/sglang/srt/models/arcee.py index 20d0ecc7c..79a6beff4 100644 --- a/python/sglang/srt/models/arcee.py +++ b/python/sglang/srt/models/arcee.py @@ -46,7 +46,7 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, 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 logger = logging.getLogger(__name__) @@ -405,7 +405,7 @@ class ArceeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 1f1297c57..3258d849d 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -77,13 +77,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers LoraConfig = None @@ -823,7 +817,7 @@ class BailingMoEForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/bailing_moe_linear.py b/python/sglang/srt/models/bailing_moe_linear.py index bc62cc009..66de4b208 100644 --- a/python/sglang/srt/models/bailing_moe_linear.py +++ b/python/sglang/srt/models/bailing_moe_linear.py @@ -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.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import ( - get_device, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_device, get_forward, get_parallel, get_stream from sglang.srt.utils import ( BumpAllocator, add_prefix, @@ -1090,7 +1084,7 @@ class BailingMoELinearForCausalLM(nn.Module): config.hidden_size, params_dtype=torch.float32, 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) diff --git a/python/sglang/srt/models/bailing_moe_nextn.py b/python/sglang/srt/models/bailing_moe_nextn.py index 5741f81c4..dab75ef02 100644 --- a/python/sglang/srt/models/bailing_moe_nextn.py +++ b/python/sglang/srt/models/bailing_moe_nextn.py @@ -42,7 +42,7 @@ from sglang.srt.models.bailing_moe_linear import ( BailingMoeV2_5ForCausalLM, ) 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 LoraConfig = None @@ -208,7 +208,7 @@ class BailingMoeForCausalLMNextN(nn.Module): config.hidden_size, quant_config=quant_config, 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) if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid": diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index d92b18c03..da8ee0793 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -359,7 +359,7 @@ class DeepseekMLAForwardMixin: # --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). 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 not self.use_deep_gemm_bmm and self.w_kc_qrep is not None @@ -1029,7 +1029,7 @@ class DeepseekMLAForwardMixin: self.num_local_heads * get_parallel().attn_dcp_size, 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"): # A2A exchange of head partials + LSE, then local Triton combine. # MLA decode LSE is base-2 (FlashInfer-MLA/FlashMLA) -> base_on_e=False. diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index be937318c..a722e6076 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -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_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import ( - get_model, - get_parallel, - get_server_args, - get_spec, -) +from sglang.srt.runtime_context import get_model, get_parallel, get_spec from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu @@ -388,7 +383,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 095dad38d..61a4be735 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -2938,7 +2938,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): config.hidden_size, quant_config=quant_config, 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: # ranks other than the last rank will have a placeholder layer diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index e64cf2c1a..52fb8c098 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -138,13 +138,7 @@ from sglang.srt.models.deepseek_v2 import ( _is_npu, _is_xpu, ) -from sglang.srt.runtime_context import ( - get_device, - get_exec, - get_forward, - get_parallel, - get_server_args, -) +from sglang.srt.runtime_context import get_device, get_exec, get_forward, get_parallel if not _is_hip: from sglang.srt.layers.utils.cp_utils import ( @@ -2501,7 +2495,7 @@ class DeepseekV4ForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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: self.lm_head = PPMissingLayer() diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index 1dd326c6e..f8d039524 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -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_context import get_attn_backend 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 logger = logging.getLogger(__name__) @@ -233,7 +233,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/exaone4.py b/python/sglang/srt/models/exaone4.py index 648f70417..d01dad06c 100644 --- a/python/sglang/srt/models/exaone4.py +++ b/python/sglang/srt/models/exaone4.py @@ -28,7 +28,7 @@ from sglang.srt.model_loader.weight_utils import ( default_weight_loader, 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.utils import get_exception_traceback, logger @@ -439,7 +439,7 @@ class Exaone4ForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/exaone_moe.py b/python/sglang/srt/models/exaone_moe.py index 88e707642..93e8de9fb 100755 --- a/python/sglang/srt/models/exaone_moe.py +++ b/python/sglang/srt/models/exaone_moe.py @@ -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.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import ( - get_exec, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_parallel, get_stream from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -648,7 +643,7 @@ class ExaoneMoEForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) # For EAGLE3 support diff --git a/python/sglang/srt/models/exaone_moe_mtp.py b/python/sglang/srt/models/exaone_moe_mtp.py index 439a4c354..10dea5461 100644 --- a/python/sglang/srt/models/exaone_moe_mtp.py +++ b/python/sglang/srt/models/exaone_moe_mtp.py @@ -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.model_executor.forward_batch_info import ForwardBatch 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 logger = logging.getLogger(__name__) @@ -63,7 +63,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/falcon_h1.py b/python/sglang/srt/models/falcon_h1.py index 5283f2798..8e3060cfb 100644 --- a/python/sglang/srt/models/falcon_h1.py +++ b/python/sglang/srt/models/falcon_h1.py @@ -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_context import get_attn_backend from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import ( - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_forward, get_parallel, get_stream from sglang.srt.utils import add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -477,7 +472,7 @@ class FalconH1ForCausalLM(nn.Module): quant_config=quant_config, org_num_embeddings=config.vocab_size, 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_multiplier = config.lm_head_multiplier diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 1e48115c3..b5bb69d2f 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -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_v2 import DeepseekV2ForCausalLM from sglang.srt.models.utils import WeightsMapper, apply_qk_norm -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -1171,7 +1165,7 @@ class Glm4MoeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index 30e2c7f6a..c5ed8243d 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -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_v2 import DeepseekV2AttentionMLA -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import ( BumpAllocator, LazyValue, @@ -911,7 +905,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/glm4_moe_lite_nextn.py b/python/sglang/srt/models/glm4_moe_lite_nextn.py index 9c91bbae2..d3dbb95be 100644 --- a/python/sglang/srt/models/glm4_moe_lite_nextn.py +++ b/python/sglang/srt/models/glm4_moe_lite_nextn.py @@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import ( Glm4MoeLiteDecoderLayer, 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 logger = logging.getLogger(__name__) @@ -151,7 +151,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/glm4_moe_nextn.py b/python/sglang/srt/models/glm4_moe_nextn.py index 5804bf241..c836ae19e 100644 --- a/python/sglang/srt/models/glm4_moe_nextn.py +++ b/python/sglang/srt/models/glm4_moe_nextn.py @@ -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.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 logger = logging.getLogger(__name__) @@ -137,7 +137,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py index 38fdc0a64..0b5e21f58 100644 --- a/python/sglang/srt/models/glm4v_moe.py +++ b/python/sglang/srt/models/glm4v_moe.py @@ -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.models.glm4_moe import Glm4MoeModel 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.hf_transformers_utils import get_processor @@ -69,7 +69,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): config.hidden_size, quant_config=quant_config, 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: # ranks other than the last rank will have a placeholder layer diff --git a/python/sglang/srt/models/glm_ocr_nextn.py b/python/sglang/srt/models/glm_ocr_nextn.py index a4d0566b4..cdd9ca18c 100644 --- a/python/sglang/srt/models/glm_ocr_nextn.py +++ b/python/sglang/srt/models/glm_ocr_nextn.py @@ -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.models.glm4 import Glm4DecoderLayer 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 logger = logging.getLogger(__name__) @@ -134,7 +134,7 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 7d99e74cf..4fa6c351e 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -259,7 +259,7 @@ class GptOssSparseMoeBlock(nn.Module): hidden_states: torch.Tensor, forward_batch: Optional[ForwardBatch] = None, ) -> torch.Tensor: - if get_server_args().dwdp_size > 1: + if get_parallel().dwdp_size > 1: return self.forward_dwdp(hidden_states) if not get_moe_a2a_backend().is_deepep(): @@ -787,7 +787,7 @@ class GptOssForCausalLM(nn.Module): config.hidden_size, # quant_config=quant_config, 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.capture_aux_hidden_states = False diff --git a/python/sglang/srt/models/laguna.py b/python/sglang/srt/models/laguna.py index 03377b4b4..3287d9776 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -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_loader.weight_utils import default_weight_loader from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel from sglang.srt.utils import LazyValue, add_prefix, make_layers logger = logging.getLogger(__name__) @@ -645,7 +640,7 @@ class LagunaForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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: self.lm_head = PPMissingLayer() diff --git a/python/sglang/srt/models/llada2.py b/python/sglang/srt/models/llada2.py index 8a184db53..626914dfe 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -76,13 +76,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import ( LazyValue, add_prefix, @@ -830,7 +824,7 @@ class LLaDA2MoeModelLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/llama.py b/python/sglang/srt/models/llama.py index 6771c6136..9209979e7 100644 --- a/python/sglang/srt/models/llama.py +++ b/python/sglang/srt/models/llama.py @@ -53,7 +53,7 @@ from sglang.srt.model_loader.weight_utils import ( maybe_remap_kv_scale_name, ) 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.utils import get_exception_traceback @@ -530,7 +530,7 @@ class LlamaForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) diff --git a/python/sglang/srt/models/longcat_flash.py b/python/sglang/srt/models/longcat_flash.py index 43860282c..4bed1df35 100644 --- a/python/sglang/srt/models/longcat_flash.py +++ b/python/sglang/srt/models/longcat_flash.py @@ -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.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 ( BumpAllocator, add_prefix, @@ -721,7 +721,7 @@ class LongcatFlashForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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.capture_aux_hidden_states = False diff --git a/python/sglang/srt/models/mellum.py b/python/sglang/srt/models/mellum.py index c792e37f8..625bef1cf 100644 --- a/python/sglang/srt/models/mellum.py +++ b/python/sglang/srt/models/mellum.py @@ -51,7 +51,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, 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 _is_cuda = is_cuda() @@ -520,7 +520,7 @@ class MellumForCausalLM(Qwen3MoeForCausalLM): cfg.hidden_size, quant_config=quant_config, 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.capture_aux_hidden_states = False diff --git a/python/sglang/srt/models/mimo_v2.py b/python/sglang/srt/models/mimo_v2.py index a16afd50f..315adf08a 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -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_vl import MiMoVisionTransformer, MiMoVLVisionConfig -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel from sglang.srt.utils import ( LazyValue, add_prefix, @@ -1192,7 +1187,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin): config.hidden_size, quant_config=quant_config, 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: self.lm_head = PPMissingLayer() diff --git a/python/sglang/srt/models/mimo_v2_nextn.py b/python/sglang/srt/models/mimo_v2_nextn.py index 49d32b66f..1eaa1f453 100644 --- a/python/sglang/srt/models/mimo_v2_nextn.py +++ b/python/sglang/srt/models/mimo_v2_nextn.py @@ -44,7 +44,7 @@ from sglang.srt.models.mimo_v2 import ( MiMoV2MLP, 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 MiMoV2Config = None @@ -259,7 +259,7 @@ class MiMoV2MTP(MiMoV2ForCausalLM): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/minimax_m3.py b/python/sglang/srt/models/minimax_m3.py index 4b9017b48..eb12e1b88 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -80,7 +80,7 @@ from sglang.srt.model_loader.weight_utils import ( ) from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP 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 ( add_prefix, get_device_sm, @@ -1453,7 +1453,7 @@ class MiniMaxM3SparseForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/minimax_m3_vl.py b/python/sglang/srt/models/minimax_m3_vl.py index bddc034be..120a0d6fc 100644 --- a/python/sglang/srt/models/minimax_m3_vl.py +++ b/python/sglang/srt/models/minimax_m3_vl.py @@ -120,7 +120,7 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module): text_config.hidden_size, quant_config=quant_config, 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: self.lm_head = PPMissingLayer() diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index bcb36a496..29332555f 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -89,12 +89,7 @@ from sglang.srt.models.nemotron_h_utils import ( pad_to_original_num_tokens, ) from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel from sglang.srt.utils import ( add_prefix, get_current_device_stream_fast, @@ -944,7 +939,7 @@ class NemotronHForCausalLM(nn.Module): else lora_config.lora_vocab_padding_size ), 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), ) else: diff --git a/python/sglang/srt/models/nemotron_h_mtp.py b/python/sglang/srt/models/nemotron_h_mtp.py index d38de9b9d..b3cd79fad 100644 --- a/python/sglang/srt/models/nemotron_h_mtp.py +++ b/python/sglang/srt/models/nemotron_h_mtp.py @@ -38,7 +38,7 @@ from sglang.srt.models.nemotron_h import ( NemotronHMoEDecoderLayer, ) 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 @@ -338,7 +338,7 @@ class NemotronHForCausalLMMTP(NemotronHForCausalLM): self.config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 942b38d30..9d17e74d1 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -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.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -1028,7 +1023,7 @@ class Qwen2MoeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) # For EAGLE3 support diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 2f3000d29..9fbd04670 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -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 Qwen2Model from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import ( - get_exec, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_parallel, get_stream from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu Qwen3Config = None @@ -497,7 +492,7 @@ class Qwen3ForCausalLM(nn.Module): config.vocab_size, config.hidden_size, 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), ) else: diff --git a/python/sglang/srt/models/qwen3_5_text.py b/python/sglang/srt/models/qwen3_5_text.py index b1e226549..7366f9f88 100644 --- a/python/sglang/srt/models/qwen3_5_text.py +++ b/python/sglang/srt/models/qwen3_5_text.py @@ -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.models import qwen3_5 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 logger = logging.getLogger(__name__) @@ -73,7 +73,7 @@ class Qwen3_5ForCausalLM(nn.Module): quant_config=quant_config, org_num_embeddings=config.vocab_size, 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: self.lm_head = PPMissingLayer() diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 2f20631e3..00e654521 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -72,13 +72,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import ( LazyValue, add_prefix, @@ -966,7 +960,7 @@ class Qwen3MoeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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.capture_aux_hidden_states = False diff --git a/python/sglang/srt/models/qwen3_moe_mtp.py b/python/sglang/srt/models/qwen3_moe_mtp.py index 6f6ec6091..e351fb4d7 100644 --- a/python/sglang/srt/models/qwen3_moe_mtp.py +++ b/python/sglang/srt/models/qwen3_moe_mtp.py @@ -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.model_executor.forward_batch_info import ForwardBatch 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 logger = logging.getLogger(__name__) @@ -63,7 +63,7 @@ class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index 5022107d7..3aca4ec7c 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -49,12 +49,7 @@ from sglang.srt.model_loader.weight_utils import ( sharded_weight_loader, ) from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock -from sglang.srt.runtime_context import ( - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_forward, get_parallel, get_stream from sglang.srt.utils import ( LazyValue, add_prefix, @@ -1032,7 +1027,7 @@ class Qwen3NextForCausalLM(nn.Module): quant_config=quant_config, org_num_embeddings=config.vocab_size, 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) # For EAGLE3 support diff --git a/python/sglang/srt/models/qwen3_next_mtp.py b/python/sglang/srt/models/qwen3_next_mtp.py index 3d86e5f94..dd36afb51 100644 --- a/python/sglang/srt/models/qwen3_next_mtp.py +++ b/python/sglang/srt/models/qwen3_next_mtp.py @@ -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.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel -from sglang.srt.runtime_context import ( - get_model, - get_parallel, - get_server_args, - get_spec, -) +from sglang.srt.runtime_context import get_model, get_parallel, get_spec from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -85,7 +80,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): config.hidden_size, quant_config=quant_config, 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) # Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index c66bcb547..83697f140 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -73,7 +73,7 @@ from sglang.srt.multimodal.mm_utils import ( run_dp_sharded_mrope_vision_model, ) 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 ( add_prefix, cpu_has_amx_support, @@ -1280,7 +1280,7 @@ class Qwen3VLForConditionalGeneration(nn.Module): self.config.vocab_size, self.config.hidden_size, 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), ) else: diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index 6906820e3..087cf2e5f 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -1231,7 +1231,7 @@ class SarvamMLAForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, 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) diff --git a/python/sglang/srt/models/sdar.py b/python/sglang/srt/models/sdar.py index 7ae5d8b5d..3d5912d97 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -41,13 +41,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -475,7 +469,7 @@ class SDARForCausalLM(nn.Module): config.vocab_size, config.hidden_size, 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), ) else: diff --git a/python/sglang/srt/models/sdar_moe.py b/python/sglang/srt/models/sdar_moe.py index 14c53099b..9ac15f3fe 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -57,13 +57,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -562,7 +556,7 @@ class SDARMoeForCausalLM(nn.Module): config.vocab_size, config.hidden_size, 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), ) else: diff --git a/python/sglang/srt/models/step3p5.py b/python/sglang/srt/models/step3p5.py index 287df9f62..c9220f2f3 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -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_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import ( - get_exec, - get_forward, - get_parallel, - get_server_args, - get_stream, -) +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel, get_stream from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers Step3p5Config = None @@ -822,7 +816,7 @@ class Step3p5ForCausalLM(nn.Module): config.vocab_size, config.hidden_size, 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), ) else: diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index e5f71f96a..1d0dbab38 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -125,12 +125,19 @@ class ParallelContext: def __getattr__(self, name): # Reached only for names that are neither a live @property nor a slot: - # serve parallel config leaves from the published bag. - try: - config = object.__getattribute__(self, "_config") - except AttributeError: - config = None - if config is not None and name in config: + # serve parallel config leaves from the published bag. The body must + # stay dynamo-traceable — config-leaf reads such as + # ``get_parallel().moe_dense_tp_size`` run inside compiled model + # forwards, and ``object.__getattribute__`` graph-breaks. + if name.startswith("_"): + # 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) detail = ( "not a published parallel config leaf" diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index 785bca205..a72ea6360 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -15,7 +15,7 @@ from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, 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.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 @@ -389,7 +389,7 @@ class DSparkWorkerV2(BaseSpecWorker): self, batch: ScheduleBatch, on_publish ) -> GenerationBatchResult: 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( batch, capture_hidden_mode=CaptureHiddenMode.FULL ) @@ -457,7 +457,7 @@ class DSparkWorkerV2(BaseSpecWorker): def _dp_verify_tier_num_tokens(self, batch: ScheduleBatch) -> Optional[int]: if not ( 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 self._verify_planner.is_compact_mode ): @@ -501,7 +501,7 @@ class DSparkWorkerV2(BaseSpecWorker): if batch.forward_mode.is_idle(): self._observers.note_idle_decode_step() - if self.server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: if self._draft_is_moe: self._proposer.run_idle_participation(batch) self._verify_executor.run_idle_participation( @@ -563,7 +563,7 @@ class DSparkWorkerV2(BaseSpecWorker): global_num_reqs = ( max(batch.global_num_tokens) 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 else None ) @@ -731,14 +731,14 @@ class DSparkWorkerV2(BaseSpecWorker): # Chain layout only: step index = commit_lens - 1. A tree (topk > 1) # layout would need the accept-index mapping the shared spec_utils # 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 last_correct_step_indices = commit_lens.to(torch.int64) - 1 mamba_steps_to_track = 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 = ( seq_lens_pre_verify // mamba_track_interval != seq_lens_post_verify // mamba_track_interval diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 6f51e67a5..23e8cdee0 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -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 ( 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_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.multi_layer_eagle_utils import ( @@ -150,7 +150,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): self.device = model_runner.device self.device_module = torch.get_device_module(self.device) 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.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding diff --git a/python/sglang/srt/state_capturer/routed_experts.py b/python/sglang/srt/state_capturer/routed_experts.py index 5dccd46c4..1207b13f9 100644 --- a/python/sglang/srt/state_capturer/routed_experts.py +++ b/python/sglang/srt/state_capturer/routed_experts.py @@ -69,15 +69,14 @@ class RoutedExpertsCapturer(BaseTopkCapturer): topk_size = model_config.hf_text_config.num_experts_per_tok 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. # _get_local_slice indexes into [attention_dp_rank * cuda_graph_batch, ...) # and otherwise overflows on dp_rank > 0 when max_running_requests > # chunked_prefill_size. # FIXME: spec decoding's num_verify_tokens is still not accounted for. max_batch_size = max( - get_schedule().chunked_prefill_size * server_args.dp_size, - max_running_requests * server_args.dp_size, + get_schedule().chunked_prefill_size * get_parallel().dp_size, + max_running_requests * get_parallel().dp_size, ) super().__init__( diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 63e110d25..1b01e9269 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -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. """ 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: - assert server_args.dp_size > 1, "dp_size must be greater than 1" - if server_args.elastic_ep_backend is not None: + # elastic-EP scale-up rewrites dp_size on the published config + if get_parallel().enable_dp_attention: + 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 ( elastic_expanded_world_enabled, ) @@ -3552,10 +3554,10 @@ def require_mlp_tp_gather(server_args: ServerArgs): if elastic_expanded_world_enabled(): return True 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 return True - elif not server_args.enable_dp_lm_head: + elif not get_parallel().enable_dp_lm_head: return True elif get_moe_a2a_backend().is_none(): return True @@ -3571,8 +3573,8 @@ def require_mlp_tp_gather(server_args: ServerArgs): return True else: return ( - server_args.moe_dense_tp_size - > server_args.tp_size // server_args.dp_size + get_parallel().moe_dense_tp_size + > server_args.tp_size // get_parallel().dp_size ) else: 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 # cuda graph runner pads num_tokens to attn_tp_size, which can cause # 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 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 server_args.enable_dp_attention: - return server_args.dp_size < server_args.tp_size + if ( + not get_moe_a2a_backend().is_none() + 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: return True else: @@ -3605,7 +3612,9 @@ def require_gathered_buffer(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: diff --git a/python/sglang/srt/weight_cache/ipc_loader.py b/python/sglang/srt/weight_cache/ipc_loader.py index 02e7ae772..a761df838 100644 --- a/python/sglang/srt/weight_cache/ipc_loader.py +++ b/python/sglang/srt/weight_cache/ipc_loader.py @@ -484,9 +484,7 @@ class IpcModelLoader(BaseModelLoader): ep_size = ps.moe_ep_size - from sglang.srt.runtime_context import get_server_args - - dp_size = get_server_args().dp_size + dp_size = get_parallel().dp_size quant_method, quant_config = self._resolve_engine_quant(model_config) diff --git a/test/registered/ops/test_aiter_allreduce_fusion_amd.py b/test/registered/ops/test_aiter_allreduce_fusion_amd.py index 5a0f1fb41..69ed9a194 100755 --- a/test/registered/ops/test_aiter_allreduce_fusion_amd.py +++ b/test/registered/ops/test_aiter_allreduce_fusion_amd.py @@ -389,7 +389,6 @@ class TestAiterAllreduceFusionGate(CustomTestCase): tp_size=8, ): """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) with ExitStack() as stack: @@ -417,10 +416,14 @@ class TestAiterAllreduceFusionGate(CustomTestCase): 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( - 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( diff --git a/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py b/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py index fa2c729fb..0eb0e4278 100644 --- a/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py +++ b/test/registered/unit/eplb/test_compute_logical_to_rank_dispatch_physical_map.py @@ -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") -import types import unittest import torch @@ -14,20 +13,26 @@ from sglang.srt.eplb.expert_location import ( append_trivial_expert_slots, compute_logical_to_rank_dispatch_physical_map, ) +from sglang.srt.runtime_context import get_context from sglang.test.test_utils import CustomTestCase -def _make_server_args(ep_size: int, nnodes: int, moe_a2a_backend: str = "deepep"): - """Minimal server_args stub for expert placement tests. +def _published( + 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 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, nnodes=nnodes, - ep_join_mode=None, moe_a2a_backend=moe_a2a_backend, + ep_join_mode=ep_join_mode, ) @@ -74,7 +79,6 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase): NUM_LAYERS = 2 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( num_layers=self.NUM_LAYERS, num_logical_experts=self.NUM_LOGICAL, @@ -83,14 +87,14 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase): ) def _call(self, ep_rank, seed=42): - return compute_logical_to_rank_dispatch_physical_map( - server_args=self.server_args, - logical_to_all_physical_map=self.logical_to_all_physical.clone(), - ep_size=self.EP_SIZE, - num_physical_experts=self.NUM_PHYSICAL, - ep_rank=ep_rank, - seed=seed, - ) + with _published(self.EP_SIZE, self.NNODES): + return compute_logical_to_rank_dispatch_physical_map( + logical_to_all_physical_map=self.logical_to_all_physical.clone(), + ep_size=self.EP_SIZE, + num_physical_experts=self.NUM_PHYSICAL, + ep_rank=ep_rank, + seed=seed, + ) # ------------------------------------------------------------------ shape & range @@ -164,26 +168,25 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase): num_physical_experts=self.NUM_PHYSICAL, replicas_per_logical=2, ) - result = compute_logical_to_rank_dispatch_physical_map( - server_args=self.server_args, - logical_to_all_physical_map=logical_to_all_physical, - ep_size=self.EP_SIZE, - num_physical_experts=self.NUM_PHYSICAL, - ep_rank=0, - ) + with _published(self.EP_SIZE, self.NNODES): + result = compute_logical_to_rank_dispatch_physical_map( + logical_to_all_physical_map=logical_to_all_physical, + ep_size=self.EP_SIZE, + num_physical_experts=self.NUM_PHYSICAL, + ep_rank=0, + ) self.assertEqual(result.shape, (1, self.NUM_LOGICAL)) self.assertTrue(torch.all(result >= 0)) def test_single_node(self): """With nnodes=1, all GPUs are on the same node.""" - server_args = _make_server_args(ep_size=4, nnodes=1) - result = compute_logical_to_rank_dispatch_physical_map( - server_args=server_args, - logical_to_all_physical_map=self.logical_to_all_physical.clone(), - ep_size=self.EP_SIZE, - num_physical_experts=self.NUM_PHYSICAL, - ep_rank=0, - ) + with _published(ep_size=4, nnodes=1): + result = compute_logical_to_rank_dispatch_physical_map( + logical_to_all_physical_map=self.logical_to_all_physical.clone(), + ep_size=self.EP_SIZE, + num_physical_experts=self.NUM_PHYSICAL, + ep_rank=0, + ) self.assertEqual(result.shape, (self.NUM_LAYERS, self.NUM_LOGICAL)) self.assertTrue(torch.all(result >= 0)) 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) ) mapping = mapping.expand(self.NUM_LAYERS, 1, self.NUM_PHYSICAL).clone() - result = compute_logical_to_rank_dispatch_physical_map( - server_args=self.server_args, - logical_to_all_physical_map=mapping, - ep_size=self.EP_SIZE, - num_physical_experts=self.NUM_PHYSICAL, - ep_rank=0, - ) + with _published(self.EP_SIZE, self.NNODES): + result = compute_logical_to_rank_dispatch_physical_map( + logical_to_all_physical_map=mapping, + ep_size=self.EP_SIZE, + num_physical_experts=self.NUM_PHYSICAL, + ep_rank=0, + ) self.assertEqual(result.shape, (self.NUM_LAYERS, 1)) self.assertTrue(torch.all(result >= 0)) @@ -210,16 +213,13 @@ class TestComputeLogicalToRankDispatchPhysicalMap(CustomTestCase): physical_to_logical = append_trivial_expert_slots( physical_to_logical, count=16, num_logical_experts=64 ) - server_args = _make_server_args(ep_size=5, nnodes=1) - server_args.ep_join_mode = "scale" - - logical_to_physical = _compute_logical_to_all_physical_map( - server_args=server_args, - physical_to_logical_map=physical_to_logical, - num_logical_experts=64, - ep_size=5, - moe_ep_rank=4, - ) + with _published(ep_size=5, nnodes=1, ep_join_mode="scale"): + logical_to_physical = _compute_logical_to_all_physical_map( + physical_to_logical_map=physical_to_logical, + num_logical_experts=64, + ep_size=5, + moe_ep_rank=4, + ) self.assertEqual(logical_to_physical[0, :16, 0].tolist(), list(range(64, 80))) diff --git a/test/registered/unit/model_loader/test_presharded_loader.py b/test/registered/unit/model_loader/test_presharded_loader.py index 8d0c241d8..8a1391723 100644 --- a/test/registered/unit/model_loader/test_presharded_loader.py +++ b/test/registered/unit/model_loader/test_presharded_loader.py @@ -787,9 +787,7 @@ class TestShardConfig(unittest.TestCase): # moe_dense_tp_size / LM-head flags out of the cache key before. loader = object.__new__(PreshardedModelLoader) server_args = SimpleNamespace( - moe_dense_tp_size=1, moe_dp_size=2, - enable_dp_lm_head=True, enable_fp32_lm_head=True, ep_num_redundant_experts=4, enable_eplb=True, @@ -812,7 +810,14 @@ class TestShardConfig(unittest.TestCase): "init_expert_location", "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( "sglang.srt.model_loader.loader.get_server_args", return_value=server_args, diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 669422c30..d3ac25300 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -760,6 +760,33 @@ class TestForwardFlags(_IsolatedServerArgs): self.assertEqual(probe(torch.zeros(())).item(), 28) 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): # Documented divergence from the contextvar-backed flags: plain slots # are process-global (the storage form these flags had before the diff --git a/test/registered/unit/test_server_args_writer_ratchet.py b/test/registered/unit/test_server_args_writer_ratchet.py index 426bca4f2..b021f6b5c 100644 --- a/test/registered/unit/test_server_args_writer_ratchet.py +++ b/test/registered/unit/test_server_args_writer_ratchet.py @@ -49,7 +49,7 @@ _EXCLUDED = ( "multimodal_gen", ) -_BASELINE = 38 +_BASELINE = 34 class TestServerArgsWriterRatchet(CustomTestCase):