diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index e2c410996..3829c4beb 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -43,7 +43,6 @@ from sglang.srt.runtime_context import ( get_device, get_exec, get_parallel, - get_server_args, ) from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip @@ -762,7 +761,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 f4a373ce5..165db4315 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -572,7 +572,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()): @@ -624,7 +624,7 @@ class CommonKVManager(BaseKVManager): "rank_port": self.rank_port, "page_size": self.kv_args.page_size, "kv_cache_dtype": get_model().kv_cache_dtype, - "load_balance_method": self.server_args.load_balance_method, + "load_balance_method": get_parallel().load_balance_method, "enable_dsa_cache_layer_split": getattr( self.server_args, "enable_dsa_cache_layer_split", False ), diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 2ba2323c2..8226bf887 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -55,7 +55,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 @@ -948,7 +948,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..88ae03ba7 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 @@ -16,6 +16,8 @@ import torch.distributed._symmetric_memory as symm_mem import triton import triton.language as tl +from sglang.srt.runtime_context import get_parallel + logger = logging.getLogger(__name__) # Each thread moves _NUMEL_PER_THREAD bf16 via one 128-bit multimem op; the @@ -466,7 +468,6 @@ 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 tp_group = get_tp_group() # Only probe node topology when the deployment can actually span @@ -477,7 +478,7 @@ class MultimemAllGatherer: # 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/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index d9975d981..528e3c4ef 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -93,6 +93,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa from sglang.srt.parser.template_detection import resolve_auto_parsers from sglang.srt.parser.template_manager import TemplateManager from sglang.srt.plugins import load_plugins +from sglang.srt.runtime_context import get_parallel from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import ( MultiprocessingSerializer, @@ -253,7 +254,7 @@ class Engine(EngineScoreMixin, EngineBase): # Initialize ZMQ sockets context = zmq.Context(2) - if self.server_args.node_rank == 0: + if server_args.node_rank == 0: self.send_to_rpc = get_zmq_socket( context, zmq.DEALER, self.port_args.rpc_ipc_name, True ) @@ -301,7 +302,7 @@ class Engine(EngineScoreMixin, EngineBase): routed_dp_rank = data_parallel_rank if routed_dp_rank is not None: - dp_size = self.server_args.dp_size + dp_size = get_parallel().dp_size if dp_size <= 1 and routed_dp_rank == 0: logger.debug( f"routed_dp_rank={routed_dp_rank} is ignored because dp_size={dp_size}" diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index d5b0b72b3..753005208 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -12,7 +12,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( is_in_tc_piecewise_cuda_graph, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip from sglang.srt.utils.common import ceil_align, ceil_div @@ -76,20 +76,20 @@ def should_use_dsa_fused_topk( def is_dsa_enable_prefill_cp(): - return get_server_args().enable_dsa_prefill_context_parallel + return get_parallel().enable_dsa_prefill_context_parallel 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 774ccdcf2..9faa78a28 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -76,7 +76,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_spec, ) from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -271,7 +270,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 @@ -282,7 +281,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" @@ -440,11 +439,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/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index a377a6054..df03ae977 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -47,7 +47,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, @@ -345,7 +345,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/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/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index 0f289f602..7f8e3e4f2 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -58,13 +58,13 @@ 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" ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 0dfdc8207..07e77b7aa 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1198,7 +1198,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, @@ -4313,7 +4313,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 8503c7b09..25b30c4f3 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -24,7 +24,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 @@ -377,7 +377,7 @@ class SchedulerDPAttnAdapter: def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch): return prepare_mlp_sync_batch_raw( local_batch, - 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, diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index 47c145adc..9dcfbaaed 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -16,7 +16,7 @@ from sglang.srt.managers.io_struct import ( sock_recv, ) 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 from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method @@ -127,7 +127,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: @@ -156,7 +156,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: @@ -233,7 +233,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 a9b9e1875..4524240ff 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/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index cdb9cfea5..47069bb76 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -74,7 +74,7 @@ from sglang.srt.managers.io_struct import ( UpdateWeightsFromTensorReqOutput, ) from sglang.srt.managers.load_snapshot import LoadSnapshot -from sglang.srt.runtime_context import get_lora +from sglang.srt.runtime_context import get_lora, get_parallel from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.utils import ( get_bool_env_var, @@ -146,8 +146,8 @@ class TokenizerControlMixin: def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int): primary_group_control = ( - self.server_args.enable_dp_attention - and not self.server_args.enable_dp_attention_local_control_broadcast + get_parallel().enable_dp_attention + and not get_parallel().enable_dp_attention_local_control_broadcast ) if primary_group_control: control_fan_out = ( @@ -397,7 +397,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" results = await self.init_weights_update_group_communicator(obj) @@ -410,7 +410,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for destroy parameter update group" results = await self.destroy_weights_update_group_communicator(obj) @@ -423,7 +423,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" if obj.abort_all_requests: @@ -454,7 +454,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() # TODO: support DP assert ( - self.server_args.dp_size == 1 + get_parallel().dp_size == 1 ), "dp_size must be 1 for init_weights_send_group_for_remote_instance" result = ( await self.init_weights_send_group_for_remote_instance_communicator(obj) @@ -469,7 +469,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() # TODO: support DP assert ( - self.server_args.dp_size == 1 + get_parallel().dp_size == 1 ), "dp_size must be 1 for send_weights_to_remote_instance" result = (await self.send_weights_to_remote_instance_communicator(obj))[0] return result.success, result.message @@ -481,7 +481,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from tensor" if obj.abort_all_requests: @@ -517,7 +517,7 @@ class TokenizerControlMixin: try: # For now, we only support single data parallel instance assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from IPC" logger.info("Starting IPC weight update") @@ -578,7 +578,7 @@ class TokenizerControlMixin: # TODO (lifuhuang): Remove this after we verify that dynamic lora loading works # with dp_size > 1. assert ( - self.server_args.dp_size == 1 + get_parallel().dp_size == 1 ), "dp_size must be 1 for dynamic lora loading" logger.info( "Start load Lora adapter. Lora name=%s, path=%s", @@ -654,7 +654,7 @@ class TokenizerControlMixin: ) assert ( - self.server_args.dp_size == 1 + get_parallel().dp_size == 1 ), "dp_size must be 1 for dynamic lora loading" logger.info( "Start load Lora adapter from tensors. Lora name=%s", @@ -730,7 +730,7 @@ class TokenizerControlMixin: # TODO (lifuhuang): Remove this after we verify that dynamic lora loading works # with dp_size > 1. assert ( - self.server_args.dp_size == 1 + get_parallel().dp_size == 1 ), "dp_size must be 1 for dynamic lora loading" logger.info( "Start unload Lora adapter. Lora name=%s", @@ -750,7 +750,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() results = await self.get_weights_by_name_communicator(obj) all_parameters = [r.parameter for r in results] - if self.server_args.dp_size == 1: + if get_parallel().dp_size == 1: return all_parameters[0] else: return all_parameters diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 17036d6ed..4d12e3fe5 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -116,6 +116,7 @@ from sglang.srt.runtime_context import ( get_lora, get_model, get_observability, + get_parallel, get_serving, ) from sglang.srt.sampling.sampling_params import SamplingParams @@ -1365,7 +1366,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): return batch_size > 0 and ( self.server_args.enable_tokenizer_batch_encode or ( - (not self.server_args.enable_dp_attention) + (not get_parallel().enable_dp_attention) and (not self._batch_has_text(batch_size, requests)) ) ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 414286793..91344dc45 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -150,6 +150,7 @@ from sglang.srt.runtime_context import ( get_global_dwdp_manager, get_lora, get_model, + get_parallel, get_schedule, set_global_dwdp_manager, ) @@ -392,12 +393,12 @@ class ModelRunner: is_scale_join = get_exec().moe.ep_join_mode == "scale" if is_scale_join: join_effective_ep_size = ( - self.server_args.ep_join_rank_offset + self.ps.tp_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() @@ -407,7 +408,7 @@ class ModelRunner: else: 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, @@ -441,9 +442,9 @@ 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 - ) + from sglang.srt.runtime_context import get_context + + 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 after elastic EP scale-up" @@ -617,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 ) @@ -797,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, @@ -1608,7 +1609,7 @@ 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) @@ -1627,7 +1628,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 _report_elastic_scale_failure(self, error: str, effective_size: int) -> None: if self.ps.tp_rank != 0 or self.server_args.is_ep_scale_joiner: @@ -1704,7 +1705,9 @@ 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) + from sglang.srt.runtime_context import get_context + + 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/models/apertus.py b/python/sglang/srt/models/apertus.py index b9f79000f..b2342e251 100644 --- a/python/sglang/srt/models/apertus.py +++ b/python/sglang/srt/models/apertus.py @@ -26,9 +26,7 @@ import torch from torch import nn from transformers import ApertusConfig -from sglang.srt.distributed import ( - get_pp_group, -) +from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import XIELU from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( @@ -52,7 +50,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 +440,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..c1d085159 100644 --- a/python/sglang/srt/models/arcee.py +++ b/python/sglang/srt/models/arcee.py @@ -20,9 +20,7 @@ import torch from torch import nn from transformers import LlamaConfig -from sglang.srt.distributed import ( - get_pp_group, -) +from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( @@ -46,7 +44,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 +403,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 cc58be6c0..3a8c8cef3 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -79,7 +79,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers @@ -821,7 +820,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 7044f0a9a..424a8d311 100644 --- a/python/sglang/srt/models/bailing_moe_linear.py +++ b/python/sglang/srt/models/bailing_moe_linear.py @@ -57,7 +57,6 @@ from sglang.srt.runtime_context import ( get_device, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import ( @@ -1085,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_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 0fa9172ff..4776ea19c 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -62,7 +62,6 @@ 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.utils import BumpAllocator, add_prefix, is_cuda, is_npu @@ -381,7 +380,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 10488e61f..a3395aed2 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -2719,7 +2719,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 733744257..81ba91cfb 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -134,7 +134,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, ) if not _is_hip: @@ -2404,7 +2403,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 888edb932..665d52833 100755 --- a/python/sglang/srt/models/exaone_moe.py +++ b/python/sglang/srt/models/exaone_moe.py @@ -60,7 +60,6 @@ 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.utils import LazyValue, add_prefix, is_cuda, make_layers @@ -643,7 +642,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..0bf3c2414 100644 --- a/python/sglang/srt/models/falcon_h1.py +++ b/python/sglang/srt/models/falcon_h1.py @@ -13,9 +13,7 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( ) from sglang.srt.layers.attention.mamba.mamba import MambaMixer2 from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -36,7 +34,6 @@ 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.utils import add_prefix, is_cuda, make_layers @@ -477,7 +474,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..9db2034e0 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -87,7 +87,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import ( @@ -1171,7 +1170,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 2b987ab42..0df3664b9 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -78,7 +78,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import ( @@ -908,7 +907,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 bd671c571..c53f54592 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -70,7 +70,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, ) from sglang.srt.utils import ( LazyValue, @@ -257,7 +256,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(): @@ -775,7 +774,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 3669a60a9..1c942fe61 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -49,7 +49,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, ) from sglang.srt.utils import LazyValue, add_prefix, make_layers @@ -637,7 +636,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 52dde19a5..e9eab5137 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -78,7 +78,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import ( @@ -828,7 +827,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..faa6bb37c 100644 --- a/python/sglang/srt/models/llama.py +++ b/python/sglang/srt/models/llama.py @@ -25,10 +25,7 @@ import torch from torch import nn from transformers import LlamaConfig -from sglang.srt.distributed import ( - get_pp_group, - get_pp_indices, -) +from sglang.srt.distributed import get_pp_group, get_pp_indices from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( @@ -53,7 +50,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 +527,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 c0662e1c0..0e991147f 100644 --- a/python/sglang/srt/models/longcat_flash.py +++ b/python/sglang/srt/models/longcat_flash.py @@ -41,17 +41,13 @@ from sglang.jit_kernel.dsv4 import linear_bf16_fp32 from sglang.kernels.ops.moe.ep_moe_kernels import zero_experts_compute_triton from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.configs import LongcatFlashConfig -from sglang.srt.distributed import ( - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -87,7 +83,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, @@ -714,7 +710,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 5962f1733..b20ce3a09 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 3a292b1ce..a3f4bb3c9 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -77,7 +77,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, ) from sglang.srt.utils import ( LazyValue, @@ -1187,7 +1186,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..9207938d6 100644 --- a/python/sglang/srt/models/mimo_v2_nextn.py +++ b/python/sglang/srt/models/mimo_v2_nextn.py @@ -26,9 +26,7 @@ from sglang.srt.layers.communicator import ( LayerScatterModes, enable_moe_dense_fully_dp, ) -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -44,7 +42,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 +257,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 9735e2fd7..e02f5d0ab 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -76,7 +76,7 @@ from sglang.srt.model_loader.weight_utils import ( maybe_remap_kv_scale_name, ) from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP -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, @@ -1438,7 +1438,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 811b9b4ec..3b0fe7b0b 100644 --- a/python/sglang/srt/models/minimax_m3_vl.py +++ b/python/sglang/srt/models/minimax_m3_vl.py @@ -105,7 +105,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 b76425ad6..476603ad3 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -90,7 +90,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, ) from sglang.srt.utils import ( add_prefix, @@ -939,7 +938,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..ef414c3f9 100644 --- a/python/sglang/srt/models/nemotron_h_mtp.py +++ b/python/sglang/srt/models/nemotron_h_mtp.py @@ -19,10 +19,7 @@ from torch import nn from sglang.srt.configs import NemotronHConfig from sglang.srt.distributed import get_pp_group -from sglang.srt.layers.dp_attention import ( - attn_tp_all_reduce, - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import attn_tp_all_reduce, is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ColumnParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor @@ -38,7 +35,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 +335,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 786624eda..060c839a7 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -93,7 +93,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, ) from sglang.srt.utils import ( add_prefix, @@ -1011,7 +1010,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 ff1209b23..7c05fda69 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -34,7 +34,6 @@ 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.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu @@ -495,7 +494,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_moe.py b/python/sglang/srt/models/qwen3_moe.py index 0c509fa1d..e9f132fcd 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -76,7 +76,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import ( @@ -956,7 +955,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 42e38db81..3d13c216a 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -15,9 +15,7 @@ from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_r from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers.attention.mamba.mamba import mamba_v2_sharded_weight_loader from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import GemmaRMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -50,7 +48,6 @@ 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.utils import ( @@ -831,9 +828,7 @@ class Qwen3HybridAttentionDecoderLayer(nn.Module): if self.attn_output_gate: if _is_hip: - from sglang.jit_kernel.triton.sigmoid_gate_mul import ( - sigmoid_gate_mul, - ) + from sglang.jit_kernel.triton.sigmoid_gate_mul import sigmoid_gate_mul attn_output = sigmoid_gate_mul(attn_output, gate) else: @@ -1030,7 +1025,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..ce3149661 100644 --- a/python/sglang/srt/models/qwen3_next_mtp.py +++ b/python/sglang/srt/models/qwen3_next_mtp.py @@ -35,7 +35,6 @@ 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.utils import add_prefix, is_npu @@ -85,7 +84,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 9dca3f92e..7454d7aad 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -68,7 +68,7 @@ from sglang.srt.models.utils import ( ) 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, is_cpu, is_npu, round_up from sglang.srt.utils.hf_transformers_utils import get_processor @@ -1269,7 +1269,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 cac464596..a668ca8a0 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -1226,7 +1226,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 ed9dcb852..7c5d22a56 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -43,7 +43,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import add_prefix, is_cuda, make_layers @@ -473,7 +472,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 7e715b789..436634f1f 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -56,7 +56,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers @@ -557,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 e3c0e720e..103d0bf78 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -45,7 +45,6 @@ from sglang.srt.runtime_context import ( get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers @@ -817,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/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index b4eca59e5..a8b29eb73 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -375,7 +375,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 ) @@ -443,7 +443,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 ): @@ -487,7 +487,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( @@ -549,7 +549,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 )