diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 864dc52de..2795460dc 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -380,7 +380,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: ) # Deferred import to avoid a circular import at module-load - # time (dsa.utils imports get_global_server_args). + # time (dsa.utils imports the runtime-context accessors). from sglang.srt.layers.attention.dsa.utils import ( aiter_can_use_preshuffle_paged_mqa, ) diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index ee755d373..d443b8cfd 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -40,7 +40,6 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip @@ -184,7 +183,7 @@ def _update_device_and_sum_field_from_cpu_field( cpu_value if isinstance(cpu_value, torch.Tensor) else torch.tensor(cpu_value, dtype=old_device_value.dtype) - ).to(device=get_global_server_args().device, non_blocking=True) + ).to(device=get_server_args().device, non_blocking=True) setattr(batch, device_field, new_device_value) if sum_field is not None: @@ -336,7 +335,7 @@ def compute_split_indices_for_cuda_graph_replay( class TboCudaGraphRunnerPlugin: def __init__(self): self._tbo_children_num_token_non_padded = torch.zeros( - (2,), dtype=torch.int32, device=get_global_server_args().device + (2,), dtype=torch.int32, device=get_server_args().device ) def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int): @@ -760,7 +759,7 @@ class TboForwardBatchPreparer: # TODO improve, e.g. unify w/ `init_raw` if ( - get_global_server_args().moe_dense_tp_size == 1 + get_server_args().moe_dense_tp_size == 1 and batch.global_dp_buffer_len is not None ): sum_len = end_token_index - start_token_index @@ -835,7 +834,7 @@ class TboForwardBatchPreparer: value_a = min(tbo_split_token_index, num_token_non_padded) value_b = max(0, num_token_non_padded - tbo_split_token_index) return torch.tensor([value_a, value_b], dtype=torch.int32).to( - device=get_global_server_args().device, non_blocking=True + device=get_server_args().device, non_blocking=True ) @classmethod diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index ef1f8feff..45d2baf2b 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -1783,9 +1783,9 @@ class _SGLangPlugin(_FrameworkPlugin): return None try: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - args = get_global_server_args() + args = get_server_args() if args is None: return None diff --git a/python/sglang/srt/distributed/device_communicators/pymscclpp.py b/python/sglang/srt/distributed/device_communicators/pymscclpp.py index 45395876f..261e8d6cd 100644 --- a/python/sglang/srt/distributed/device_communicators/pymscclpp.py +++ b/python/sglang/srt/distributed/device_communicators/pymscclpp.py @@ -14,7 +14,7 @@ from sglang.srt.compilation.compile_phase import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( is_in_tc_piecewise_cuda_graph, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args logger = logging.getLogger(__name__) @@ -25,7 +25,7 @@ class PyMscclppCommunicator: def _is_symm_mem_enabled(self) -> bool: try: - return get_global_server_args().enable_symm_mem + return get_server_args().enable_symm_mem except ValueError: return False diff --git a/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py b/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py index efab33daf..aba729908 100644 --- a/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py +++ b/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py @@ -15,7 +15,7 @@ from torch.cuda.memory import ( from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.environ import envs -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils.common import torch_release after_2_8_0 = torch_release >= (2, 8) @@ -159,7 +159,7 @@ _register_func = None def is_symmetric_memory_enabled(): try: - return get_global_server_args().enable_symm_mem + return get_server_args().enable_symm_mem except ValueError: return False 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 874ae121b..758e9c965 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 @@ -423,9 +423,9 @@ def recommended_max_tokens(include_prefill: bool, floor: int = 0) -> int: NCCL. Covers the spec-decode batch plus, if ``include_prefill``, a prefill chunk. Returns ``floor`` if server args are unavailable.""" try: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - sa = get_global_server_args() + sa = get_server_args() def g(name: str) -> int: v = getattr(sa, name, 0) diff --git a/python/sglang/srt/distributed/utils.py b/python/sglang/srt/distributed/utils.py index 8d4d2c443..61f2bb1f4 100644 --- a/python/sglang/srt/distributed/utils.py +++ b/python/sglang/srt/distributed/utils.py @@ -19,19 +19,13 @@ from torch.distributed import TCPStore logger = logging.getLogger(__name__) -# Global TCPStore that is created during distributed initialization -# This is the single shared store that all components should use -_global_tcp_store: Optional[TCPStore] = None - def set_global_tcp_store(store: TCPStore) -> None: - """Set the global TCPStore instance. + """Install the shared TCPStore created during distributed initialization; + the handle lives on ``ctx.resources``.""" + from sglang.srt.runtime_context import get_resources - This should be called during distributed initialization to make - the store available to all components that need it. - """ - global _global_tcp_store - _global_tcp_store = store + get_resources().tcp_store = store logger.info("Global TCPStore has been set") @@ -45,15 +39,15 @@ def get_global_tcp_store() -> Optional[TCPStore]: Returns: The global TCPStore instance, or None if not initialized yet. """ - global _global_tcp_store + from sglang.srt.runtime_context import get_resources - if _global_tcp_store is None: + store = get_resources().tcp_store + if store is None: logger.warning( "Global TCPStore not found. Make sure init_distributed_environment " "was called with a tcp:// init method." ) - - return _global_tcp_store + return store def ensure_divisibility(numerator, denominator): diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 83483c515..9fa1bfff0 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -166,7 +166,7 @@ def init_tokenizer_manager( if getattr(server_args, attr) != "auto": continue if suggested is not None: - setattr(server_args, attr, suggested) + server_args.override(source="template-detection", **{attr: suggested}) logger.info( f"Auto-detected --{attr.replace('_', '-')} as '{suggested}' from chat template" ) @@ -175,7 +175,7 @@ def init_tokenizer_manager( f"--{attr.replace('_', '-')}=auto specified but could not detect " f"{label} from chat template. Disabling {label}." ) - setattr(server_args, attr, None) + server_args.override(source="template-detection", **{attr: None}) return tokenizer_manager, template_manager diff --git a/python/sglang/srt/eplb/expert_location_dispatch.py b/python/sglang/srt/eplb/expert_location_dispatch.py index 9deb69b60..6484089bc 100644 --- a/python/sglang/srt/eplb/expert_location_dispatch.py +++ b/python/sglang/srt/eplb/expert_location_dispatch.py @@ -18,7 +18,7 @@ from typing import Literal, Optional import torch from sglang.srt.eplb.expert_location import get_global_expert_location_metadata -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import is_hip _is_hip = is_hip() @@ -37,7 +37,7 @@ class ExpertLocationDispatchInfo: @classmethod def init_new(cls, layer_id: int): - ep_dispatch_algorithm = get_global_server_args().ep_dispatch_algorithm + ep_dispatch_algorithm = get_server_args().ep_dispatch_algorithm expert_location_metadata = get_global_expert_location_metadata() assert expert_location_metadata is not None diff --git a/python/sglang/srt/eplb/expert_location_updater.py b/python/sglang/srt/eplb/expert_location_updater.py index 52d80a15d..694e62487 100644 --- a/python/sglang/srt/eplb/expert_location_updater.py +++ b/python/sglang/srt/eplb/expert_location_updater.py @@ -25,7 +25,7 @@ from sglang.srt.eplb.expert_location import ( ExpertLocationMetadata, get_global_expert_location_metadata, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import get_bool_env_var, get_int_env_var, is_hip logger = logging.getLogger(__name__) @@ -108,7 +108,7 @@ def _update_expert_weights_with_canary( canary_tensor = ( _get_canary_value(old_expert_location_metadata, layer_id) .clone() - .to(device=get_global_server_args().device, non_blocking=True) + .to(device=get_server_args().device, non_blocking=True) ) routed_experts_weights_of_layer[layer_id].append(canary_tensor) diff --git a/python/sglang/srt/hardware_backend/mlx/model_runner.py b/python/sglang/srt/hardware_backend/mlx/model_runner.py index d76671c94..86674adee 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner.py @@ -49,7 +49,7 @@ from sglang.srt.hardware_backend.mlx.kv_cache import ( uses_sliding_window_attention, ) from sglang.srt.mem_cache.memory_pool import ReqToTokenPool -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args logger = logging.getLogger(__name__) @@ -281,7 +281,7 @@ class MlxModelRunner: ): return None - chunk_size = get_global_server_args().mamba_cache_chunk_size + chunk_size = get_server_args().mamba_cache_chunk_size track_len = prefix_len + (new_token_count // chunk_size) * chunk_size branching_len = getattr(req, "mamba_branching_seqlen", None) if ( diff --git a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py index ab62b0145..6b823eaec 100644 --- a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py +++ b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py @@ -23,7 +23,7 @@ from sglang.srt.layers.utils.cp_utils import ( cp_allgather_and_save_kv_cache, ) from sglang.srt.mem_cache.memory_pool import KVWriteLoc -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -515,7 +515,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): and not forward_batch.forward_mode.is_draft_extend_v2() ): if forward_batch.attn_attend_prefix_cache: - assert not get_global_server_args().disable_chunked_prefix_cache + assert not get_server_args().disable_chunked_prefix_cache assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None assert forward_batch.prefix_chunk_max_seq_lens is not None diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py index 86b30924d..9c2b9c2c9 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py @@ -1362,9 +1362,9 @@ class DeepseekV4AscendAttnBackend( or forward_batch.forward_mode.is_draft_extend_v2() ): B = forward_batch.batch_size - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - n_draft = get_global_server_args().speculative_num_draft_tokens or 1 + n_draft = get_server_args().speculative_num_draft_tokens or 1 actual_q = torch.arange( n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device ) @@ -1409,9 +1409,9 @@ class DeepseekV4AscendAttnBackend( forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_draft_extend_v2() ): - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - max_seqlen_q = get_global_server_args().speculative_num_draft_tokens or 1 + max_seqlen_q = get_server_args().speculative_num_draft_tokens or 1 else: max_seqlen_q = 1 return self._kernel_metadata_from_parts( diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py b/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py index 87cf281a9..b3a909b22 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py @@ -27,7 +27,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( ) from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args class ViTNpuGraphRunner(ViTCudaGraphRunner): @@ -70,7 +70,7 @@ class ViTNpuGraphRunner(ViTCudaGraphRunner): graph = torch_npu.npu.NPUGraph() vit = self.vit - override_backend = get_global_server_args().mm_attention_backend + override_backend = get_server_args().mm_attention_backend with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool): y = None deepstack_outs: List[torch.Tensor] = [] diff --git a/python/sglang/srt/layers/activation.py b/python/sglang/srt/layers/activation.py index 3ea6ceeae..a7db9d27a 100644 --- a/python/sglang/srt/layers/activation.py +++ b/python/sglang/srt/layers/activation.py @@ -33,8 +33,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( Phase, check_cuda_graph_backend, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -90,7 +89,7 @@ logger = logging.getLogger(__name__) class SiluAndMul(MultiPlatformOp): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: self._forward_method = self.forward_native elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get(): self._forward_method = self.forward_aiter diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index c367779cb..97ed54de6 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -112,7 +112,7 @@ from sglang.srt.model_executor.forward_context import ( get_token_to_kv_pool, ) from sglang.srt.model_executor.runner import get_is_capture_mode -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args _use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get() if TYPE_CHECKING: @@ -134,7 +134,7 @@ def _is_in_piecewise_or_breakable_cuda_graph() -> bool: def _uses_dsa_attention_backend(forward_batch: ForwardBatch) -> bool: attn_backend = get_attn_backend() - server_args = get_global_server_args() + server_args = get_server_args() prefill_backend, decode_backend = server_args.get_attention_backends() prefill_backend = ( getattr(attn_backend, "prefill_attention_backend_str", None) or prefill_backend @@ -394,7 +394,7 @@ class Indexer(MultiPlatformOp): if _is_cuda: self.sm_count = deep_gemm.get_num_sms() self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) - pp_size = get_global_server_args().pp_size + pp_size = get_server_args().pp_size self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank else: self.logits_with_pp_recv = False @@ -446,7 +446,7 @@ class Indexer(MultiPlatformOp): base=rope_theta, # type: ignore rope_scaling=rope_scaling, is_neox_style=is_neox_style, - device=get_global_server_args().device, + device=get_server_args().device, ) self.block_size = block_size self.scale_fmt = scale_fmt @@ -1032,7 +1032,7 @@ class Indexer(MultiPlatformOp): total_mem = torch.cuda.get_device_properties(device_index).total_memory total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION) - mem_fraction_static = get_global_server_args().mem_fraction_static + mem_fraction_static = get_server_args().mem_fraction_static if mem_fraction_static is None: static_budget = total_mem_budget else: diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index 380cad4ed..21a684541 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -15,8 +15,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 -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip from sglang.srt.utils.common import ceil_align, ceil_div @@ -69,20 +68,20 @@ def compute_dsa_seqlens(original_seq_lens, dsa_index_topk: int): def is_dsa_enable_prefill_cp(): - return get_global_server_args().enable_dsa_prefill_context_parallel + return get_server_args().enable_dsa_prefill_context_parallel def is_dsa_prefill_cp_in_seq_split(): return ( is_dsa_enable_prefill_cp() - and get_global_server_args().dsa_prefill_cp_mode == "in-seq-split" + and get_server_args().dsa_prefill_cp_mode == "in-seq-split" ) def is_dsa_prefill_cp_round_robin_split(): return ( is_dsa_enable_prefill_cp() - and get_global_server_args().dsa_prefill_cp_mode == "round-robin-split" + and get_server_args().dsa_prefill_cp_mode == "round-robin-split" ) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index cd1801062..8db3b0fc8 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -861,9 +861,9 @@ class C4Indexer(nn.Module): self.rotary_emb = rotary_emb self.freqs_cis = freqs_cis self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5 - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - self.use_fp4_indexer = get_global_server_args().enable_deepseek_v4_fp4_indexer + self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer self.alt_streams = alt_streams def compute_q( diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 0f41886fc..f9dc630fb 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -28,7 +28,7 @@ from sglang.srt.layers.utils.cp_utils import ( from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm from sglang.srt.utils import get_compiler_backend @@ -1342,7 +1342,7 @@ class FlashAttentionBackend(AttentionBackend): ): # Do multi-head attention with chunked prefix cache if forward_batch.attn_attend_prefix_cache: - assert not get_global_server_args().disable_chunked_prefix_cache + assert not get_server_args().disable_chunked_prefix_cache # MHA for chunked prefix kv cache when running model with MLA assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index 3c172abc3..fb664a815 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -34,8 +34,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo 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_buffer -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_buffer, get_server_args from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_utils import ( draft_kv_indices_buffer_width, @@ -226,9 +225,9 @@ class FlashInferMLAAttnBackend(AttentionBackend): self.token_to_kv_pool = model_runner.token_to_kv_pool self.enable_chunk_kv = ( not skip_prefill - and get_global_server_args().disaggregation_mode != "decode" - and not get_global_server_args().disable_chunked_prefix_cache - and not get_global_server_args().flashinfer_mla_disable_ragged + and get_server_args().disaggregation_mode != "decode" + and not get_server_args().disable_chunked_prefix_cache + and not get_server_args().flashinfer_mla_disable_ragged ) self.page_size = model_runner.page_size @@ -404,7 +403,7 @@ class FlashInferMLAAttnBackend(AttentionBackend): prefix_lens = forward_batch.extend_prefix_lens extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) use_ragged = ( - not get_global_server_args().flashinfer_mla_disable_ragged + not get_server_args().flashinfer_mla_disable_ragged and extend_no_prefix # Piecewise cuda graph should use paged prefill to be compatible with prefix cache and not is_in_tc_piecewise_cuda_graph() diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 101b50287..f6ae4179d 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -19,7 +19,7 @@ from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.spec_info import SpecInput @@ -246,7 +246,7 @@ class MambaAttnBackendBase(AttentionBackend): lens_to_track = ( forward_batch.mamba_track_seqlens - forward_batch.extend_prefix_lens ) - mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size + mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size aligned_len = (lens_to_track // mamba_cache_chunk_size) * mamba_cache_chunk_size start_indices = query_start_loc[:-1] + aligned_len - conv_state_len start_indices = start_indices[forward_batch.mamba_track_mask] @@ -265,7 +265,7 @@ class MambaAttnBackendBase(AttentionBackend): """src/dst indices to track SSM states for prefix caching: aligned seqs cache last_recurrent_state, unaligned cache intermediate `h` at the last chunk boundary.""" - mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size + mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size # CPU to avoid kernel launches for the masking ops mamba_track_mask = forward_batch.mamba_track_mask.cpu() extend_seq_lens = forward_batch.extend_seq_lens.cpu() @@ -336,7 +336,7 @@ class MambaAttnBackendBase(AttentionBackend): """Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu % mamba_track_interval == 0, so force-flush and snapshot fire on the same steps (no off-by-one).""" - interval = get_global_server_args().mamba_track_interval + interval = get_server_args().mamba_track_interval if seq_lens_cpu is None: # Should not happen for the supported config; stay safe and never flush. return torch.zeros((bs,), dtype=torch.bool) @@ -748,8 +748,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase): # Page-major stores state strided; only the stride-aware Triton causal-conv # reads it (CUDA causal_conv1d garbles it). A model may also force Triton. use_triton_causal_conv = ( - use_triton_causal_conv - or get_global_server_args().enable_page_major_kv_layout + use_triton_causal_conv or get_server_args().enable_page_major_kv_layout ) layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) mixer_out, intermediate_states = mixer.forward( diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 1d628f10e..28177c251 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -38,8 +38,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo 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_buffer, get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_buffer, get_parallel, get_server_args from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2 if is_flashinfer_available(): @@ -199,7 +198,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None self.disable_chunked_prefix_cache = ( - get_global_server_args().disable_chunked_prefix_cache + get_server_args().disable_chunked_prefix_cache ) self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 12a5199d9..87405bb25 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -77,7 +77,7 @@ from sglang.srt.layers.linear import ( ) from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import add_prefix, get_bool_env_var _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip @@ -964,10 +964,7 @@ class VisionAttention(nn.Module): # Select attention backend via a unified method _passed_backend = qkv_backend qkv_backend = self._determine_attention_backend(_passed_backend) - if ( - get_global_server_args().mm_attention_backend is None - and _passed_backend is None - ): + if get_server_args().mm_attention_backend is None and _passed_backend is None: print_info_once(f"Multimodal attention backend not set. Use {qkv_backend}.") print_info_once(f"Using {qkv_backend} as multimodal attention backend.") @@ -1047,7 +1044,7 @@ class VisionAttention(nn.Module): weight_dtype=torch.float32, cast_x_before_out_mul=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) q_norm = RMSNorm( @@ -1075,7 +1072,7 @@ class VisionAttention(nn.Module): - CUDA (other): "triton_attn" - Non-CUDA: "sdpa" """ - override_backend = get_global_server_args().mm_attention_backend + override_backend = get_server_args().mm_attention_backend if override_backend is not None: backend = override_backend elif passed_backend is not None: @@ -1179,7 +1176,7 @@ class VisionAttention(nn.Module): x = x.unsqueeze(0) assert x.dim() == 3, x.shape if ( - get_global_server_args().rl_on_policy_target is not None + get_server_args().rl_on_policy_target is not None and position_embeddings is not None ): assert isinstance(position_embeddings, tuple), ( diff --git a/python/sglang/srt/layers/attention/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index 984e4c6bf..8d956c6ef 100644 --- a/python/sglang/srt/layers/attention/xpu_backend.py +++ b/python/sglang/srt/layers/attention/xpu_backend.py @@ -12,10 +12,10 @@ from sglang.srt.layers.attention.flashattention_backend import ( merge_state_v2_wrapper, prepare_swa_spec_page_table_triton, ) -from sglang.srt.managers.schedule_batch import get_global_server_args from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.runtime_context import get_server_args if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -638,7 +638,7 @@ class XPUAttentionBackend(AttentionBackend): ): # Do multi-head attention with chunked prefix cache if forward_batch.attn_attend_prefix_cache: - assert not get_global_server_args().disable_chunked_prefix_cache + assert not get_server_args().disable_chunked_prefix_cache # MHA for chunked prefix kv cache when running model with MLA assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index a85e522bd..ea84ae9a6 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -72,8 +72,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_forward, get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( get_bool_env_var, @@ -171,7 +170,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int): and batch_size > 0 and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE and not is_dp_attention_enabled() - and get_global_server_args().flashinfer_allreduce_fusion_backend is not None + and get_server_args().flashinfer_allreduce_fusion_backend is not None and not is_flashinfer_allreduce_unavailable() ) @@ -187,7 +186,7 @@ def apply_aiter_all_reduce_fusion(input_tensor: torch.Tensor): and total_bytes <= 8 * 1024 * 8192 and get_parallel().tp_size != 6 and not is_dp_attention_enabled() - and get_global_server_args().enable_aiter_allreduce_fusion + and get_server_args().enable_aiter_allreduce_fusion ) @@ -266,7 +265,7 @@ class AttnTpContext: def init_context(self, q_lora_rank, is_dsa): self.is_dsa = is_dsa self.allow_input_scattered = ( - get_global_server_args().enable_attn_tp_input_scattered + get_server_args().enable_attn_tp_input_scattered and (_is_cuda or _is_npu) and q_lora_rank is not None and not is_dsa @@ -275,9 +274,9 @@ class AttnTpContext: and get_moe_a2a_backend().is_none() and not enable_moe_dense_fully_dp() and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - and get_global_server_args().speculative_algorithm != "EAGLE3" + and get_server_args().speculative_algorithm != "EAGLE3" ) - if get_global_server_args().enable_attn_tp_input_scattered: + if get_server_args().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" @@ -407,7 +406,7 @@ class LayerScatterModes: not context.is_layer_sparse and context.is_next_layer_sparse and enable_moe_dense_fully_dp() - and get_global_server_args().enable_two_batch_overlap + and get_server_args().enable_two_batch_overlap ) @classmethod @@ -434,7 +433,7 @@ class LayerScatterModes: def enable_moe_dense_fully_dp(): - return get_global_server_args().moe_dense_tp_size == 1 + return get_server_args().moe_dense_tp_size == 1 class LayerCommunicator: @@ -463,7 +462,7 @@ class LayerCommunicator: ) self._post_init_communicate() self._speculative_algo = SpeculativeAlgorithm.from_string( - get_global_server_args().speculative_algorithm + get_server_args().speculative_algorithm ) def _post_init_communicate(self): @@ -811,7 +810,7 @@ class LayerCommunicator: and get_parallel().tp_size != 6 and not is_dp_attention_enabled() and get_moe_a2a_backend().is_none() - and get_global_server_args().enable_aiter_allreduce_fusion + and get_server_args().enable_aiter_allreduce_fusion ) ) and (not self.is_last_layer) @@ -1116,7 +1115,7 @@ class CommunicateWithAllReduceAndLayerNormFn: if not handled: quantize_communications = ( not forward_batch.forward_mode.is_decode_or_idle() - and get_global_server_args().enable_quant_communications + and get_server_args().enable_quant_communications ) if quantize_communications: hidden_states = attention_tensor_model_parallel_quant_all_reduce( diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index 4be8498a0..59825c9ce 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -195,9 +195,9 @@ class ContextParallelStrategy(ABC): def _is_dsa_active() -> bool: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - sa = get_global_server_args() + sa = get_server_args() return bool( getattr(sa, "enable_prefill_cp", False) and getattr(sa, "_is_dsa_model_arch", False) @@ -247,10 +247,10 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]: global _STRATEGY if _STRATEGY is None: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args try: - server_args = get_global_server_args() + server_args = get_server_args() except ValueError: return None if server_args is not None and getattr(server_args, "enable_prefill_cp", False): diff --git a/python/sglang/srt/layers/cp/zigzag.py b/python/sglang/srt/layers/cp/zigzag.py index d50ae25b6..b53236994 100644 --- a/python/sglang/srt/layers/cp/zigzag.py +++ b/python/sglang/srt/layers/cp/zigzag.py @@ -204,10 +204,10 @@ class ZigzagCPStrategy(ContextParallelStrategy): actual_seq_q_prev_list.append(block_sizes[cp_rank]) actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1]) - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args try: - device = torch.device(get_global_server_args().device) + device = torch.device(get_server_args().device) except Exception: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") cu_prev = [0] + list(accumulate(actual_seq_q_prev_list)) diff --git a/python/sglang/srt/layers/dcp/__init__.py b/python/sglang/srt/layers/dcp/__init__.py index fb87c1de6..60e94c526 100644 --- a/python/sglang/srt/layers/dcp/__init__.py +++ b/python/sglang/srt/layers/dcp/__init__.py @@ -46,7 +46,7 @@ from sglang.srt.layers.dcp.layout import ( from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata # NOTE: planner.py is intentionally NOT imported here. It depends on server_args -# (get_global_server_args), whereas this package-init executes at module-load time +# (get_server_args), whereas this package-init executes at module-load time # for every eager importer of the DCP primitives — triton_backend, # mem_cache.memory_pool, mem_cache.triton_ops.mla_buffer, mem_cache.kv_cache_builder, # the FlashInfer-MLA / FlashMLA backends, and the deepseek forward methods. Keeping diff --git a/python/sglang/srt/layers/dcp/planner.py b/python/sglang/srt/layers/dcp/planner.py index d8c0c08bb..de9f19fba 100644 --- a/python/sglang/srt/layers/dcp/planner.py +++ b/python/sglang/srt/layers/dcp/planner.py @@ -27,8 +27,7 @@ from sglang.srt.layers.dcp.kernels import ( ) from sglang.srt.layers.dcp.layout import update_local_kv_lens_for_dcp from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args def prepare_decode_context_parallel_metadata( @@ -54,12 +53,12 @@ def prepare_decode_context_parallel_metadata( extend_prefix_starts = torch.zeros( len(seq_lens), dtype=torch.int32, - device=get_global_server_args().device, + device=get_server_args().device, ) extend_cu_prefix_lens = torch.zeros( len(seq_lens) + 1, dtype=torch.int32, - device=get_global_server_args().device, + device=get_server_args().device, ) extend_cu_prefix_lens[1:] = torch.cumsum(extend_prefix_lens, dim=0) extend_cu_prefix_lens = extend_cu_prefix_lens[:-1] @@ -68,7 +67,7 @@ def prepare_decode_context_parallel_metadata( dcp_prefix_kv_indices = torch.empty( sum(extend_prefix_lens_cpu), dtype=torch.int32, - device=get_global_server_args().device, + device=get_server_args().device, ) create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)]( req_to_token, @@ -82,20 +81,20 @@ def prepare_decode_context_parallel_metadata( dcp_kv_indptr = torch.zeros( len(seq_lens) + 1, dtype=torch.int32, - device=get_global_server_args().device, + device=get_server_args().device, ) dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0) dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)] dcp_kv_indices = torch.zeros( seq_lens_sum, dtype=torch.int32, - device=get_global_server_args().device, + device=get_server_args().device, ) extend_cu_lens = torch.zeros( len(seq_lens) + 1, dtype=torch.int32, - device=get_global_server_args().device, + device=get_server_args().device, ) extend_cu_lens[1:] = torch.cumsum(extend_seq_lens, dim=0) extend_cu_lens = extend_cu_lens[:-1] diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 292d78907..8432a62a9 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -778,9 +778,9 @@ def get_moe_cp_size() -> int: def is_enable_moe_cp_allgather() -> bool: """True when moe_dp_size < attn_cp_size, requiring allgather across CP ranks before MoE.""" - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - sa = get_global_server_args() + sa = get_server_args() return sa.attn_cp_size > sa.moe_dp_size diff --git a/python/sglang/srt/layers/flashinfer_comm_fusion.py b/python/sglang/srt/layers/flashinfer_comm_fusion.py index 06abf4ca1..34c2b1373 100644 --- a/python/sglang/srt/layers/flashinfer_comm_fusion.py +++ b/python/sglang/srt/layers/flashinfer_comm_fusion.py @@ -13,8 +13,7 @@ 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_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import ( ceil_align, get_cuda_driver_bindings, @@ -673,7 +672,7 @@ def ensure_workspace_initialized( token_num = token_num or max_token_num group_key = (device_group, cpu_group) effective_dtype = dtype or torch.bfloat16 - server_args = get_global_server_args() + server_args = get_server_args() backend = resolve_flashinfer_allreduce_fusion_backend(server_args) if backend is None: return False diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 23c46113a..fd1bf385f 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -31,8 +31,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( Phase, check_cuda_graph_backend, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -193,7 +192,7 @@ def _forward_with_allreduce_fusion( return fused_result # For AITER route, preserve correctness when fused path is unavailable. - if _use_aiter and get_global_server_args().enable_aiter_allreduce_fusion: + if _use_aiter and get_server_args().enable_aiter_allreduce_fusion: x = tensor_model_parallel_all_reduce(x) return norm_module.forward(x, residual, None) @@ -271,7 +270,7 @@ class RMSNorm(MultiPlatformOp): if ( residual is not None or self.cast_x_before_out_mul - or get_global_server_args().rl_on_policy_target == "fsdp" + or get_server_args().rl_on_policy_target == "fsdp" ): return self.forward_native(x, residual, post_residual_addition) return rms_norm_batch_invariant( @@ -371,7 +370,7 @@ class RMSNorm(MultiPlatformOp): if ( residual is not None or self.cast_x_before_out_mul - or get_global_server_args().rl_on_policy_target == "fsdp" + or get_server_args().rl_on_policy_target == "fsdp" or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0) ): return self.forward_native(x, residual, post_residual_addition) @@ -432,7 +431,7 @@ class RMSNorm(MultiPlatformOp): if ( residual is not None or self.cast_x_before_out_mul - or get_global_server_args().rl_on_policy_target == "fsdp" + or get_server_args().rl_on_policy_target == "fsdp" ): return self.forward_native(x, residual, post_residual_addition) return rms_norm_batch_invariant( @@ -559,10 +558,7 @@ class RMSNorm(MultiPlatformOp): if self.variance_size_override is not None: return self.forward_native(x, residual, post_residual_addition) if is_batch_invariant_mode_enabled(): - if ( - residual is not None - or get_global_server_args().rl_on_policy_target == "fsdp" - ): + if residual is not None or get_server_args().rl_on_policy_target == "fsdp": return self.forward_native(x, residual, post_residual_addition) return rms_norm_batch_invariant( x, diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index a521e379d..310cacb7e 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -37,8 +37,7 @@ from sglang.srt.layers.parameter import ( _ColumnvLLMParameter, ) from sglang.srt.layers.utils import pad_or_narrow_weight -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs if TYPE_CHECKING: @@ -1545,7 +1544,7 @@ class RowParallelLinear(LinearBase): quantize_communications = ( ( not forward_batch.forward_mode.is_decode_or_idle() - and get_global_server_args().enable_quant_communications + and get_server_args().enable_quant_communications ) if forward_batch is not None else False diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 2ade16095..a413b1836 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -48,7 +48,6 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, ) from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils.common import ( is_cpu, is_npu, @@ -336,7 +335,7 @@ class LogitsProcessor(nn.Module): 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_fp32_lm_head = get_global_server_args().enable_fp32_lm_head + self.use_fp32_lm_head = get_server_args().enable_fp32_lm_head if self.use_attn_tp_group: self.attn_tp_size = get_parallel().attn_tp_size self.do_tensor_parallel_all_gather = ( @@ -360,7 +359,7 @@ class LogitsProcessor(nn.Module): self.final_logit_softcapping = None self.return_full_logits = return_full_logits - self.enable_mis = get_global_server_args().enable_mis + self.enable_mis = get_server_args().enable_mis self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer( max_tokens=triton_symm_mem_ag.recommended_max_tokens( @@ -970,7 +969,7 @@ class LogitsProcessor(nn.Module): None, # bias True, # is_vnni ) - elif get_global_server_args().rl_on_policy_target is not None: + elif get_server_args().rl_on_policy_target is not None: # Due to tie-weight, we may not be able to change lm_head's weight dtype logits = torch.matmul( hidden_states.bfloat16(), lm_head.weight.T.bfloat16() diff --git a/python/sglang/srt/layers/mhc.py b/python/sglang/srt/layers/mhc.py index 4c56521d1..efa3a980f 100644 --- a/python/sglang/srt/layers/mhc.py +++ b/python/sglang/srt/layers/mhc.py @@ -521,10 +521,10 @@ def prewarm_mhc_pre( the TileLang/DeepGEMM on-disk JIT cache, so this cost is paid only on a cold cache; later server runs hit the cache. Driven once per process from load_weights. """ - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args hc_mult, hidden_size = residual.shape[-2], residual.shape[-1] - max_num_tokens = get_global_server_args().chunked_prefill_size + max_num_tokens = get_server_args().chunked_prefill_size buckets = get_mhc_pre_token_count_representatives( max_num_tokens, hc_mult * hidden_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 3d4cbfc44..289907dd3 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -66,8 +66,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo is_in_tc_piecewise_cuda_graph, ) from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -276,7 +275,7 @@ class FusedMoE(torch.nn.Module): ) self.quant_method: Optional[FusedMoEMethodBase] = None - server_args = get_global_server_args() + server_args = get_server_args() kt_config = create_kt_config_from_server_args(server_args, layer_id) if kt_config is not None: if quant_config is not None: diff --git a/python/sglang/srt/layers/moe/hash_topk.py b/python/sglang/srt/layers/moe/hash_topk.py index 2be451e2d..49299deab 100644 --- a/python/sglang/srt/layers/moe/hash_topk.py +++ b/python/sglang/srt/layers/moe/hash_topk.py @@ -44,11 +44,10 @@ class HashTopK(nn.Module): ): super().__init__() self.layer_id = layer_id - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args self.enable_deepep_waterfill = ( - num_fused_shared_experts > 0 - and get_global_server_args().enable_deepep_waterfill + num_fused_shared_experts > 0 and get_server_args().enable_deepep_waterfill ) self.deepep_waterfill_balancer = None 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 02c781fa4..04d4b139d 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py @@ -244,14 +244,14 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: "Install with: pip install flashinfer" ) from e - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args assert layer.intermediate_size_per_partition > 0, ( f"CuteDSL MoE: intermediate_size_per_partition must be > 0, " f"got {layer.intermediate_size_per_partition}. Check EP/TP configuration." ) - server_args = get_global_server_args() + server_args = get_server_args() # CuteDSL wrapper preallocates CG buffers used by any captured graph # that routes through this MoE — decode and prefill alike. use_cuda_graph = not cuda_graph_fully_disabled() diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py index 95fc3ec64..30d1f1deb 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py @@ -17,7 +17,7 @@ from sglang.srt.batch_invariant_ops import is_batch_invariant_mode_enabled from sglang.srt.environ import envs from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig from sglang.srt.layers.moe.utils import get_moe_padding_size -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -488,7 +488,7 @@ def _fused_moe_kernel_sequence( out_hidden_states = torch.empty_like(hidden_states) use_fused_moe_sum_all_reduce = ( - get_global_server_args().enable_fused_moe_sum_all_reduce + get_server_args().enable_fused_moe_sum_all_reduce and (not no_combine) and (topk > 2) and (not use_int8_w8a16) diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py index 5f402913e..2f1735cf4 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py @@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional, Tuple import torch import triton -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import get_device_name, is_hip logger = logging.getLogger(__name__) @@ -69,7 +69,7 @@ def get_moe_configs( kernel on a given batch size bs, the closest batch size in the grid should be picked and the associated configuration chosen to invoke the kernel. """ - if get_global_server_args().enable_deterministic_inference: + if get_server_args().enable_deterministic_inference: logger.warning( "Deterministic inference is enabled, using default MoE kernel config." ) @@ -170,7 +170,7 @@ def get_default_config( is_marlin: bool, block_shape: Optional[List[int]] = None, ) -> Dict[str, int]: - if get_global_server_args().enable_deterministic_inference: + if get_server_args().enable_deterministic_inference: config = { "BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 64, diff --git a/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py b/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py index 029e382bc..4b4f9d015 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py @@ -23,7 +23,7 @@ from sglang.srt.layers.moe.token_dispatcher.flashinfer_utils import ( ) from sglang.srt.layers.moe.topk import StandardTopKOutput, TopKOutput from sglang.srt.layers.moe.utils import get_moe_runner_backend -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import get_int_env_var @@ -119,7 +119,7 @@ class FlashinferDispatcher(BaseDispatcher): # max_running_requests is not yet resolved at model-construction time, # so we use 4096 as a floor to cover decode batches and _dummy_run # (which warms up at batch_size = req_to_token_pool.size). - cps = get_global_server_args().chunked_prefill_size + cps = get_server_args().chunked_prefill_size default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096) self.max_num_tokens = get_int_env_var( "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", @@ -128,7 +128,7 @@ class FlashinferDispatcher(BaseDispatcher): # Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized. speculative_algo = SpeculativeAlgorithm.from_string( - get_global_server_args().speculative_algorithm + get_server_args().speculative_algorithm ) if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle(): total_dispatch_payload_size_per_token = ( diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index 27933521b..7a3c98b5a 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -395,11 +395,10 @@ class TopK(MultiPlatformOp): assert num_expert_group is not None and topk_group is not None self.layer_id = layer_id - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args self.enable_deepep_waterfill = ( - num_fused_shared_experts > 0 - and get_global_server_args().enable_deepep_waterfill + num_fused_shared_experts > 0 and get_server_args().enable_deepep_waterfill ) self.deepep_waterfill_balancer = None @@ -475,9 +474,9 @@ class TopK(MultiPlatformOp): # ===== TO BE REFACTORED ==== elif get_moe_runner_backend().is_experimental_sgl_trtllm(): try: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - use_standard_for_lora = bool(get_global_server_args().enable_lora) + use_standard_for_lora = bool(get_server_args().enable_lora) except ValueError: use_standard_for_lora = False output_format = ( @@ -1256,10 +1255,10 @@ def _eplb_remap_enabled() -> bool: # initial expert placement is non-trivial, or there are redundant physical # experts. Otherwise the map is identity and the remap must be skipped (it is # both unnecessary and not well-defined over the padded region of topk_ids). - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args try: - server_args = get_global_server_args() + server_args = get_server_args() except ValueError: # Global server args are not initialized outside the server runtime # (e.g. in unit tests that call select_experts directly). In that case diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 94f86ce85..fe8eb3545 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -20,7 +20,7 @@ _is_npu = is_npu() if TYPE_CHECKING: from sglang.srt.server_args import ServerArgs -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args logger = logging.getLogger(__name__) @@ -206,7 +206,7 @@ def get_deepep_output_dtype(self) -> DeepEPOutputDtype: """ # 0. Parse server argument. - server_args = get_global_server_args() + server_args = get_server_args() if server_args and server_args.deepep_dispatcher_output_dtype != "auto": return DeepEPOutputDtype(server_args.deepep_dispatcher_output_dtype) diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 72d4942fa..b96c90b97 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -33,7 +33,7 @@ from sglang.srt.layers.quantization.fp8_kernel import ( w8a8_block_fp8_matmul_deepgemm, w8a8_block_fp8_matmul_triton, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( ceil_align, ceil_div, @@ -1691,8 +1691,7 @@ def apply_fp8_linear( if ( input_scale is not None and input_scale.numel() == 1 - and get_global_server_args().cuda_graph_config.prefill.tc_compiler - == "inductor" + and get_server_args().cuda_graph_config.prefill.tc_compiler == "inductor" ): qinput = ( (input_2d * input_scale.reciprocal()) diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py index 186e32e93..bca97962b 100644 --- a/python/sglang/srt/layers/quantization/mxfp4.py +++ b/python/sglang/srt/layers/quantization/mxfp4.py @@ -48,7 +48,7 @@ from sglang.srt.layers.quantization.base_config import ( QuantizeMethodBase, ) from sglang.srt.layers.quantization.utils import is_layer_skipped -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( cpu_has_amx_support, is_cpu, @@ -334,7 +334,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase): self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4() self.use_marlin = get_moe_runner_backend().is_marlin() self.flashinfer_mxfp4_moe_precision = ( - get_global_server_args().flashinfer_mxfp4_moe_precision + get_server_args().flashinfer_mxfp4_moe_precision ) # When `flashinfer_mxfp4` is enabled, dispatch to one of two FlashInfer # entry points depending on the GPU: diff --git a/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py b/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py index 7a2e34d47..c318b5f8c 100644 --- a/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py +++ b/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py @@ -15,7 +15,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( ) from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.utils import RoutingMethodType -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( is_flashinfer_available, log_info_on_rank0, @@ -128,7 +128,7 @@ class Mxfp4FlashinferTrtllmMoEMethod: self._fp8 = fp8_method self.prefix = prefix self.flashinfer_mxfp4_moe_precision = ( - get_global_server_args().flashinfer_mxfp4_moe_precision + get_server_args().flashinfer_mxfp4_moe_precision ) def create_moe_runner(self, layer, moe_runner_config): diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index e04cd876f..95e5b2e7e 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -11,7 +11,7 @@ from sglang.srt.environ import envs from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.platforms import current_platform -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -127,7 +127,7 @@ class RotaryEmbedding(MultiPlatformOp): self._apply_rotary_emb_wrapped = apply_rotary_emb # XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend - if get_global_server_args().rl_on_policy_target is not None or _is_musa: + if get_server_args().rl_on_policy_target is not None or _is_musa: self._forward_method = self.forward_native self._apply_rotary_emb_wrapped = torch.compile( dynamic=True, @@ -151,7 +151,7 @@ class RotaryEmbedding(MultiPlatformOp): # create the cache on GPU for faster initialization. This may cause # a slight numerical difference between the HF implementation and ours. init_device = ( - "cpu" if get_global_server_args().rl_on_policy_target is not None else None + "cpu" if get_server_args().rl_on_policy_target is not None else None ) inv_freq = 1.0 / ( base @@ -162,7 +162,7 @@ class RotaryEmbedding(MultiPlatformOp): / self.rotary_dim ) ) - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: inv_freq = inv_freq.cuda() return inv_freq diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index bc9f29c5f..609488e7d 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import ( yarn_get_mscale_simple, yarn_linear_ramp_mask, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( cpu_has_amx_support, is_cuda, @@ -216,7 +216,7 @@ class MRotaryEmbedding(RotaryEmbedding): self.register_buffer("axis_map", axis_map, persistent=False) else: self.axis_map = None - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: self._forward_method = self.forward_native def get_cos_sin_with_position(self, positions): diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index 990652f04..1f7d5b8b1 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -15,7 +15,6 @@ from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logp from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils.async_probe import sanitize_nan_logits from sglang.srt.utils.common import ( get_bool_env_var, @@ -72,11 +71,9 @@ class Sampler(nn.Module): if is_dp_attention_enabled(): self.tp_sync_group = get_parallel().attn_tp_group.device_group - self.rl_on_policy_target = get_global_server_args().rl_on_policy_target + self.rl_on_policy_target = get_server_args().rl_on_policy_target # In RL on-policy mode, deterministic inference is automatically enabled. - self.enable_deterministic = ( - get_global_server_args().enable_deterministic_inference - ) + self.enable_deterministic = get_server_args().enable_deterministic_inference # In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. self.use_log_softmax_logprob = self.rl_on_policy_target is not None self.use_ascend_backend = get_server_args().sampling_backend == "ascend" @@ -460,7 +457,7 @@ def register_sampler_backend(backend: str, factory: Callable[[], "Sampler"]) -> def create_sampler(backend: Optional[str] = None) -> "Sampler": """Create a sampler honoring custom backend registrations.""" - server_args = get_global_server_args() + server_args = get_server_args() backend = backend or (server_args.sampling_backend if server_args else None) if backend in _CUSTOM_SAMPLER_FACTORIES: diff --git a/python/sglang/srt/layers/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index 559ec89d0..237665d57 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -15,8 +15,7 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.model_executor.forward_context import get_token_to_kv_pool -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args @dataclass @@ -59,13 +58,13 @@ class ContextParallelMetadata: def is_prefill_context_parallel_enabled(): - return get_global_server_args().enable_prefill_context_parallel + return get_server_args().enable_prefill_context_parallel def is_prefill_cp_in_seq_split(): return ( is_prefill_context_parallel_enabled() - and get_global_server_args().prefill_cp_mode == "in-seq-split" + and get_server_args().prefill_cp_mode == "in-seq-split" ) @@ -85,7 +84,7 @@ def get_cp_padding_align_size() -> int: def is_mla_prefill_cp_enabled() -> bool: - sa = get_global_server_args() + sa = get_server_args() return sa.enable_prefill_context_parallel and sa.use_mla_backend() diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index b1e08d29b..26bb916d8 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -32,7 +32,7 @@ from sglang.srt.managers.schedule_batch import ( from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.multimodal.evs import EVSEmbeddingResult -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import flatten_nested_list, is_npu, print_warning_once from sglang.srt.utils.stale_shm_cleanup import make_shm_name from sglang.utils import logger @@ -714,7 +714,7 @@ def _adjust_embedding_length( f"tokens from multimodal embeddings." ) if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding: - chunked_prefill_size = get_global_server_args().chunked_prefill_size + chunked_prefill_size = get_server_args().chunked_prefill_size if chunked_prefill_size != -1: logger.warning( "You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill" @@ -1073,7 +1073,7 @@ def general_mm_embed_routine( for i, seq_len in enumerate(forward_batch.extend_seq_lens_cpu) if forward_batch.mm_inputs[i] is not None ] - server_args = get_global_server_args() + server_args = get_server_args() if server_args and server_args.enable_adaptive_dispatch_to_encoder: # Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches input_embeds, other_info = _embed_mm_inputs_with_split( @@ -1119,7 +1119,7 @@ def general_mm_embed_routine( feature = getattr(mm_item, "feature", None) if isinstance(feature, torch.Tensor) and feature.is_cuda: mm_item.feature = feature.to("cpu", non_blocking=True) - if get_global_server_args().language_only: + if get_server_args().language_only: precomputed_embeddings = getattr( mm_item, "precomputed_embeddings", None ) @@ -1756,7 +1756,7 @@ def _get_is_default_transport(): ) _is_default_tensor_transport = ( - _determine_tensor_transport_mode(get_global_server_args()) == "default" + _determine_tensor_transport_mode(get_server_args()) == "default" ) return _is_default_tensor_transport @@ -1798,7 +1798,7 @@ def wrap_shm_features(obj): """ Scan the object for multimodal tensors and wrap them in SHM pointers. """ - if _get_is_default_transport() or get_global_server_args().skip_tokenizer_init: + if _get_is_default_transport() or get_server_args().skip_tokenizer_init: return obj if obj.mm_inputs: @@ -1859,7 +1859,7 @@ def unwrap_shm_features(obj): Restore ShmPointerMMData wrappers back into standard torch.Tensors. Handles both single requests and batch requests. """ - if _get_is_default_transport() or get_global_server_args().skip_tokenizer_init: + if _get_is_default_transport() or get_server_args().skip_tokenizer_init: return obj # Handle batch requests if isinstance(obj, BaseBatchReq): diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 20f4f2146..085b8798d 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -106,10 +106,10 @@ from sglang.srt.observability.req_time_stats import ( DPControllerReqTimeStats, SchedulerReqTimeStats, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import SamplingParams -from sglang.srt.server_args import ServerArgs, get_global_server_args +from sglang.srt.server_args import ServerArgs from sglang.srt.utils import flatten_nested_list from sglang.srt.utils.cuda_ipc_transport_utils import CudaIpcTensorTransportProxy @@ -1035,7 +1035,7 @@ class Req(ReqDllmMixin): """Check if this request is prefill-only (no token generation needed).""" # NOTE: when spec is enabled, prefill_only optimizations are disabled - spec_alg = get_global_server_args().speculative_algorithm + spec_alg = get_server_args().speculative_algorithm return self.sampling_params.max_new_tokens == 0 and spec_alg is None @property @@ -1056,7 +1056,7 @@ class Req(ReqDllmMixin): def _cache_commit_len(self) -> int: # Report only the prompt prefix so thinking + answer fall into the # overallocated range and are reclaimed by release_kv_cache. #22373. - if get_global_server_args().strip_thinking_cache and self.reasoning_tokens > 0: + if get_server_args().strip_thinking_cache and self.reasoning_tokens > 0: return min(self.kv_committed_len, len(self.origin_input_ids)) return self.kv_committed_len @@ -2205,7 +2205,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): req.already_computed = seq_len req.is_retracted = False - if get_global_server_args().enable_mamba_extra_buffer(): + if get_server_args().enable_mamba_extra_buffer(): track_entry = self._mamba_radix_cache_v2_req_prepare_for_extend(req) mamba_track_mask_cpu.append(track_entry.track_mask) mamba_track_indices_cpu.append(track_entry.track_index) @@ -2310,7 +2310,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.extend_logprob_start_lens = extend_logprob_start_lens self.extend_input_logprob_token_ids = extend_input_logprob_token_ids - if get_global_server_args().enable_mamba_extra_buffer(): + if get_server_args().enable_mamba_extra_buffer(): self.mamba_track_indices = torch.tensor( mamba_track_indices_cpu, dtype=torch.int64, @@ -2344,7 +2344,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self, req: Req, ) -> _MambaRadixCacheV2TrackEntry: - mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size + mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size def _force_track_h(i: int) -> int: assert i % mamba_cache_chunk_size == 0 @@ -2395,7 +2395,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # In lazy mode, skip the swap — the second ping-pong slot is not # allocated yet; it will be allocated on demand at the track boundary # in mamba_lazy_prealloc_at_boundary during prepare_for_decode. - if not get_global_server_args().enable_mamba_extra_buffer_lazy(): + if not get_server_args().enable_mamba_extra_buffer_lazy(): req.mamba_next_track_idx = ( self.req_to_token_pool.get_mamba_ping_pong_other_idx( req.mamba_next_track_idx @@ -2736,15 +2736,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.req_pool_indices_cpu, ) - if get_global_server_args().enable_mamba_extra_buffer(): - mamba_track_interval = get_global_server_args().mamba_track_interval + if get_server_args().enable_mamba_extra_buffer(): + mamba_track_interval = get_server_args().mamba_track_interval if len(self.reqs) == 0: self.mamba_track_indices = torch.empty( (0,), dtype=torch.int64, device=self.device ) else: - if get_global_server_args().enable_mamba_extra_buffer_lazy(): + if get_server_args().enable_mamba_extra_buffer_lazy(): self.mamba_lazy_prealloc_at_boundary(mamba_track_interval) set_mamba_track_indices_from_reqs(self) @@ -2932,7 +2932,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): def maybe_evict_swa(self): if self.tree_cache.supports_swa(): sliding_window_size = self.tree_cache.sliding_window_size - server_args = get_global_server_args() + server_args = get_server_args() release_leaf_lock = ( envs.SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW.get() diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index ef31b32c7..b270a0f50 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -56,7 +56,8 @@ from sglang.srt.mem_cache.multi_ended_allocator import ( UnifiedMambaTokenToKVPoolAllocator, ) from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode -from sglang.srt.server_args import ServerArgs, get_global_server_args +from sglang.srt.runtime_context import get_server_args +from sglang.srt.server_args import ServerArgs if TYPE_CHECKING: from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator @@ -185,7 +186,7 @@ class SchedulePolicy: if ( not isinstance(policy, CacheAwarePolicy) and self.tree_cache.supports_fast_match_prefix() - and get_global_server_args().disaggregation_mode != "decode" + and get_server_args().disaggregation_mode != "decode" ): for r in waiting_queue: match_prefix_for_req(self.tree_cache, r, include_req=True) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 90d548f0f..97d3a21df 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -237,9 +237,9 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.platforms import current_platform from sglang.srt.plugins import load_plugins -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo -from sglang.srt.server_args import PortArgs, ServerArgs, get_global_server_args +from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.session.session_controller import SessionController from sglang.srt.speculative.dflash_utils import validate_dflash_request from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec @@ -892,8 +892,8 @@ class Scheduler( self.min_free_slots_delayer = MinFreeSlotsDelayer( min_free_slots=min_free_slots ) - if not get_global_server_args().pp_max_micro_batch_size: - get_global_server_args().override( + if not get_server_args().pp_max_micro_batch_size: + get_server_args().override( "scheduler.pp_max_micro_batch_size_default", pp_max_micro_batch_size=max( self.max_running_requests // self.ps.pp_size, 1 @@ -2730,7 +2730,7 @@ class Scheduler( return ret def get_num_allocatable_reqs(self, running_bs): - res = get_global_server_args().pp_max_micro_batch_size - running_bs + res = get_server_args().pp_max_micro_batch_size - running_bs res = min(res, self.req_to_token_pool.available_size()) return res @@ -3766,7 +3766,7 @@ class Scheduler( return success def get_internal_state(self, recv_req: GetInternalStateReq): - ret = dict(vars(get_global_server_args())) # vars returns a ref to obj.__dict__ + ret = dict(vars(get_server_args())) # vars returns a ref to obj.__dict__ ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput ret["memory_usage"] = { "weight": round(self.tp_worker.model_runner.weight_load_mem_usage, 2), @@ -3833,11 +3833,10 @@ class Scheduler( self.metrics_reporter.spec_total_num_accept_tokens = ( self.metrics_reporter.spec_total_num_forward_ct ) = 0 - for k, v in server_args_dict.items(): - setattr(get_global_server_args(), k, v) - logger.info(f"Global server args updated! {get_global_server_args()=}") + get_server_args().override(source="update_server_args", **server_args_dict) + logger.info(f"Global server args updated! {get_server_args()=}") - server_args = dict(vars(get_global_server_args())) + server_args = dict(vars(get_server_args())) # This field is not serializable. server_args.pop("model_config", None) return SetInternalStateReqOutput( diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 6e20e6532..22343ee2e 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -25,7 +25,7 @@ from sglang.srt.mem_cache.common import ( maybe_cache_unfinished_req, release_kv_cache, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer @@ -848,7 +848,7 @@ class SchedulerBatchResultProcessor: prepare_release(req) is_insert = ( req.mamba_lazy_is_insert - if get_global_server_args().enable_mamba_extra_buffer_lazy() + if get_server_args().enable_mamba_extra_buffer_lazy() else True ) release_kv_cache(req, self.tree_cache, is_insert=is_insert) @@ -883,7 +883,7 @@ class SchedulerBatchResultProcessor: if req.mamba_ping_pong_track_buffer is None: return - lazy = get_global_server_args().enable_mamba_extra_buffer_lazy() + lazy = get_server_args().enable_mamba_extra_buffer_lazy() at_boundary, track_seqlen = self._mamba_check_track_boundary( req, batch, result, i ) @@ -915,7 +915,7 @@ class SchedulerBatchResultProcessor: For spec decode, the boundary is detected by comparing the accepted seq_len range against interval boundaries. """ - interval = get_global_server_args().mamba_track_interval + interval = get_server_args().mamba_track_interval if batch.spec_algorithm.is_none(): if req.kv_committed_len % interval == 0: diff --git a/python/sglang/srt/managers/scheduler_components/profiler_manager.py b/python/sglang/srt/managers/scheduler_components/profiler_manager.py index 90945b82d..a4c10fc7e 100644 --- a/python/sglang/srt/managers/scheduler_components/profiler_manager.py +++ b/python/sglang/srt/managers/scheduler_components/profiler_manager.py @@ -18,7 +18,7 @@ import torch from sglang.srt.environ import envs from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType from sglang.srt.model_executor.forward_batch_info import ForwardMode -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import is_mps, is_npu from sglang.srt.utils.profile_merger import ProfileMerger from sglang.srt.utils.profile_utils import ProfileManager @@ -245,7 +245,7 @@ class SchedulerProfilerManager: self.profile_in_progress = True if "CUDA_PROFILER" in activities: - if self.ps.gpu_id == get_global_server_args().base_gpu_id: + if self.ps.gpu_id == get_server_args().base_gpu_id: torch.cuda.cudart().cudaProfilerStart() self.profile_in_progress = True @@ -355,7 +355,7 @@ class SchedulerProfilerManager: torch.cuda.memory._record_memory_history(enabled=None) if "CUDA_PROFILER" in self.profiler_activities: - if self.ps.gpu_id == get_global_server_args().base_gpu_id: + if self.ps.gpu_id == get_server_args().base_gpu_id: torch.cuda.cudart().cudaProfilerStop() merge_message = self._merge_profile_traces() diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index b5c66a349..2d9f54402 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -216,9 +216,9 @@ class BasePrefixCache(ABC, PrefixCacheTrait): ) def init_metrics_collector(self): - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - server_args = get_global_server_args() + server_args = get_server_args() labels = {"cache_type": self.__class__.__name__} if server_args.extra_metric_labels: labels.update(server_args.extra_metric_labels) diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 461a795e0..a58f8970b 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -26,7 +26,7 @@ from sglang.srt.mem_cache.triton_ops.common import ( write_req_to_token_pool_triton, ) from sglang.srt.runtime_context import get_server_args -from sglang.srt.server_args import ServerArgs, get_global_server_args +from sglang.srt.server_args import ServerArgs from sglang.srt.utils import is_cuda, is_hip, is_npu, support_triton from sglang.srt.utils.common import ceil_align, is_pin_memory_available @@ -214,7 +214,7 @@ def get_last_loc_torch( def get_alloc_len_per_decode(server_args: Optional[ServerArgs] = None) -> int: if server_args is None: - server_args = get_global_server_args() + server_args = get_server_args() if server_args.speculative_algorithm is None: return 1 @@ -444,7 +444,7 @@ def _alloc_page_size(batch: ScheduleBatch) -> int: # DCP swaps in an allocator whose page_size is server_args.page_size * # dcp_size, so it can be > 1 even when tree_cache.page_size is 1; branch on # the real allocator's page_size there. Elsewhere the two are equal. - if (_is_hip or _is_cuda) and get_global_server_args().dcp_size > 1: + if (_is_hip or _is_cuda) and get_server_args().dcp_size > 1: return batch.tree_cache.token_to_kv_pool_allocator.page_size return batch.tree_cache.page_size @@ -658,7 +658,7 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr start_p, end_p = req.pop_overallocated_kv_cache() - global_server_args = get_global_server_args() + global_server_args = get_server_args() page_size = global_server_args.page_size spec_algo = global_server_args.speculative_algorithm diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index f3bfef414..6176a0b12 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -21,7 +21,7 @@ from sglang.srt.layers.attention.dsv4.index_buf_accessor import NopeFp8RopeBf16P from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.memory_pool import KVCache -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ceil_div, is_hip logger = logging.getLogger(__name__) @@ -276,7 +276,7 @@ class DeepSeekV4IndexerPool(KVCache): end_layer, ) self.index_head_dim = index_head_dim - self.use_fp4_indexer = get_global_server_args().enable_deepseek_v4_fp4_indexer + self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer self._create_buffer() @@ -569,7 +569,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.swa_kv_pool = None self.c4_kv_pool = None self.c128_kv_pool = None - server_args = get_global_server_args() + server_args = get_server_args() spec_extra = ( (server_args.speculative_num_draft_tokens - 1) if server_args.speculative_algorithm is not None @@ -651,7 +651,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.full_to_swa_index_mapping = full_to_swa_index_mapping def get_ring_size(self, compress_ratio: int) -> int: - server_args = get_global_server_args() + server_args = get_server_args() is_speculative = server_args.speculative_algorithm is not None return get_compress_state_ring_size(compress_ratio, is_speculative) diff --git a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py index c317f5790..4c75e58e8 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -1030,9 +1030,9 @@ class HiMambaRadixCache(MambaRadixCache): node_update = node_update.parent if len(value) > best_value_len: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size + mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size mamba_cache_chunk_aligned_seqlen = ( sum(len(v) for v in value) // mamba_cache_chunk_size ) * mamba_cache_chunk_size @@ -1272,10 +1272,10 @@ class HiMambaRadixCache(MambaRadixCache): } if extra_metric_labels: labels.update(extra_metric_labels) - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args storage_cls = resolve_collector_class( - get_global_server_args(), + get_server_args(), STAT_LOGGER_ROLE_STORAGE, StorageMetricsCollector, ) diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index a89635a71..bfbe91ffe 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -338,10 +338,10 @@ class HiRadixCache(RadixCache): labels.update(extra_metric_labels) existing_collector = getattr(self, "storage_metrics_collector", None) if existing_collector is None: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args storage_cls = resolve_collector_class( - get_global_server_args(), + get_server_args(), STAT_LOGGER_ROLE_STORAGE, StorageMetricsCollector, ) diff --git a/python/sglang/srt/mem_cache/mamba_checkpoint_pool.py b/python/sglang/srt/mem_cache/mamba_checkpoint_pool.py index 13c30bdc9..7dad73864 100644 --- a/python/sglang/srt/mem_cache/mamba_checkpoint_pool.py +++ b/python/sglang/srt/mem_cache/mamba_checkpoint_pool.py @@ -313,10 +313,10 @@ def maybe_init_int8_mamba_checkpoint_pool( allocating, so an oversized ``--int8-mamba-ckpt-size`` fails with an actionable message instead of a cryptic mid-allocation CUDA OOM. """ - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args try: - _sa = get_global_server_args() + _sa = get_server_args() except ValueError: # Some unit-test / mock runners construct HybridReqToTokenPool directly # without a global server-args context. The int8 checkpoint pool is opt-in diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index 63f814a38..4b2961181 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -50,7 +50,7 @@ from sglang.srt.mem_cache.multi_ended_allocator import ( ) from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.utils import split_node_hash_value -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req @@ -435,7 +435,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): ) self.req_to_token_pool: HybridReqToTokenPool = params.req_to_token_pool self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator - self.mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size + self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size self.page_size = params.page_size self.disable = params.disable diff --git a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py index 84cce04a3..96e2f18b0 100644 --- a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py +++ b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py @@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import ( MatchResult, ) from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args try: from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector @@ -108,7 +108,7 @@ class LMCRadixCache(RadixCache): ): super().__init__(params) - cli_lmc_cfg = get_global_server_args().lmcache_config_file or "" + cli_lmc_cfg = get_server_args().lmcache_config_file or "" kvcache = self.token_to_kv_pool_allocator.get_kvcache() connector_kwargs = dict( @@ -438,7 +438,7 @@ class LMCRadixCache(RadixCache): self.lmcache_connector.end_session(req.rid) return - global_server_args = get_global_server_args() + global_server_args = get_server_args() topk = global_server_args.speculative_eagle_topk enable_kv_committed_len = topk is None or topk == 1 if enable_kv_committed_len: diff --git a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py index 48ec5e17d..dae605f35 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py @@ -26,7 +26,7 @@ from sglang.srt.mem_cache.unified_cache_components.tree_component import ( TreeComponent, get_and_increase_time_counter, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req @@ -84,7 +84,7 @@ class MambaComponent(TreeComponent): # states. We temporarily skip branching-state fill in that mode and can # add a HiCache-aware branching policy later. if self.cache.cache_controller is None and len(value_chunks) > best_value_len: - chunk_size = get_global_server_args().mamba_cache_chunk_size + chunk_size = get_server_args().mamba_cache_chunk_size aligned_seqlen = ( sum(len(v) for v in value_chunks) // chunk_size ) * chunk_size diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 239248e17..53db89b7b 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -2280,10 +2280,10 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): labels.update(extra_metric_labels) existing_collector = self.storage_metrics_collector if existing_collector is None: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args storage_cls = resolve_collector_class( - get_global_server_args(), + get_server_args(), STAT_LOGGER_ROLE_STORAGE, StorageMetricsCollector, ) diff --git a/python/sglang/srt/model_executor/cuda_graph_config.py b/python/sglang/srt/model_executor/cuda_graph_config.py index 271bef043..2588f2c15 100644 --- a/python/sglang/srt/model_executor/cuda_graph_config.py +++ b/python/sglang/srt/model_executor/cuda_graph_config.py @@ -16,7 +16,7 @@ cuda_graph_config, and the --cuda-graph-config JSON CLI parser. Module-level imports are pure stdlib — no torch / sglang.srt deps — so ServerArgs can import everything here without pulling in backend -classes. check_cuda_graph_backend lazy-imports get_global_server_args +classes. check_cuda_graph_backend lazy-imports get_server_args inside the function body to preserve that invariant. """ @@ -164,10 +164,10 @@ def check_cuda_graph_backend(phase: str, backend: str) -> bool: """True if cuda_graph_config[phase].backend == backend on the global server args. Returns False if the global server args have not been initialized yet (e.g. unit tests, early startup).""" - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args try: - server_args = get_global_server_args() + server_args = get_server_args() except ValueError: return False cfg = server_args.cuda_graph_config diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 464cc5b37..f6bf09c35 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -49,8 +49,7 @@ from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import ( ForwardBatchDeepSeekMHAMixin, ) from sglang.srt.model_executor.triton_ops.position import compute_position_triton -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import ( is_cuda, is_hip, @@ -910,7 +909,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # --enable-mis: every request must carry delimiter indices (the score # endpoint always produces MIS-structured requests; consumers index # without None-checking). - if get_global_server_args().enable_mis and any( + if get_server_args().enable_mis and any( r.multi_item_delimiter_indices is not None for r in batch.reqs ): assert all( @@ -1073,7 +1072,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # 3 * N if ( mm_input is None - or get_global_server_args().rl_on_policy_target is not None + or get_server_args().rl_on_policy_target is not None ): mrope_positions_list[batch_idx] = torch.full( (3, 1), @@ -1092,7 +1091,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): ) if ( mm_input is None - or get_global_server_args().rl_on_policy_target is not None + or get_server_args().rl_on_policy_target is not None ): # text only mrope_positions = torch.tensor( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 3809cabff..4cd22513e 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -780,9 +780,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): torchao_applied = getattr(self.model, "torchao_applied", False) # In layered loading, torchao may have been applied if not torchao_applied: - apply_torchao_config_to_model( - self.model, get_global_server_args().torchao_config - ) + apply_torchao_config_to_model(self.model, get_server_args().torchao_config) # Apply torch TP if the model supports it supports_torch_tp = getattr(self.model, "supports_torch_tp", False) @@ -1007,7 +1005,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): set_global_experts_capturer( RoutedExpertsCapturer.create( - enable=get_global_server_args().enable_return_routed_experts, + enable=get_server_args().enable_return_routed_experts, model_config=self.model_config, num_fused_shared_experts=num_fused_shared_experts, num_tokens=self.max_total_num_tokens + self.page_size, @@ -1017,7 +1015,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) def init_indexer_capturer(self): - enable = get_global_server_args().enable_return_indexer_topk + enable = get_server_args().enable_return_indexer_topk # Producer wiring is CUDA-only (Indexer.forward_cuda + MLA skip_topk # path); other backends would create a capturer but never feed it. if enable and self.device != "cuda": @@ -1709,8 +1707,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): else: # Load the missing weights from disk self.update_weights_from_disk( - get_global_server_args().model_path, - get_global_server_args().load_format, + get_server_args().model_path, + get_server_args().load_format, weight_name_filter=weight_name_filter, ) diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 21aba97f4..e11563564 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -44,7 +44,7 @@ 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.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import get_available_gpu_memory # Try to import accelerate (optional dependency) @@ -481,7 +481,7 @@ class DefaultModelLoader(BaseModelLoader): else: hf_folder = model_name_or_path - server_args = get_global_server_args() + server_args = get_server_args() if server_args and server_args.model_checksum is not None: from sglang.srt.utils.model_file_verifier import verify @@ -567,7 +567,7 @@ class DefaultModelLoader(BaseModelLoader): hf_weights_files, ) elif use_safetensors: - server_args = get_global_server_args() + server_args = get_server_args() weight_loader_disable_mmap = server_args.weight_loader_disable_mmap weight_loader_prefetch = server_args.weight_loader_prefetch_checkpoints prefetch_num_threads = server_args.weight_loader_prefetch_num_threads @@ -866,9 +866,9 @@ class LayeredModelLoader(DefaultModelLoader): device_config: DeviceConfig, ) -> nn.Module: from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - torchao_config = get_global_server_args().torchao_config + torchao_config = get_server_args().torchao_config target_device = torch.device(device_config.device) quant_config = _get_quantization_config(model_config, self.load_config) @@ -3078,7 +3078,7 @@ class RunaiModelStreamerLoader(BaseModelLoader): ) ) - server_args = get_global_server_args() + server_args = get_server_args() if server_args and server_args.model_checksum is not None: from sglang.srt.utils.model_file_verifier import verify diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 4ac27a69c..98063401c 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -78,7 +78,6 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers LoraConfig = None @@ -209,7 +208,7 @@ class BailingMoESparseMoeBlock(nn.Module): self.router_dtype = torch.bfloat16 # TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now - assert get_global_server_args().ep_num_redundant_experts == 0 + assert get_server_args().ep_num_redundant_experts == 0 # check group topk self.num_expert_group = getattr(config, "n_group", 0) self.topk_group = getattr(config, "topk_group", 0) @@ -224,7 +223,7 @@ class BailingMoESparseMoeBlock(nn.Module): self.use_grouped_topk = False self.num_experts = ( - config.num_experts + get_global_server_args().ep_num_redundant_experts + config.num_experts + get_server_args().ep_num_redundant_experts ) self.gate = BailingMoEGate( diff --git a/python/sglang/srt/models/bailing_moe_linear.py b/python/sglang/srt/models/bailing_moe_linear.py index bf46a1eb4..d87038d31 100644 --- a/python/sglang/srt/models/bailing_moe_linear.py +++ b/python/sglang/srt/models/bailing_moe_linear.py @@ -59,7 +59,6 @@ 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_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( BumpAllocator, add_prefix, @@ -534,7 +533,7 @@ class BailingMoELinearAttention(nn.Module): base=self.rope_theta, rope_scaling=config.rope_scaling, is_neox_style=True, - device=get_global_server_args().device, + device=get_server_args().device, dtype=torch.float32, ) @@ -695,7 +694,7 @@ class BailingMoEAttention(nn.Module): max_position=self.max_position_embeddings, base=self.rope_theta, rope_scaling=config.rope_scaling, - device=get_global_server_args().device, + device=get_server_args().device, ) self.attn = RadixAttention( self.num_heads, diff --git a/python/sglang/srt/models/bert.py b/python/sglang/srt/models/bert.py index ed81a26c4..82881395f 100644 --- a/python/sglang/srt/models/bert.py +++ b/python/sglang/srt/models/bert.py @@ -16,8 +16,7 @@ from sglang.srt.layers.radix_attention import AttentionType, RadixAttention from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import add_prefix BertConfig = None @@ -367,9 +366,7 @@ class BertModel(nn.Module): prefix=add_prefix("encoder", prefix), ) pooling_type = ( - PoolingType.CLS - if get_global_server_args().is_embedding - else PoolingType.LAST + PoolingType.CLS if get_server_args().is_embedding else PoolingType.LAST ) self.pooler = ( BertPooler(config) diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index 15862b3af..3c36909a1 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -11,7 +11,7 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods AttnForwardMethod, ) from sglang.srt.models.deepseek_common.utils import _is_hip -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import use_intel_amx_backend MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"] @@ -114,7 +114,7 @@ def handle_attention_flashinfer(attn, forward_batch): def handle_attention_fa3(attn, forward_batch): # when deterministic inference is enabled, use MLA - if get_global_server_args().enable_deterministic_inference: + if get_server_args().enable_deterministic_inference: return _dispatch_mla_subtype(attn, forward_batch) else: return _handle_attention_backend(attn, forward_batch, "fa3") @@ -183,7 +183,7 @@ def handle_attention_triton(attn, forward_batch): return AttnForwardMethod.MLA # when deterministic inference is enabled, use MLA - if get_global_server_args().enable_deterministic_inference: + if get_server_args().enable_deterministic_inference: return _dispatch_mla_subtype(attn, forward_batch) if ( diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index 7c3cd571e..3cf0142c1 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -31,7 +31,7 @@ from sglang.srt.models.deepseek_common.utils import ( _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2 _use_fp8_prefill_attn = ( @@ -115,7 +115,7 @@ class DeepseekMHAForwardMixin: def init_mha_forward(self: DeepseekV2AttentionMLA): self.disable_chunked_prefix_cache = ( - get_global_server_args().disable_chunked_prefix_cache + get_server_args().disable_chunked_prefix_cache ) # TODO: Design a finer way to determine the threshold @@ -279,8 +279,8 @@ class DeepseekMHAForwardMixin: self.use_dsa and self.kv_cache_dtype == "fp8_e4m3" and ( - not get_global_server_args().dsa_decode_backend == "trtllm" - or not get_global_server_args().dsa_prefill_backend == "trtllm" + not get_server_args().dsa_decode_backend == "trtllm" + or not get_server_args().dsa_prefill_backend == "trtllm" ) ): # FP8 path: dequantize DSA-specific FP8 format to BF16 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 4d2d52427..f4ac0bc65 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 @@ -67,7 +67,7 @@ from sglang.srt.models.deepseek_common.utils import ( _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.state_capturer.indexer_topk import ( maybe_capture_indexer_topk, ) @@ -178,7 +178,7 @@ def _should_defer_dsa_cp_kv_gather( class DeepseekMLAForwardMixin: def init_mla_forward(self: DeepseekV2AttentionMLA): self.flashinfer_mla_disable_ragged = ( - get_global_server_args().flashinfer_mla_disable_ragged + get_server_args().flashinfer_mla_disable_ragged ) def should_run_indexer( @@ -1002,8 +1002,8 @@ class DeepseekMLAForwardMixin: """ if self.current_attention_backend in ("dsa", "nsa"): return ( - get_global_server_args().dsa_decode_backend == "trtllm" - or get_global_server_args().dsa_prefill_backend == "trtllm" + get_server_args().dsa_decode_backend == "trtllm" + or get_server_args().dsa_prefill_backend == "trtllm" ) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn return ( @@ -1020,7 +1020,7 @@ class DeepseekMLAForwardMixin: """ Check if we should skip rope and use fused rope+cache path for TileLang DSA on gfx95. """ - server_args = get_global_server_args() + server_args = get_server_args() return ( _use_aiter_gfx95 and self.current_attention_backend in ("dsa", "nsa") diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index fb98511f1..7e8cc0b9b 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -58,7 +58,6 @@ from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_t from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM from sglang.srt.models.utils import WeightsMapper from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu logger = logging.getLogger(__name__) @@ -119,7 +118,7 @@ class DeepseekModelNextN(nn.Module): self.rot_weight = None if _is_npu: - rot_weight_path = get_global_server_args().model_path + "/rot.safetensors" + rot_weight_path = get_server_args().model_path + "/rot.safetensors" if os.path.isfile(rot_weight_path): self.rot_weight = load_file(rot_weight_path) self.rot_weight = self.rot_weight["rot.weight"].npu() @@ -132,8 +131,8 @@ class DeepseekModelNextN(nn.Module): layer_name = "decoder" if _is_npu and ( - get_global_server_args().speculative_draft_model_path - == get_global_server_args().model_path + get_server_args().speculative_draft_model_path + == get_server_args().model_path ): layer_name = "layers." + str(config.num_hidden_layers) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 51dd8da8f..670d2adcb 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -180,7 +180,6 @@ from sglang.srt.models.deepseek_common.utils import ( _use_aiter_gfx95, ) from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( BumpAllocator, @@ -484,7 +483,7 @@ class MoEGate(nn.Module): True, # is_vnni ) - if get_global_server_args().enable_deterministic_inference: + if get_server_args().enable_deterministic_inference: return F.linear(hidden_states, self.weight, None) if ( @@ -621,7 +620,7 @@ class DeepseekV2MoE(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=num_experts_for_moe - + get_global_server_args().ep_num_redundant_experts, + + get_server_args().ep_num_redundant_experts, num_fused_shared_experts=self.num_fused_shared_experts, top_k=top_k_for_moe, hidden_size=config.hidden_size, @@ -794,8 +793,7 @@ class DeepseekV2MoE(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts - + get_global_server_args().ep_num_redundant_experts + config.n_routed_experts + get_server_args().ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group @@ -836,7 +834,7 @@ class DeepseekV2MoE(nn.Module): self, hidden_states: torch.Tensor, server_args=None ) -> bool: if server_args is None: - server_args = get_global_server_args() + server_args = get_server_args() return ( _enable_pcg_dsv2_dual_stream and (is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph()) @@ -874,7 +872,7 @@ class DeepseekV2MoE(nn.Module): ) if not self._enable_a2a_moe: - server_args = get_global_server_args() + server_args = get_server_args() if self._can_dual_stream_graph(hidden_states, server_args): return dsv2_flashinfer_moe_dual_stream_graph( hidden_states, @@ -935,7 +933,7 @@ class DeepseekV2MoE(nn.Module): # into the decode CUDA graph and replays from null. current_stream = torch.cuda.current_stream() self.alt_stream.wait_stream(current_stream) - server_args = get_global_server_args() + server_args = get_server_args() dispatch_info = ( ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id) if server_args.enable_eplb @@ -1032,7 +1030,7 @@ class DeepseekV2MoE(nn.Module): self.shared_experts.gate_up_proj ): return self.forward_cpu(hidden_states, should_allreduce_fusion) - server_args = get_global_server_args() + server_args = get_server_args() dispatch_info = ( ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id) if server_args.enable_eplb @@ -1599,7 +1597,7 @@ class DeepseekV2AttentionMLA( self.scaling = self.qk_head_dim**-0.5 self.rope_theta = rope_theta self.max_position_embeddings = max_position_embeddings - self.kv_cache_dtype = get_global_server_args().kv_cache_dtype + self.kv_cache_dtype = get_server_args().kv_cache_dtype # NOTE modification to rope_scaling must be done early enough, b/c e.g. Indexer needs it if rope_scaling: @@ -1711,7 +1709,7 @@ class DeepseekV2AttentionMLA( base=rope_theta, rope_scaling=rope_scaling, is_neox_style=is_neox_style, - device=get_global_server_args().device, + device=get_server_args().device, ) if rope_scaling and rope_scaling.get("apply_yarn_scaling", True): @@ -1798,7 +1796,7 @@ class DeepseekV2AttentionMLA( # Determine attention backend name for current forward batch: prefer the # name stamped per-runner on the backend object, else resolve from server args. backend = get_attn_backend() - server_args = get_global_server_args() + server_args = get_server_args() default_prefill_str, default_decode_str = server_args.get_attention_backends() prefill_backend_str = ( backend.prefill_attention_backend_str or default_prefill_str @@ -2071,7 +2069,7 @@ class DeepseekV2DecoderLayer(nn.Module): rope_scaling = config.rope_scaling max_position_embeddings = config.max_position_embeddings self.speculative_algorithm = SpeculativeAlgorithm.from_string( - get_global_server_args().speculative_algorithm + get_server_args().speculative_algorithm ) self.dsa_enable_prefill_cp = dsa_enable_prefill_cp self.mla_enable_prefill_cp = mla_enable_prefill_cp @@ -2749,7 +2747,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): self, architecture: str = "DeepseekV3ForCausalLM" ): self.num_fused_shared_experts = 0 - server_args = get_global_server_args() + server_args = get_server_args() if get_server_args().disable_shared_experts_fusion: return diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 21ffd8d76..e98136c1f 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -135,7 +135,7 @@ if _is_xpu: else: from sglang.srt.layers.mhc import hc_split_sinkhorn, mhc_fused_post_pre, npu_hc_pre -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( LazyValue, add_prefix, @@ -346,7 +346,7 @@ class MQALayer(nn.Module): base=rope_base, rope_scaling=rope_scaling, is_neox_style=False, - device=get_global_server_args().device, + device=get_server_args().device, ) from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis @@ -2210,7 +2210,7 @@ class DeepseekV4ForCausalLM(nn.Module): return disable_reason = None - if get_global_server_args().enforce_shared_experts_fusion: + if get_server_args().enforce_shared_experts_fusion: if self.config.n_shared_experts != 1: raise ValueError( "DeepSeek V4 shared-experts fusion expects exactly one shared " diff --git a/python/sglang/srt/models/exaone_moe.py b/python/sglang/srt/models/exaone_moe.py index 32195bea3..c6a851a92 100755 --- a/python/sglang/srt/models/exaone_moe.py +++ b/python/sglang/srt/models/exaone_moe.py @@ -63,7 +63,6 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTe 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_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -171,8 +170,7 @@ class ExaoneMoESparseMoEBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts - + get_global_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, @@ -213,7 +211,7 @@ class ExaoneMoESparseMoEBlock(nn.Module): if get_moe_a2a_backend().is_deepep(): self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_global_server_args().ep_num_redundant_experts + config.num_experts + get_server_args().ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok diff --git a/python/sglang/srt/models/gemma4_causal.py b/python/sglang/srt/models/gemma4_causal.py index 339257ca1..b8c800f5c 100644 --- a/python/sglang/srt/models/gemma4_causal.py +++ b/python/sglang/srt/models/gemma4_causal.py @@ -58,8 +58,7 @@ from sglang.srt.models.gemma3_causal import Gemma3MLP, Gemma3TextScaledWordEmbed from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import add_prefix, make_layers logger = logging.getLogger(__name__) @@ -255,8 +254,7 @@ class Gemma4MoE(nn.Module): experts_type = get_moe_impl_class(quant_config) self.experts = experts_type( - num_experts=config.num_experts - + get_global_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, layer_id=layer_id, @@ -789,7 +787,7 @@ class Gemma4TextModel(PreTrainedModel): # combination until the runner becomes schema-aware; users can run # PP + PLE eagerly with --disable-cuda-graph. if self.pp_group.world_size > 1 and self.hidden_size_per_layer_input > 0: - sa = get_global_server_args() + sa = get_server_args() if sa is not None and not sa.disable_cuda_graph: raise ValueError( "Pipeline parallelism is currently incompatible with " diff --git a/python/sglang/srt/models/gemma4_vision.py b/python/sglang/srt/models/gemma4_vision.py index c4c6fb1c2..7e440555c 100644 --- a/python/sglang/srt/models/gemma4_vision.py +++ b/python/sglang/srt/models/gemma4_vision.py @@ -181,9 +181,9 @@ class Gemma4VisionAttention(nn.Module): @staticmethod def _select_backend() -> str: """Mirror VisionAttention._determine_attention_backend for consistency.""" - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - override = get_global_server_args().mm_attention_backend + override = get_server_args().mm_attention_backend if override is not None: return override if is_cuda(): diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 2cfceb612..cab180951 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -88,7 +88,6 @@ from sglang.srt.runtime_context import ( get_server_args, get_stream, ) -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -530,8 +529,7 @@ class Glm4MoeSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts - + get_global_server_args().ep_num_redundant_experts + config.n_routed_experts + get_server_args().ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index ecc1df2e2..0ad8e73af 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -75,7 +75,6 @@ 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_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( BumpAllocator, LazyValue, @@ -216,7 +215,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.n_routed_experts + self.num_fused_shared_experts - + get_global_server_args().ep_num_redundant_experts, + + get_server_args().ep_num_redundant_experts, num_fused_shared_experts=self.num_fused_shared_experts, top_k=config.num_experts_per_tok + self.num_fused_shared_experts, hidden_size=config.hidden_size, @@ -284,8 +283,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts - + get_global_server_args().ep_num_redundant_experts + config.n_routed_experts + get_server_args().ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group diff --git a/python/sglang/srt/models/glm4_moe_lite_nextn.py b/python/sglang/srt/models/glm4_moe_lite_nextn.py index 9cb52ff9b..a8ad68b18 100644 --- a/python/sglang/srt/models/glm4_moe_lite_nextn.py +++ b/python/sglang/srt/models/glm4_moe_lite_nextn.py @@ -36,7 +36,6 @@ from sglang.srt.models.glm4_moe_lite import ( Glm4MoeLiteForCausalLM, ) from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import BumpAllocator, add_prefix, is_npu logger = logging.getLogger(__name__) @@ -140,10 +139,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM): nn.Module.__init__(self) self.config = config self.tp_size = get_parallel().tp_size - if ( - is_npu() - and get_global_server_args().speculative_draft_model_quantization is None - ): + if is_npu() and get_server_args().speculative_draft_model_quantization is None: quant_config = None self.quant_config = quant_config diff --git a/python/sglang/srt/models/glm4_moe_nextn.py b/python/sglang/srt/models/glm4_moe_nextn.py index 732c20aa3..3126fd026 100644 --- a/python/sglang/srt/models/glm4_moe_nextn.py +++ b/python/sglang/srt/models/glm4_moe_nextn.py @@ -33,7 +33,6 @@ 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_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -126,10 +125,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM): nn.Module.__init__(self) self.config = config self.tp_size = get_parallel().tp_size - if ( - is_npu() - and get_global_server_args().speculative_draft_model_quantization is None - ): + if is_npu() and get_server_args().speculative_draft_model_quantization is None: quant_config = None self.quant_config = quant_config diff --git a/python/sglang/srt/models/glm4v.py b/python/sglang/srt/models/glm4v.py index 7fafa200d..8ca87e2e1 100644 --- a/python/sglang/srt/models/glm4v.py +++ b/python/sglang/srt/models/glm4v.py @@ -53,8 +53,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.glm4 import Glm4Model from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils.hf_transformers_utils import get_processor @@ -548,7 +547,7 @@ class Glm4vForConditionalGeneration(nn.Module): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_server_args().mm_enable_dp_encoder vision_utils.update_vit_attn_dummy_heads_config(self.config) self.visual = Glm4vVisionModel( config.vision_config, diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py index bc1d931a2..c69899003 100644 --- a/python/sglang/srt/models/glm4v_moe.py +++ b/python/sglang/srt/models/glm4v_moe.py @@ -19,7 +19,6 @@ 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_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args 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 @@ -42,7 +41,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_server_args().mm_enable_dp_encoder vision_utils.update_vit_attn_dummy_heads_config(self.config) self.tp_size = get_parallel().tp_size self.quant_config = quant_config diff --git a/python/sglang/srt/models/glm_ocr.py b/python/sglang/srt/models/glm_ocr.py index bb74461d5..9a9935b61 100644 --- a/python/sglang/srt/models/glm_ocr.py +++ b/python/sglang/srt/models/glm_ocr.py @@ -50,7 +50,7 @@ from sglang.srt.models.glm4v import ( Glm4vVisionModel, Glm4vVisionPatchEmbed, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import add_prefix from sglang.srt.utils.hf_transformers_utils import get_processor @@ -272,7 +272,7 @@ class GlmOcrForConditionalGeneration(Glm4vForConditionalGeneration): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_server_args().mm_enable_dp_encoder self.visual = GlmOcrVisionModel( vision_config=config.vision_config, text_config=config.text_config, diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index b6dc2994a..2ed27d170 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -69,7 +69,6 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, add_prefix, @@ -228,7 +227,7 @@ class GptOssSparseMoeBlock(nn.Module): self.experts = experts_type( num_experts=config.num_local_experts - + get_global_server_args().ep_num_redundant_experts, + + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, diff --git a/python/sglang/srt/models/internvl.py b/python/sglang/srt/models/internvl.py index a59fc1164..e8f1d00be 100644 --- a/python/sglang/srt/models/internvl.py +++ b/python/sglang/srt/models/internvl.py @@ -40,8 +40,7 @@ from sglang.srt.multimodal.internvl_vit_cuda_graph_runner import ( InternViTCudaGraphRunner, ) from sglang.srt.multimodal.mm_utils import run_dp_sharded_vision_model -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import is_cuda from sglang.utils import logger @@ -498,7 +497,7 @@ class InternVLChatModel(nn.Module): ) -> None: super().__init__() self.config = config - self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_server_args().mm_enable_dp_encoder self.quant_config = quant_config vision_utils.update_vit_attn_dummy_heads_config(self.config) image_size = config.force_image_size or config.vision_config.image_size diff --git a/python/sglang/srt/models/kimi_k25.py b/python/sglang/srt/models/kimi_k25.py index 1c500f8fe..022d1aee2 100644 --- a/python/sglang/srt/models/kimi_k25.py +++ b/python/sglang/srt/models/kimi_k25.py @@ -31,7 +31,7 @@ from sglang.srt.models.deepseek_v2 import DeepseekV3ForCausalLM from sglang.srt.models.kimi_vl_moonvit import MLP2 from sglang.srt.models.utils import WeightsMapper from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -633,7 +633,7 @@ class KimiK25ForConditionalGeneration(nn.Module): super().__init__() self.config = config self.quant_config = quant_config - self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_server_args().mm_enable_dp_encoder # Create vision tower self.vision_tower = MoonViT3dPretrainedModel( config.vision_config, diff --git a/python/sglang/srt/models/laguna.py b/python/sglang/srt/models/laguna.py index ef7d9c041..a2dd566d8 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -54,7 +54,6 @@ 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.utils import apply_qk_norm from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import LazyValue, add_prefix, make_layers logger = logging.getLogger(__name__) @@ -161,8 +160,7 @@ class LagunaMoE(nn.Module): self.gate = LagunaMoEGate(config, prefix=add_prefix("gate", prefix)) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts - + get_global_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, diff --git a/python/sglang/srt/models/llada2.py b/python/sglang/srt/models/llada2.py index 81ecc3290..3c93ce986 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -77,7 +77,6 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, is_cuda, @@ -220,7 +219,7 @@ class LLaDA2MoeSparseMoeBlock(nn.Module): self.router_dtype = torch.bfloat16 # TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now - assert get_global_server_args().ep_num_redundant_experts == 0 + assert get_server_args().ep_num_redundant_experts == 0 # check group topk self.num_expert_group = getattr(config, "n_group", 0) self.topk_group = getattr(config, "topk_group", 0) @@ -235,7 +234,7 @@ class LLaDA2MoeSparseMoeBlock(nn.Module): self.use_grouped_topk = False self.num_experts = ( - config.num_experts + get_global_server_args().ep_num_redundant_experts + config.num_experts + get_server_args().ep_num_redundant_experts ) self.gate = LLaDA2MoeGate( diff --git a/python/sglang/srt/models/llama_eagle3.py b/python/sglang/srt/models/llama_eagle3.py index d177c8eed..294710d11 100644 --- a/python/sglang/srt/models/llama_eagle3.py +++ b/python/sglang/srt/models/llama_eagle3.py @@ -38,7 +38,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.llama import LlamaDecoderLayer, LlamaForCausalLM, LlamaMLP -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args class LlamaDecoderLayer(LlamaDecoderLayer): @@ -258,7 +258,7 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM): # Cache draft SWA size from server args once; consumed both by the post-init # attention patch below and by `get_attention_sliding_window_size` later. self._draft_window_size: Optional[int] = ( - get_global_server_args().speculative_draft_window_size + get_server_args().speculative_draft_window_size ) self.model = LlamaModel( diff --git a/python/sglang/srt/models/mimo_audio.py b/python/sglang/srt/models/mimo_audio.py index 24b18faac..941e616fc 100644 --- a/python/sglang/srt/models/mimo_audio.py +++ b/python/sglang/srt/models/mimo_audio.py @@ -21,7 +21,7 @@ from transformers.models.qwen2.configuration_qwen2 import Qwen2Config from transformers.models.qwen2.modeling_qwen2 import Qwen2Model from sglang.srt.layers.quantization.base_config import QuantizationConfig -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import is_cuda if is_cuda(): @@ -1223,7 +1223,7 @@ class AudioEncoderMixin: else: raise ValueError(f"Invalid projection layers: {config.projection_layers}") - model_path = get_global_server_args().model_path + model_path = get_server_args().model_path if not os.path.isdir(model_path): from huggingface_hub import snapshot_download diff --git a/python/sglang/srt/models/mimo_v2.py b/python/sglang/srt/models/mimo_v2.py index d74b9ecfe..c9deb0f40 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -77,7 +77,6 @@ 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_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, add_prefix, @@ -253,7 +252,7 @@ class MiMoV2MoE(nn.Module): experts_type = get_moe_impl_class(quant_config) self.experts = experts_type( num_experts=config.n_routed_experts - + get_global_server_args().ep_num_redundant_experts, + + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, @@ -288,8 +287,7 @@ class MiMoV2MoE(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts - + get_global_server_args().ep_num_redundant_experts + config.n_routed_experts + get_server_args().ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group diff --git a/python/sglang/srt/models/mimo_vl.py b/python/sglang/srt/models/mimo_vl.py index 838b5983a..7d178a1a9 100644 --- a/python/sglang/srt/models/mimo_vl.py +++ b/python/sglang/srt/models/mimo_vl.py @@ -18,7 +18,7 @@ from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.models.qwen2_5_vl import Qwen2_5_VisionPatchMerger, Qwen2_5_VLMLP -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import add_prefix @@ -232,7 +232,7 @@ class MiMoVisionTransformer(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.server_args = get_global_server_args() + self.server_args = get_server_args() self.vit_window_attn_types = vision_config.vit_window_attn_types patch_size: int = vision_config.patch_size temporal_patch_size: int = vision_config.temporal_patch_size diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index a97df8211..81cfd590a 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -80,8 +80,7 @@ from sglang.srt.model_loader.weight_utils import ( maybe_remap_kv_scale_name, narrow_padded_param_and_loaded_weight, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args # get_bool_env_var is defined in sglang.srt.utils.common, not sglang.srt.distributed. # Importing from the wrong module causes this file to fail import, which prevents the @@ -426,7 +425,7 @@ class MiniMaxM2QKRMSNorm: props = torch.cuda.get_device_properties(device) # probe the maximum tokens for one prefill - server_args = get_global_server_args() + server_args = get_server_args() max_tokens = server_args.chunked_prefill_size if max_tokens is None: max_tokens = server_args.model_config.context_len @@ -514,7 +513,7 @@ class MiniMaxM2MoE(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.num_local_experts - + get_global_server_args().ep_num_redundant_experts, + + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, diff --git a/python/sglang/srt/models/mllama4.py b/python/sglang/srt/models/mllama4.py index 77f89eca8..68fc8f6db 100644 --- a/python/sglang/srt/models/mllama4.py +++ b/python/sglang/srt/models/mllama4.py @@ -33,7 +33,7 @@ from sglang.srt.managers.schedule_batch import ( MultimodalInputs, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import is_cpu _is_cpu = is_cpu() @@ -477,7 +477,7 @@ class Llama4ForConditionalGeneration(nn.Module): ) self.has_vision = ( - self.has_vision_weights and get_global_server_args().enable_multimodal + self.has_vision_weights and get_server_args().enable_multimodal ) if self.has_vision: diff --git a/python/sglang/srt/models/moss_vl.py b/python/sglang/srt/models/moss_vl.py index 6886551ab..d958b6f4d 100644 --- a/python/sglang/srt/models/moss_vl.py +++ b/python/sglang/srt/models/moss_vl.py @@ -44,8 +44,7 @@ from sglang.srt.managers.schedule_batch import MultimodalInputs from sglang.srt.model_executor.forward_batch_info import ForwardBatch 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_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -989,7 +988,7 @@ class MossVLSelfAttentionDecoderLayer(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index 15362b00e..ee98ffed4 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -88,7 +88,6 @@ from sglang.srt.models.nemotron_h_utils import ( ) from sglang.srt.models.utils import WeightsMapper from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, get_current_device_stream_fast, @@ -191,7 +190,7 @@ class NemotronHMoE(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.n_routed_experts - + get_global_server_args().ep_num_redundant_experts, + + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=self.moe_hidden_size, intermediate_size=config.moe_intermediate_size, diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py index ad19bdaa7..98ae54aeb 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -50,8 +50,7 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, ) from sglang.srt.platforms import current_platform -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -97,7 +96,7 @@ class Qwen2MLP(nn.Module): x: torch.Tensor, forward_batch: ForwardBatch = None, ) -> torch.Tensor: - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: x = x.bfloat16() gate_up, _ = self.gate_up_proj(x) @@ -292,7 +291,7 @@ class Qwen2Model(nn.Module): prefix=add_prefix("embed_tokens", prefix), params_dtype=( torch.float32 - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else None ), ) @@ -328,7 +327,7 @@ class Qwen2Model(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.norm = RMSNorm( diff --git a/python/sglang/srt/models/qwen2_5_vl.py b/python/sglang/srt/models/qwen2_5_vl.py index c38b59043..ce2fa5141 100644 --- a/python/sglang/srt/models/qwen2_5_vl.py +++ b/python/sglang/srt/models/qwen2_5_vl.py @@ -72,8 +72,7 @@ from sglang.srt.models.qwen2 import Qwen2Model from sglang.srt.models.utils import RotaryPosMixin, WeightsMapper, permute_inv 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_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import add_prefix, is_cpu, is_cuda, is_npu _is_cuda = is_cuda() @@ -604,7 +603,7 @@ class Qwen2_5_VLForConditionalGeneration(nn.Module): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_server_args().mm_enable_dp_encoder if not self.config.encoder_only: self.model = Qwen2Model( diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index af1d28c61..0de257976 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -92,7 +92,6 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTe 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_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -276,10 +275,10 @@ class Qwen2MoeSparseMoeBlock(nn.Module): else config.num_experts_per_tok + self.num_fused_shared_experts ), num_experts=( - config.num_experts + get_global_server_args().ep_num_redundant_experts + config.num_experts + get_server_args().ep_num_redundant_experts if not self.enable_shared_expert_fusion else config.num_experts - + get_global_server_args().ep_num_redundant_experts + + get_server_args().ep_num_redundant_experts + self.num_fused_shared_experts ), hidden_size=config.hidden_size, @@ -338,7 +337,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_global_server_args().ep_num_redundant_experts + config.num_experts + get_server_args().ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok self.is_nextn = is_nextn diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 8543bd4b0..daaca2276 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -34,7 +34,6 @@ 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_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu Qwen3Config = None @@ -112,7 +111,7 @@ class Qwen3Attention(nn.Module): weight_dtype=torch.float32, cast_x_before_out_mul=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) @@ -272,14 +271,14 @@ class Qwen3Attention(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: hidden_states = hidden_states.bfloat16() save_kv_cache = True use_aiter_fused = ( self.use_fused_qk_norm_mrope and forward_batch.forward_mode.is_decode() - and get_global_server_args().rl_on_policy_target is None + and get_server_args().rl_on_policy_target is None ) if use_aiter_fused: @@ -299,7 +298,7 @@ class Qwen3Attention(nn.Module): forward_batch=forward_batch, ) - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -363,7 +362,7 @@ class Qwen3DecoderLayer(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py index d08ed72e6..ebda7517f 100644 --- a/python/sglang/srt/models/qwen3_5_mtp.py +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -35,7 +35,6 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM from sglang.srt.runtime_context import get_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -64,10 +63,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module): "modelopt_mixed", ): quant_config = None - if ( - is_npu() - and get_global_server_args().speculative_draft_model_quantization is None - ): + if is_npu() and get_server_args().speculative_draft_model_quantization is None: quant_config = None # Quark-quantized Qwen3.5 MXFP4 checkpoints ship the MTP module in diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 9702e11dd..1f6391e07 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -73,7 +73,6 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, add_prefix, @@ -257,8 +256,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts - + get_global_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, @@ -280,7 +278,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_global_server_args().ep_num_redundant_experts + config.num_experts + get_server_args().ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok @@ -521,7 +519,7 @@ class Qwen3MoeAttention(nn.Module): ) and self.head_dim in (64, 128, 256) _yarn_factor, _, _, _ = compute_yarn_parameters(config) self.use_fused_qk_norm_rope = ( - get_global_server_args().enable_fused_qk_norm_rope + get_server_args().enable_fused_qk_norm_rope and self.compatible_with_fused_qk_norm_rope and _is_cuda and can_use_fused_qk_norm_rope( diff --git a/python/sglang/srt/models/qwen3_next_mtp.py b/python/sglang/srt/models/qwen3_next_mtp.py index 50becff5e..ef33eb9a1 100644 --- a/python/sglang/srt/models/qwen3_next_mtp.py +++ b/python/sglang/srt/models/qwen3_next_mtp.py @@ -33,7 +33,6 @@ 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_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -52,10 +51,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): config = copy.deepcopy(config) self.config = config self.tp_size = get_parallel().tp_size - if ( - is_npu() - and get_global_server_args().speculative_draft_model_quantization is None - ): + if is_npu() and get_server_args().speculative_draft_model_quantization is None: quant_config = None self.quant_config = quant_config # if not set, model load will be broken in Qwen3NextForCausalLM load_weights() diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index b21a63a7d..aa7f7c9c1 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -69,7 +69,6 @@ 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_parallel, get_server_args -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -324,9 +323,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): self.num_position_embeddings = vision_config.num_position_embeddings self.num_grid_per_side = int(self.num_position_embeddings**0.5) self.num_grid = self.num_grid_per_side * self.num_grid_per_side - self.align_corners = ( - get_global_server_args().enable_precise_embedding_interpolation - ) + self.align_corners = get_server_args().enable_precise_embedding_interpolation self.patch_size = vision_config.patch_size self.spatial_merge_size = vision_config.spatial_merge_size self.spatial_merge_unit = self.spatial_merge_size**2 @@ -369,7 +366,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): ) workspace_buffer = None - if get_global_server_args().mm_attention_backend == "flashinfer_cudnn": + if get_server_args().mm_attention_backend == "flashinfer_cudnn": if torch.cuda.is_available() and (not _is_npu): ws_device = torch.device("cuda", torch.cuda.current_device()) else: @@ -914,7 +911,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): flashinfer_max_seqlen = 0 cu_seqlens = None - if get_global_server_args().mm_attention_backend == "flashinfer_cudnn": + if get_server_args().mm_attention_backend == "flashinfer_cudnn": # real token lens (B,) real_seq_lens = token_cu_seqlens[1:] - token_cu_seqlens[:-1] flashinfer_max_seqlen = self.bucket_flashinfer_max_seqlen( @@ -1234,7 +1231,7 @@ class Qwen3VLForConditionalGeneration(nn.Module): self.pp_group = get_pp_group() self.quant_config = quant_config - self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_server_args().mm_enable_dp_encoder self.visual = Qwen3VLMoeVisionModel( config.vision_config, diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index afa7565c0..f6e40d155 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -61,7 +61,6 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha imp DeepseekMHAForwardMixin, ) from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( BumpAllocator, add_prefix, @@ -272,8 +271,7 @@ class SarvamMoESparseMoeBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts - + get_global_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, @@ -464,7 +462,7 @@ class SarvamMoEMLAAttention(nn.Module): self.scaling = self.qk_head_dim**-0.5 self.rope_theta = rope_theta self.max_position_embeddings = max_position_embeddings - self.kv_cache_dtype = get_global_server_args().kv_cache_dtype + self.kv_cache_dtype = get_server_args().kv_cache_dtype self._server_args = None self.current_attention_backend = None @@ -622,7 +620,7 @@ class SarvamMoEMLAAttention(nn.Module): def _set_current_attention_backend(self, forward_batch: ForwardBatch) -> None: if self._server_args is None: - self._server_args = get_global_server_args() + self._server_args = get_server_args() if forward_batch.forward_mode.is_decode_or_idle(): self.current_attention_backend = ( self._server_args.decode_attention_backend @@ -778,7 +776,7 @@ class SarvamMoEMLAAttention(nn.Module): k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1) if self._server_args is None: - self._server_args = get_global_server_args() + self._server_args = get_server_args() self._set_current_attention_backend(forward_batch) forward_method = get_attn_forward_method(self._server_args, forward_batch) @@ -892,7 +890,7 @@ class SarvamMoEMLAAttention(nn.Module): k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1) if self._server_args is None: - self._server_args = get_global_server_args() + self._server_args = get_server_args() self._set_current_attention_backend(forward_batch) forward_method = get_attn_forward_method(self._server_args, forward_batch) @@ -950,7 +948,7 @@ class SarvamMoEMLAAttention(nn.Module): q_nope_out, k_nope, q_pe, k_pe, forward_batch, zero_allocator = inner_state if self._server_args is None: - self._server_args = get_global_server_args() + self._server_args = get_server_args() self._set_current_attention_backend(forward_batch) forward_method = get_attn_forward_method(self._server_args, forward_batch) diff --git a/python/sglang/srt/models/sdar.py b/python/sglang/srt/models/sdar.py index 68bc12ce7..4db26c8e4 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -42,7 +42,6 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -205,7 +204,7 @@ class SDARAttention(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, ): - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: hidden_states = hidden_states.bfloat16() qkv, _ = self.qkv_proj(hidden_states) @@ -233,7 +232,7 @@ class SDARAttention(nn.Module): ), ) - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -268,7 +267,7 @@ class SDARBlock(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( @@ -392,7 +391,7 @@ class SDARModel(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.norm = RMSNorm(self.embed_dim, eps=config.rms_norm_eps, **norm_kwargs) diff --git a/python/sglang/srt/models/sdar_moe.py b/python/sglang/srt/models/sdar_moe.py index 42c569d3f..e7fe81718 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -58,7 +58,6 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -97,8 +96,7 @@ class SDARMoeSparseMoeBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts - + get_global_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, @@ -120,7 +118,7 @@ class SDARMoeSparseMoeBlock(nn.Module): if get_moe_a2a_backend().is_deepep(): self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_global_server_args().ep_num_redundant_experts + config.num_experts + get_server_args().ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok @@ -281,7 +279,7 @@ class SDARMoeAttention(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: hidden_states = hidden_states.bfloat16() qkv, _ = self.qkv_proj(hidden_states) @@ -309,7 +307,7 @@ class SDARMoeAttention(nn.Module): ), ) - if get_global_server_args().rl_on_policy_target is not None: + if get_server_args().rl_on_policy_target is not None: q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -345,7 +343,7 @@ class SDARMoeBlock(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( @@ -483,7 +481,7 @@ class SDARMoeModel(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else {} ) self.norm = RMSNorm(self.embed_dim, eps=config.rms_norm_eps, **norm_kwargs) diff --git a/python/sglang/srt/models/step3p5.py b/python/sglang/srt/models/step3p5.py index a55ce02f5..d28929a8e 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -47,7 +47,6 @@ 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_parallel, get_server_args, get_stream -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers Step3p5Config = None @@ -149,7 +148,7 @@ class Step3p5MoEMLP(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.moe_num_experts - + get_global_server_args().ep_num_redundant_experts, + + get_server_args().ep_num_redundant_experts, top_k=config.moe_top_k, layer_id=layer_id, hidden_size=config.hidden_size, @@ -172,8 +171,7 @@ class Step3p5MoEMLP(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.moe_num_experts = ( - config.moe_num_experts - + get_global_server_args().ep_num_redundant_experts + config.moe_num_experts + get_server_args().ep_num_redundant_experts ) self.top_k = config.moe_top_k @@ -681,7 +679,7 @@ class Step3p5Model(nn.Module): prefix=add_prefix("embed_tokens", prefix), params_dtype=( torch.float32 - if get_global_server_args().rl_on_policy_target is not None + if get_server_args().rl_on_policy_target is not None else None ), ) diff --git a/python/sglang/srt/models/transformers.py b/python/sglang/srt/models/transformers.py index ced5394bf..e8e801477 100644 --- a/python/sglang/srt/models/transformers.py +++ b/python/sglang/srt/models/transformers.py @@ -65,8 +65,7 @@ from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInp 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 AutoWeightsLoader, WeightsMapper -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import get_device from sglang.srt.utils.common import direct_register_custom_op from sglang.srt.utils.hf_transformers_utils import get_hf_text_config @@ -353,7 +352,7 @@ class TransformersFusedMoE(nn.Module): expert_mapping: list, ) -> None: super().__init__() - num_redundant = get_global_server_args().ep_num_redundant_experts + num_redundant = get_server_args().ep_num_redundant_experts experts_cls = get_moe_impl_class(quant_config) self.experts = experts_cls( num_experts=num_experts + num_redundant, @@ -1232,7 +1231,7 @@ class MoEMixin: expert_mapping = self._get_expert_mapping(num_experts) # EPLB / EP tracking - num_redundant = get_global_server_args().ep_num_redundant_experts + num_redundant = get_server_args().ep_num_redundant_experts ep_size = get_parallel().moe_ep_size self.mlp_moe_layers: list[nn.Module] = [] diff --git a/python/sglang/srt/models/utils.py b/python/sglang/srt/models/utils.py index 0f6e79d00..08ad51701 100644 --- a/python/sglang/srt/models/utils.py +++ b/python/sglang/srt/models/utils.py @@ -33,7 +33,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_context import get_token_to_kv_pool 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.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import get_current_device_stream_fast, is_cuda, is_hip from sglang.srt.utils.custom_op import register_custom_op @@ -434,7 +434,7 @@ def _reshape_for_qk_norm(x: torch.Tensor, head_dim: int) -> torch.Tensor: if ( _is_cuda - and get_global_server_args().cuda_graph_config.prefill.tc_compiler == "inductor" + and get_server_args().cuda_graph_config.prefill.tc_compiler == "inductor" ): return x.view(*x.shape[:-1], -1, head_dim) return x.reshape(-1, head_dim) @@ -475,7 +475,7 @@ def apply_qk_norm( and allow_inplace # TODO(dark): this can be relaxed if needed and (q_eps == k_eps) # TODO(dark): this can also be relaxed and not envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get() - and get_global_server_args().cuda_graph_config.prefill.tc_compiler + and get_server_args().cuda_graph_config.prefill.tc_compiler != "inductor" # let inductor fuse QK norm and can_use_fused_inplace_qknorm(head_dim, q.dtype) ): diff --git a/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py b/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py index 8a5611679..10461fe86 100644 --- a/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py +++ b/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py @@ -22,7 +22,7 @@ import torch import torch.nn as nn from sglang.srt.layers.attention.vision import VisionAttention -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args class InternViTCudaGraphRunner: @@ -95,7 +95,7 @@ class InternViTCudaGraphRunner: def _warmup_once(self, key: Hashable) -> None: """Run a tiny eager warmup on the preallocated buffers to trigger lazy init.""" - override_backend = get_global_server_args().mm_attention_backend + override_backend = get_server_args().mm_attention_backend cu = self.cu[key] cu_kk = self.cu_kk[key] max_len = int(cu_kk.max().item()) if cu_kk.numel() else 0 @@ -115,7 +115,7 @@ class InternViTCudaGraphRunner: def _capture_graph(self, key: Hashable) -> None: g = torch.cuda.CUDAGraph() - override_backend = get_global_server_args().mm_attention_backend + override_backend = get_server_args().mm_attention_backend cu = self.cu[key] cu_kk = self.cu_kk[key] diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index cc2b1002e..93f4ac348 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -19,7 +19,7 @@ from sglang.srt.managers.schedule_batch import ( MultimodalInputFormat, MultimodalProcessorOutput, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( envs, is_cpu, @@ -438,12 +438,12 @@ class BaseMultimodalProcessor(ABC): and isinstance(processor.image_processor, BaseImageProcessor) and not self.server_args.disable_fast_image_processor ): - if _is_cpu or get_global_server_args().rl_on_policy_target is not None: + if _is_cpu or get_server_args().rl_on_policy_target is not None: kwargs["device"] = "cpu" elif _is_xpu: kwargs["device"] = "xpu" elif not _is_npu: - base_gpu_id = get_global_server_args().base_gpu_id + base_gpu_id = get_server_args().base_gpu_id kwargs["device"] = f"cuda:{base_gpu_id}" elif processor.__class__.__name__ not in { "Glm4vProcessor", diff --git a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py index 8819cfdab..28d50c7df 100644 --- a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py +++ b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py @@ -25,7 +25,7 @@ import torch.nn as nn from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.layers.attention.vision import VisionAttention -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args class ViTCudaGraphRunner: @@ -139,7 +139,7 @@ class ViTCudaGraphRunner: cu_full_kk = self.cu_full_len_kk[graph_key] max_full_len = int(cu_full_kk.max().item()) - override_backend = get_global_server_args().mm_attention_backend + override_backend = get_server_args().mm_attention_backend tp_group = get_tp_group() ca_comm = tp_group.ca_comm diff --git a/python/sglang/srt/observability/trace.py b/python/sglang/srt/observability/trace.py index fb74fc4bf..51fd513c5 100644 --- a/python/sglang/srt/observability/trace.py +++ b/python/sglang/srt/observability/trace.py @@ -32,7 +32,6 @@ opentelemetry_initialized = False _trace_context_propagator = None tracer: Optional[trace.Tracer] = None -global_trace_level = get_int_env_var("SGLANG_TRACE_LEVEL", 3) # Modules allowed to emit spans (from --trace-modules); None means no filtering. global_trace_modules: Optional[List[str]] = None @@ -74,9 +73,19 @@ def extract_trace_headers(headers: Mapping[str, str]) -> Optional[Dict]: return {h: headers[h] for h in TRACE_HEADERS if h in headers} +def get_global_trace_level() -> int: + from sglang.srt.runtime_context import get_resources + + resources = get_resources() + if resources.trace_level is None: + resources.trace_level = get_int_env_var("SGLANG_TRACE_LEVEL", 3) + return resources.trace_level + + def set_global_trace_level(level: int): - global global_trace_level - global_trace_level = level + from sglang.srt.runtime_context import get_resources + + get_resources().trace_level = level @dataclass @@ -268,7 +277,7 @@ class TraceReqContext: external_trace_header: Optional[Dict[str, str]] = None, ): self.rid: str = str(rid) - self.trace_level = global_trace_level + self.trace_level = get_global_trace_level() self.tracing_enable: bool = opentelemetry_initialized and self.trace_level > 0 # Filter by --trace-modules only for explicitly named modules; contexts diff --git a/python/sglang/srt/parser/template_detection.py b/python/sglang/srt/parser/template_detection.py index f6ea7d05a..24a1a1b50 100644 --- a/python/sglang/srt/parser/template_detection.py +++ b/python/sglang/srt/parser/template_detection.py @@ -552,7 +552,7 @@ def _resolve_auto_parser( """Resolve a single auto parser, updating server_args in place.""" detected = match_rules(ctx, rules, label) if detected: - setattr(server_args, attr, detected) + server_args.override(source="template-detection", **{attr: detected}) logger.info( f"Auto-detected --{attr.replace('_', '-')} as '{detected}' from chat template" ) @@ -561,7 +561,7 @@ def _resolve_auto_parser( f"--{attr.replace('_', '-')}=auto specified but could not detect " f"{label} from chat template. Disabling {label}." ) - setattr(server_args, attr, None) + server_args.override(source="template-detection", **{attr: None}) def _load_explicit_jinja_template(chat_template_arg: Optional[str]) -> Optional[str]: @@ -580,7 +580,7 @@ def _disable_auto_parser(server_args, attr: str, label: str) -> None: f"--{attr.replace('_', '-')}=auto specified but could not detect " f"{label} from chat template. Disabling {label}." ) - setattr(server_args, attr, None) + server_args.override(source="template-detection", **{attr: None}) def _resolve_architecture_auto_parsers(server_args) -> None: @@ -607,7 +607,7 @@ def _resolve_architecture_auto_parsers(server_args) -> None: ("tool_call_parser", tool_call_parser), ): if getattr(server_args, attr) == "auto": - setattr(server_args, attr, detected) + server_args.override(source="template-detection", **{attr: detected}) logger.info( f"Auto-detected --{attr.replace('_', '-')} as '{detected}' " f"from model architecture '{arch}'" diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 8fe4e21df..284814c95 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -366,6 +366,13 @@ class Resources(_FlagGroupBase): # Persistent reusable CUDA events for non-EP DP TBO, keyed by # (kind, subbatch) — see dp_attention._tbo_event for why reuse matters. tbo_event_pool: dict = dataclasses.field(default_factory=dict) + # State capturers (installed by their subsystems when capture is on). + indexer_capturer: Any = None + experts_capturer: Any = None + # The shared TCPStore created during distributed initialization. + tcp_store: Any = None + # Trace verbosity; the accessor seeds it lazily from SGLANG_TRACE_LEVEL. + trace_level: Any = None class ForwardFlags: @@ -551,6 +558,87 @@ class RuntimeContext: ) self._server_args = server_args + def override_server_args(self, **fields) -> _ServerArgsOverride: + """Test-only scoped override for the config tier — the sibling of + ``get_parallel().override()`` and the flag groups' ``override()``: + tests force execution paths by overriding the context instead of + hand-building config objects. + + ``install()`` (or entering it as a context manager) publishes a fresh + dummy-boundary ``ServerArgs`` carrying ``fields`` and returns it; + ``restore()`` (or exiting) reinstates whatever the slot held before. + + Transitional — to be deprecated: it exists because production code + still branches on raw ``server_args`` fields at runtime, so forcing a + path needs a full config in the slot. As those readers migrate onto + the named runtime tiers (flags / resources / forward), prefer the + finer-grained overrides; once they cover the branching surface this + override loses its clients and goes away. + """ + return _ServerArgsOverride(self, fields) + + +class _ServerArgsOverride: + """Scoped config override (see ``RuntimeContext.override_server_args``). + + Deliberately a plain class rather than a generator context manager: + fixtures that live for a whole test case install the override without a + ``with`` block, and a suspended generator would run its restore whenever + the garbage collector closes it — un-publishing the active config at a + nondeterministic point. + """ + + __slots__ = ("_context", "_fields", "_previous", "_previous_capture", "_installed") + + def __init__(self, context: RuntimeContext, fields: dict): + self._context = context + self._fields = fields + self._previous: ServerArgs | None = None + self._previous_capture = False + self._installed = False + + def install(self) -> ServerArgs: + """Publish a fresh dummy-boundary ``ServerArgs`` carrying the + overrides (written through ``ServerArgs.override`` for provenance); + returns the published instance.""" + from sglang.srt.server_args import ServerArgs + + assert not self._installed, "override_server_args already installed" + self._previous = self._context._server_args + self._previous_capture = self._context.flags.capture.enable_torch_compile + server_args = ServerArgs(model_path="dummy") + if self._fields: + server_args.override(source="test-override", **self._fields) + # The dummy boundary skips materialization, which would leave the + # strict mutation guard unarmed on the published object — mark it + # materialized so bare post-publish writes raise like they do on a + # fully resolved config. + object.__setattr__(server_args, "_declarations_materialized", True) + self._context.set_server_args(server_args) + self._installed = True + return server_args + + def restore(self) -> None: + """Reinstate the previously published config (or the empty slot).""" + if not self._installed: + return + self._installed = False + previous, self._previous = self._previous, None + if previous is None: + self._context._server_args = None + else: + self._context.set_server_args(previous) + # set_server_args reseeds the capture tier from the published object + # (and the empty-slot path does not touch it at all); the snapshot + # puts back the exact pre-install runtime state either way. + self._context.flags.capture.enable_torch_compile = self._previous_capture + + def __enter__(self) -> ServerArgs: + return self.install() + + def __exit__(self, *exc) -> None: + self.restore() + _PARALLEL = ParallelContext() _CONTEXT = RuntimeContext(parallel=_PARALLEL) diff --git a/python/sglang/srt/sampling/sampling_batch_info.py b/python/sglang/srt/sampling/sampling_batch_info.py index cfb22d841..680bac645 100644 --- a/python/sglang/srt/sampling/sampling_batch_info.py +++ b/python/sglang/srt/sampling/sampling_batch_info.py @@ -7,10 +7,10 @@ from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Tuple import torch import sglang.srt.sampling.penaltylib as penaltylib +from sglang.srt.runtime_context import get_server_args from sglang.srt.sampling.custom_logit_processor import CustomLogitProcessor from sglang.srt.sampling.penaltylib.repetition_penalty import apply_scaling_penalties from sglang.srt.sampling.sampling_params import TOP_K_ALL -from sglang.srt.server_args import get_global_server_args from sglang.srt.utils.common import is_pin_memory_available if TYPE_CHECKING: @@ -75,7 +75,7 @@ class SamplingBatchInfo: @classmethod def from_schedule_batch(cls, batch: ScheduleBatch, vocab_size: int): - global_server_args = get_global_server_args() + global_server_args = get_server_args() enable_deterministic = global_server_args.enable_deterministic_inference reqs = batch.reqs diff --git a/python/sglang/srt/speculative/dflash_info_v2.py b/python/sglang/srt/speculative/dflash_info_v2.py index fa4d35942..866ea093b 100644 --- a/python/sglang/srt/speculative/dflash_info_v2.py +++ b/python/sglang/srt/speculative/dflash_info_v2.py @@ -13,7 +13,7 @@ from sglang.srt.mem_cache.common import ( alloc_token_slots, get_last_loc, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func from sglang.srt.utils.common import is_pin_memory_available @@ -138,7 +138,7 @@ class DFlashDraftInputV2(SpecInput): cur_kv_lens_cpu_t = self._prepare_cur_kv_lens_cpu_buf[:bs] # For DFLASH, each decode step needs a fixed-size verify block. - block_size = int(get_global_server_args().speculative_num_draft_tokens) + block_size = int(get_server_args().speculative_num_draft_tokens) if block_size <= 0: raise ValueError( f"DFLASH invalid speculative_num_draft_tokens={block_size}." diff --git a/python/sglang/srt/speculative/dflash_utils.py b/python/sglang/srt/speculative/dflash_utils.py index 1c3c798ea..a1caaac2f 100644 --- a/python/sglang/srt/speculative/dflash_utils.py +++ b/python/sglang/srt/speculative/dflash_utils.py @@ -634,13 +634,13 @@ def compute_dflash_sampling_correct_drafts_and_bonus( ) if threshold_single is None: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - threshold_single = get_global_server_args().speculative_accept_threshold_single + threshold_single = get_server_args().speculative_accept_threshold_single if threshold_acc is None: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - threshold_acc = get_global_server_args().speculative_accept_threshold_acc + threshold_acc = get_server_args().speculative_accept_threshold_acc threshold_single = float(threshold_single) threshold_acc = max(float(threshold_acc), 1e-9) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 11406f6d8..ecdb47880 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -8,7 +8,7 @@ from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject from sglang.srt.environ import envs from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.speculative.spec_info import SpecInput, SpecInputType logger = logging.getLogger(__name__) @@ -206,7 +206,7 @@ class EagleDraftInput(SpecInput): topk_index=torch.empty((0, topk), device=device, dtype=torch.int64), draft_probs=( torch.empty((0, vocab_size), device=device, dtype=torch.float32) - if get_global_server_args().speculative_use_rejection_sampling + if get_server_args().speculative_use_rejection_sampling else None ), capture_hidden_mode=capture_hidden_mode, diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index 875fe9959..b397a6b1a 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -594,10 +594,10 @@ def eagle_sample( from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) + from sglang.srt.runtime_context import get_server_args from sglang.srt.sampling.penaltylib.repetition_penalty import ( apply_scaling_penalties, ) - from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_utils import ( SIMULATE_ACC_LEN, SIMULATE_ACC_TOKEN_MODE, @@ -684,9 +684,7 @@ def eagle_sample( chain_speculative_sampling_triton, ) - use_rejection_sampling = ( - get_global_server_args().speculative_use_rejection_sampling - ) + use_rejection_sampling = get_server_args().speculative_use_rejection_sampling # Apply temperature and get target probs expanded_temperature = torch.repeat_interleave( @@ -752,8 +750,8 @@ def eagle_sample( uniform_samples_for_final_sampling=coins_for_final_sampling, target_probs=target_probs, draft_probs=draft_probs, - threshold_single=get_global_server_args().speculative_accept_threshold_single, - threshold_acc=get_global_server_args().speculative_accept_threshold_acc, + threshold_single=get_server_args().speculative_accept_threshold_single, + threshold_acc=get_server_args().speculative_accept_threshold_acc, deterministic=True, ) @@ -849,9 +847,9 @@ def eagle_prepare_for_decode(batch: ScheduleBatch): # (get_alloc_reserve_per_decode) outgrows the req_to_token row: the write below # would OOB and free would leak KV. The row is widened to hold it in _init_pools # (PR #26972); fail here with a clear error, not on a later cryptic CUDA assert. - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args - if page_size > 1 and (get_global_server_args().speculative_eagle_topk or 1) > 1: + if page_size > 1 and (get_server_args().speculative_eagle_topk or 1) > 1: max_alloc_len = int(nxt_kv_lens_cpu.max()) row_width = batch.req_to_token_pool.req_to_token.shape[1] assert max_alloc_len <= row_width, ( diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index b4268a09b..8d75142ab 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -15,7 +15,7 @@ from sglang.srt.distributed.parallel_state import ( ) from sglang.srt.environ import envs from sglang.srt.managers.schedule_batch import set_mamba_track_indices_from_reqs -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.speculative.triton_ops.cache_locs import ( align_evict_mask_to_page_size as align_evict_mask_to_page_size, ) @@ -204,7 +204,7 @@ def record_stream_for_v2_verify(batch, verify_input, fwd_stream): def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool: if server_args is None: - server_args = get_global_server_args() + server_args = get_server_args() # STANDALONE drafts don't consume `spec_info.hidden_states` (vanilla LLM). # multi_layer_eagle and DFLASH don't relay hidden_states through FutureMap. @@ -622,7 +622,7 @@ def prepare_mamba_track_for_verify(batch: ScheduleBatch) -> None: tracking during TARGET_VERIFY; tracking is done in commit_mamba_states_after_verify instead. """ - if not get_global_server_args().enable_mamba_extra_buffer(): + if not get_server_args().enable_mamba_extra_buffer(): return set_mamba_track_indices_from_reqs(batch) batch.mamba_track_mask = None @@ -677,7 +677,7 @@ def commit_mamba_states_after_verify( # we need to update the mamba state for the request at the crossing point. seq_lens_pre_verify = batch.seq_lens seq_lens_post_verify = batch.seq_lens + accept_lens - mamba_track_interval = get_global_server_args().mamba_track_interval + mamba_track_interval = get_server_args().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/state_capturer/indexer_topk.py b/python/sglang/srt/state_capturer/indexer_topk.py index cd8ffcb9d..e37b5f243 100644 --- a/python/sglang/srt/state_capturer/indexer_topk.py +++ b/python/sglang/srt/state_capturer/indexer_topk.py @@ -20,7 +20,7 @@ class IndexerTopkCapturer(BaseTopkCapturer): max_running_requests: int, device: str, ): - from sglang.srt.server_args import get_global_server_args + from sglang.srt.runtime_context import get_server_args self.num_indexer_layers = num_indexer_layers self.index_topk = index_topk @@ -30,7 +30,7 @@ class IndexerTopkCapturer(BaseTopkCapturer): # DP-attention capture is per-rank-local: each rank writes [:local_batch, ...] # to its own device_cache, so the buffer only needs to fit one rank's batch. - server_args = get_global_server_args() + server_args = get_server_args() max_batch_size = max(server_args.chunked_prefill_size, max_running_requests) super().__init__( @@ -43,16 +43,16 @@ class IndexerTopkCapturer(BaseTopkCapturer): ) -_global_indexer_capturer: Optional[IndexerTopkCapturer] = None - - def get_global_indexer_capturer() -> Optional[IndexerTopkCapturer]: - return _global_indexer_capturer + from sglang.srt.runtime_context import get_resources + + return get_resources().indexer_capturer def set_global_indexer_capturer(capturer: Optional[IndexerTopkCapturer]): - global _global_indexer_capturer - _global_indexer_capturer = capturer + from sglang.srt.runtime_context import get_resources + + get_resources().indexer_capturer = capturer def maybe_capture_indexer_topk( diff --git a/python/sglang/srt/state_capturer/routed_experts.py b/python/sglang/srt/state_capturer/routed_experts.py index 24479870f..4e69f5dd3 100644 --- a/python/sglang/srt/state_capturer/routed_experts.py +++ b/python/sglang/srt/state_capturer/routed_experts.py @@ -12,8 +12,7 @@ from sglang.srt.layers.dp_attention import ( ) from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.state_capturer.base import BaseTopkCapturer @@ -58,7 +57,7 @@ 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_global_server_args() + 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 > @@ -127,16 +126,16 @@ class RoutedExpertsCapturer(BaseTopkCapturer): ] -_global_expert_capturer: Optional[RoutedExpertsCapturer] = None - - def get_global_experts_capturer() -> Optional[RoutedExpertsCapturer]: - return _global_expert_capturer + from sglang.srt.runtime_context import get_resources + + return get_resources().experts_capturer def set_global_experts_capturer(capturer: Optional[RoutedExpertsCapturer]): - global _global_expert_capturer - _global_expert_capturer = capturer + from sglang.srt.runtime_context import get_resources + + get_resources().experts_capturer = capturer def extract_routed_experts_from_meta_info(data): diff --git a/python/sglang/srt/utils/cuda_ipc_transport_utils.py b/python/sglang/srt/utils/cuda_ipc_transport_utils.py index b062ddd7b..e8df3026b 100644 --- a/python/sglang/srt/utils/cuda_ipc_transport_utils.py +++ b/python/sglang/srt/utils/cuda_ipc_transport_utils.py @@ -9,7 +9,7 @@ import numpy as np import torch from sglang.srt.environ import envs -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils.stale_shm_cleanup import make_shm_name logger = logging.getLogger(__name__) @@ -104,10 +104,10 @@ class MmItemMemoryChunk: def try_to_recycle(self) -> bool: try: - tp_num = get_global_server_args().tp_size + tp_num = get_server_args().tp_size except Exception: logger.info( - "get_global_server_args has not been inited , skip this turn 's recycle" + "server_args has not been published yet, skip this turn's recycle" ) return False diff --git a/python/sglang/srt/utils/profile_utils.py b/python/sglang/srt/utils/profile_utils.py index 84e8db0b9..626f6614a 100644 --- a/python/sglang/srt/utils/profile_utils.py +++ b/python/sglang/srt/utils/profile_utils.py @@ -12,7 +12,7 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.managers.io_struct import ProfileReqOutput from sglang.srt.model_executor.forward_batch_info import ForwardMode -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import is_npu from sglang.srt.utils.torch_npu_patch_utils import apply_torch_npu_patches @@ -61,7 +61,7 @@ class ProfileManager: ) self.ps = ps self.cpu_group = cpu_group - self.first_rank_in_node = ps.gpu_id == get_global_server_args().base_gpu_id + self.first_rank_in_node = ps.gpu_id == get_server_args().base_gpu_id self.profiler_kwargs = None self.profiler = None diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py index c5f357007..a012dfeab 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py @@ -19,12 +19,9 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import set_global_server_args_for_scheduler +from sglang.srt.runtime_context import get_context, get_parallel from sglang.srt.speculative.spec_info import SpeculativeAlgorithm -from ..mock_server_args import make_mock_server_args - # Unit tests run without distributed initialization. Backends that size buffers by # attention tensor-parallel degree should see the single-rank default. _parallel_override = get_parallel().override(attn_tp_size=1) @@ -338,7 +335,7 @@ class MockModelRunner(ModelRunner): or case.forward_mode.is_draft_extend_v2() else 0 ) - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -363,7 +360,6 @@ class MockModelRunner(ModelRunner): is_embedding=False, kv_cache_dtype="auto", max_running_requests=None, - model_path=None, pp_size=1, revision=None, speculative_algorithm=None, @@ -374,7 +370,7 @@ class MockModelRunner(ModelRunner): triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, ) - set_global_server_args_for_scheduler(self.server_args) + self.server_args = self._server_args_override.install() self.req_to_token_pool = ReqToTokenPool( size=pool_batch_size, max_context_len=max_context_len, diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py index bcad2b111..d056ff6c8 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py @@ -16,10 +16,8 @@ from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import set_global_server_args_for_scheduler +from sglang.srt.runtime_context import get_context, get_parallel -from ..mock_server_args import make_mock_server_args from .dense_attention import ( DEFAULT_DEVICE, DEFAULT_HEAD_DIM, @@ -310,7 +308,7 @@ class DSAMockModelRunner(ModelRunner): self._kernel_warmed_up = True self.dp_size = 1 self.pp_size = 1 - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -341,7 +339,6 @@ class DSAMockModelRunner(ModelRunner): kv_cache_dtype="auto", max_running_requests=None, mem_fraction_static=0.8, - model_path=None, pp_size=1, revision=None, speculative_algorithm=None, @@ -352,7 +349,7 @@ class DSAMockModelRunner(ModelRunner): triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, ) - set_global_server_args_for_scheduler(self.server_args) + self.server_args = self._server_args_override.install() self.req_to_token_pool = ReqToTokenPool( size=pool_batch_size, max_context_len=max_context_len, diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py index d295b1489..5fca0dc96 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py @@ -34,10 +34,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_context import ForwardContext, forward_context -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import set_global_server_args_for_scheduler - -from ..mock_server_args import make_mock_server_args +from sglang.srt.runtime_context import get_context, get_parallel # DSV4 backend pre-resolves attention TP at construction; pin to single-rank. _parallel_override = get_parallel().override( @@ -334,7 +331,7 @@ class MockDSV4ModelRunner: self.tp_size = 1 self.dp_size = 1 self.pp_size = 1 - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -358,7 +355,6 @@ class MockDSV4ModelRunner: is_embedding=False, kv_cache_dtype="auto", max_running_requests=None, - model_path=None, pp_size=1, revision=None, speculative_algorithm=None, @@ -369,7 +365,7 @@ class MockDSV4ModelRunner: device=device, mem_fraction_static=0.8, ) - set_global_server_args_for_scheduler(self.server_args) + self.server_args = self._server_args_override.install() self.req_to_token_pool = ReqToTokenPool( size=pool_batch_size, max_context_len=max_context_len, diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py index 355e3f30d..93a09a0e5 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py @@ -18,10 +18,8 @@ from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import set_global_server_args_for_scheduler +from sglang.srt.runtime_context import get_context, get_parallel -from ..mock_server_args import make_mock_server_args from .dense_attention import ( DEFAULT_DEVICE, DEFAULT_DTYPE, @@ -325,7 +323,7 @@ class DualChunkMockModelRunner(ModelRunner): self._kernel_warmed_up = True self.dp_size = 1 self.pp_size = 1 - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -352,7 +350,7 @@ class DualChunkMockModelRunner(ModelRunner): triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, ) - set_global_server_args_for_scheduler(self.server_args) + self.server_args = self._server_args_override.install() self.req_to_token_pool = ReqToTokenPool( size=pool_batch_size, max_context_len=max_context_len, diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py index fb6bf4606..09d13572a 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py @@ -29,9 +29,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel - -from ..mock_server_args import make_mock_server_args +from sglang.srt.runtime_context import get_context, get_parallel _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() @@ -221,7 +219,7 @@ class MockGDNModelRunner(ModelRunner): or case.forward_mode.is_draft_extend_v2() else 0 ) - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -243,9 +241,7 @@ class MockGDNModelRunner(ModelRunner): linear_attn_backend="triton", linear_attn_decode_backend=None, linear_attn_prefill_backend=None, - mamba_cache_chunk_size=64, max_running_requests=None, - model_path=None, revision=None, speculative_algorithm=None, speculative_eagle_topk=1 if case.forward_mode.is_target_verify() else 0, @@ -253,7 +249,11 @@ class MockGDNModelRunner(ModelRunner): speculative_num_steps=max(0, speculative_num_draft_tokens - 1), triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, + # Pin the lazy mamba_cache_chunk_size property cache: production + # derives it from hf_config + page_size, which needs a real model. + _mamba_cache_chunk_size=64, ) + self.server_args = self._server_args_override.install() cache_shape = Mamba2StateShape.create( tp_world_size=1, intermediate_size=case.num_v_heads * head_v_dim, diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py index 9213e5f2f..89453b45c 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py @@ -29,9 +29,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel - -from ..mock_server_args import make_mock_server_args +from sglang.srt.runtime_context import get_context, get_parallel _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() @@ -227,7 +225,7 @@ class MockKDAModelRunner(ModelRunner): or case.forward_mode.is_draft_extend_v2() else 0 ) - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -249,9 +247,7 @@ class MockKDAModelRunner(ModelRunner): linear_attn_backend="triton", linear_attn_decode_backend=None, linear_attn_prefill_backend=None, - mamba_cache_chunk_size=64, max_running_requests=None, - model_path=None, revision=None, speculative_algorithm=None, speculative_eagle_topk=1 if case.forward_mode.is_target_verify() else 0, @@ -259,7 +255,11 @@ class MockKDAModelRunner(ModelRunner): speculative_num_steps=max(0, speculative_num_draft_tokens - 1), triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, + # Pin the lazy mamba_cache_chunk_size property cache: production + # derives it from hf_config + page_size, which needs a real model. + _mamba_cache_chunk_size=64, ) + self.server_args = self._server_args_override.install() # KDA uses the KimiLinear cache layout (conv_kernel-1, conv_dim) and a # temporal state of (num_heads, head_dim, head_dim). The KDA backend's # forward_extend splits conv by [q_dim, k_dim, v_dim] along the conv_dim diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py index 43b72fb4b..7eab8dcf4 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py @@ -28,9 +28,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel - -from ..mock_server_args import make_mock_server_args +from sglang.srt.runtime_context import get_context, get_parallel _parallel_override = get_parallel().override(attn_tp_size=1, attn_tp_rank=0) _parallel_override.__enter__() @@ -235,7 +233,7 @@ class MockLightningModelRunner(ModelRunner): or case.forward_mode.is_draft_extend_v2() else 0 ) - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -258,9 +256,7 @@ class MockLightningModelRunner(ModelRunner): linear_attn_backend="triton", linear_attn_decode_backend=None, linear_attn_prefill_backend=None, - mamba_cache_chunk_size=64, max_running_requests=None, - model_path=None, revision=None, speculative_algorithm=None, speculative_eagle_topk=1 if case.forward_mode.is_target_verify() else 0, @@ -268,7 +264,11 @@ class MockLightningModelRunner(ModelRunner): speculative_num_steps=max(0, speculative_num_draft_tokens - 1), triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, + # Pin the lazy mamba_cache_chunk_size property cache: production + # derives it from hf_config + page_size, which needs a real model. + _mamba_cache_chunk_size=64, ) + self.server_args = self._server_args_override.install() # Lightning seg_la temporal state is [num_heads, head_dim, head_dim]; Bailing's # mamba2_cache_params sets intermediate_size=0, n_groups=0, conv_kernel=1 # because seg_la does not use a conv state (the conv shape collapses to (0, 0)). diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index 09fadaacf..18d887338 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -7,7 +7,7 @@ from torch import nn # Patch TP world size / rank before importing modules that read them at __init__. import sglang.srt.layers.linear as _linear_mod -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_context, get_parallel _parallel_override = get_parallel().override( tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0 @@ -50,8 +50,6 @@ from sglang.srt.model_executor.forward_context import ( # noqa: E402 ) from sglang.srt.model_executor.model_runner import ModelRunner # noqa: E402 -from ..mock_server_args import make_mock_server_args - # Tiny dims chosen to be the minimum that satisfies MambaMixer2's TP/chunk asserts: # - num_heads % tp_size == 0 (tp_size=1) # - intermediate_size = num_heads * head_dim @@ -330,7 +328,7 @@ class MockMamba2ModelRunner(ModelRunner): ) else: speculative_num_draft_tokens = 0 - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -351,7 +349,7 @@ class MockMamba2ModelRunner(ModelRunner): enable_mis=False, # `RowParallelLinear.forward` (called by the production # `MambaMixer2.out_proj`) consults - # `get_global_server_args().enable_symm_mem` to decide whether + # `get_server_args().enable_symm_mem` to decide whether # to wrap allocations in a symmetric-memory context. With # `world_size=1` the wrapper short-circuits, but the # attribute read still happens, so it must exist on the mock @@ -367,9 +365,7 @@ class MockMamba2ModelRunner(ModelRunner): # `MambaMixer2.forward_decode` calls into. Set it explicitly so # the DECODE fixture path becomes reachable. mamba_backend="triton", - mamba_cache_chunk_size=64, max_running_requests=None, - model_path=None, revision=None, speculative_algorithm=None, speculative_eagle_topk=0, @@ -377,17 +373,14 @@ class MockMamba2ModelRunner(ModelRunner): speculative_num_steps=0, triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, + # Pin the lazy mamba_cache_chunk_size property cache: production + # derives it from hf_config + page_size, which needs a real model. + _mamba_cache_chunk_size=64, ) - # Install this fixture's `server_args` as the global so that - # `is_symmetric_memory_enabled()` (called from - # `RowParallelLinear.forward`) reads our `enable_symm_mem=False` - # value. Without this, a previous test in the discover sweep - # whose fixture *did* call `set_global_server_args_for_scheduler` - # would leave a SimpleNamespace without `enable_symm_mem` as the - # global, and the mamba2 forward would AttributeError. - from sglang.srt.server_args import set_global_server_args_for_scheduler - - set_global_server_args_for_scheduler(self.server_args) + # install() publishes this fixture's config, so production reads + # like `is_symmetric_memory_enabled()` (RowParallelLinear.forward) + # see our `enable_symm_mem=False` for the fixture's lifetime. + self.server_args = self._server_args_override.install() # Install the selective-state-update backend that # `MambaMixer2.forward_decode` requires. In production the diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py index 61852dcc9..7eeff9477 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py @@ -23,10 +23,7 @@ from sglang.srt.model_executor.forward_context import ( ) from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import set_global_server_args_for_scheduler - -from ..mock_server_args import make_mock_server_args +from sglang.srt.runtime_context import get_context, get_parallel _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() @@ -242,7 +239,7 @@ class MockMLAModelRunner(ModelRunner): or case.forward_mode.is_draft_extend_v2() else 0 ) - self.server_args = make_mock_server_args( + self._server_args_override = get_context().override_server_args( attention_backend=case.backend, chunked_prefill_size=-1, cuda_graph_config=CudaGraphConfig( @@ -270,7 +267,6 @@ class MockMLAModelRunner(ModelRunner): is_embedding=False, kv_cache_dtype="fp8_e4m3" if fp8_kv_cache else "auto", max_running_requests=None, - model_path=None, pp_size=1, revision=None, speculative_algorithm=None, @@ -281,7 +277,7 @@ class MockMLAModelRunner(ModelRunner): triton_attention_num_kv_splits=8, triton_attention_split_tile_size=None, ) - set_global_server_args_for_scheduler(self.server_args) + self.server_args = self._server_args_override.install() self.req_to_token_pool = ReqToTokenPool( size=pool_batch_size, max_context_len=max_context_len, diff --git a/python/sglang/test/kits/attention_unittest/mock_server_args.py b/python/sglang/test/kits/attention_unittest/mock_server_args.py deleted file mode 100644 index a347a4d15..000000000 --- a/python/sglang/test/kits/attention_unittest/mock_server_args.py +++ /dev/null @@ -1,62 +0,0 @@ -"""Mock `ServerArgs` factory for attention-backend unit tests. - -Production attention backends read many `ServerArgs` attributes and call -several `ServerArgs` methods at backend construction time. The set grows -monotonically: new attention features add new attributes/methods to -`ServerArgs`, and a fixture that mocks `server_args` as a manually- -populated `SimpleNamespace` will silently miss the new field and fail -with `AttributeError` the next time a backend looks it up. - -`make_mock_server_args` sidesteps this by instantiating a real -`ServerArgs` (the dataclass) with all defaults from the dataclass -definition, then overlaying the caller's explicit overrides. New -`ServerArgs` attributes are picked up automatically with their default -values; methods like `enable_mamba_extra_buffer()` work because the -object is a real `ServerArgs` instance, so methods are bound correctly. - -`__post_init__` is intentionally bypassed (via `object.__new__`) so -fixture callers don't have to supply a real `model_path`; the -validation it performs is irrelevant for module-level attention tests. -""" - -import dataclasses - -from sglang.srt.model_executor.cuda_graph_config import default_cuda_graph_config -from sglang.srt.server_args import ServerArgs - - -def make_mock_server_args(**overrides) -> ServerArgs: - """Return a `ServerArgs` instance with all defaults pre-populated. - - The instance is built by `object.__new__(ServerArgs)` so `__post_init__` - does not run — fixture callers do not need to supply a valid - `model_path` or other required-field values. - - Any field with a `default` or `default_factory` in the dataclass - definition is set automatically. Caller-supplied `overrides` replace - those defaults; unknown keys are also stored (matching `SimpleNamespace` - semantics) so fixtures can attach test-only attributes when needed. - - If an override name corresponds to a read-only `@property` on - `ServerArgs`, the value is stored under `_` instead — many - `ServerArgs` properties cache through `_` and return it when - set, so fixture callers can keep using the public name and let this - helper translate. - """ - sa = object.__new__(ServerArgs) - for f in dataclasses.fields(ServerArgs): - if f.default is not dataclasses.MISSING: - setattr(sa, f.name, f.default) - elif f.default_factory is not dataclasses.MISSING: - setattr(sa, f.name, f.default_factory()) - for k, v in overrides.items(): - cls_attr = getattr(type(sa), k, None) - if isinstance(cls_attr, property): - setattr(sa, f"_{k}", v) - else: - setattr(sa, k, v) - if sa.cuda_graph_config is None: - sa.cuda_graph_config = default_cuda_graph_config() - if not hasattr(sa, "_cuda_graph_config_locked"): - sa._cuda_graph_config_locked = set() - return sa diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py index ff3a57c94..1b724c4eb 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py @@ -17,7 +17,6 @@ from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.input_buffers import _forward_input_buffer_pool from sglang.srt.model_executor.runner import set_global_graph_memory_pool from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import set_global_server_args_for_scheduler from sglang.srt.speculative.draft_utils import DraftBackendFactory from sglang.srt.speculative.eagle_draft_cuda_graph_runner import ( EAGLEDraftCudaGraphRunner, @@ -327,8 +326,7 @@ def _configure_runner_for_eagle_draft( "torch_compile_max_bs": 0, "use_mla_backend": runner.use_mla_backend, } - for key, value in updates.items(): - setattr(server_args, key, value) + server_args.override(source="attention-unittest-eagle-draft", **updates) runner.spec_algorithm = SpeculativeAlgorithm.EAGLE runner.is_draft_worker = True @@ -339,7 +337,6 @@ def _configure_runner_for_eagle_draft( runner.model_config.dtype = runner.dtype runner.model_config.vocab_size = settings.vocab_size runner.model_config.hf_config.vocab_size = settings.vocab_size - set_global_server_args_for_scheduler(server_args) def _build_eagle_draft_fixture( diff --git a/test/manual/attention/test_trtllm_mla_backend.py b/test/manual/attention/test_trtllm_mla_backend.py index 7cd5abb4f..f5aec11f7 100755 --- a/test/manual/attention/test_trtllm_mla_backend.py +++ b/test/manual/attention/test_trtllm_mla_backend.py @@ -5,7 +5,7 @@ from types import SimpleNamespace import numpy as np import torch -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_server_args _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() @@ -26,7 +26,6 @@ from sglang.srt.model_executor.forward_context import ( ) from sglang.srt.server_args import ( ServerArgs, - get_global_server_args, set_global_server_args_for_scheduler, ) from sglang.srt.utils import is_flashinfer_available @@ -222,7 +221,7 @@ class MockModelRunner: self.page_size = config["page_size"] # Server args stub - needed by attention backends - self.server_args = get_global_server_args() + self.server_args = get_server_args() # Model-config stub with MLA attributes self.model_config = type( diff --git a/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py b/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py index f8b1d4726..a00851797 100644 --- a/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py +++ b/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py @@ -22,7 +22,6 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( HybridLinearAttnBackend, ) from sglang.srt.model_executor.forward_batch_info import ForwardMode -from sglang.srt.server_args import set_global_server_args_for_scheduler from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.attention_unittest.attention_methods.mla_attention import ( DEFAULT_KV_LORA_RANK, @@ -46,9 +45,13 @@ class _ChunkKVMLARunner(MockMLAModelRunner): def __init__(self, **kwargs): super().__init__(**kwargs) - self.server_args.disable_chunked_prefix_cache = False - self.server_args.flashinfer_mla_disable_ragged = False - set_global_server_args_for_scheduler(self.server_args) + # The fixture's config is already published; adjust it through the + # audited entry point (bare writes raise under the strict guard). + self.server_args.override( + source="attention-unittest", + disable_chunked_prefix_cache=False, + flashinfer_mla_disable_ragged=False, + ) def _make_case() -> MLAAttentionCase: diff --git a/test/registered/ops/test_aiter_allreduce_fusion_amd.py b/test/registered/ops/test_aiter_allreduce_fusion_amd.py index 94542c652..5a0f1fb41 100755 --- a/test/registered/ops/test_aiter_allreduce_fusion_amd.py +++ b/test/registered/ops/test_aiter_allreduce_fusion_amd.py @@ -418,7 +418,7 @@ class TestAiterAllreduceFusionGate(CustomTestCase): ) ) stack.enter_context( - mock.patch.object(comm, "get_global_server_args", lambda: server_args) + mock.patch.object(comm, "get_server_args", lambda: server_args) ) from sglang.srt.runtime_context import get_flags diff --git a/test/registered/rl/test_fp32_lm_head.py b/test/registered/rl/test_fp32_lm_head.py index d952974f6..eb2261dd0 100644 --- a/test/registered/rl/test_fp32_lm_head.py +++ b/test/registered/rl/test_fp32_lm_head.py @@ -7,9 +7,9 @@ import torch.nn as nn import torch.nn.functional as F from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.runtime_context import get_server_args from sglang.srt.server_args import ( ServerArgs, - get_global_server_args, set_global_server_args_for_scheduler, ) from sglang.srt.utils import get_device @@ -44,7 +44,7 @@ class TestLMHeadFP32(unittest.TestCase): def _make_logprocessor(self, vocab_size, enable_fp32): set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) - get_global_server_args().enable_fp32_lm_head = enable_fp32 + get_server_args().enable_fp32_lm_head = enable_fp32 cfg = SimpleNamespace(vocab_size=vocab_size, final_logit_softcapping=None) return LogitsProcessor(cfg, skip_all_gather=True, logit_scale=None) diff --git a/test/registered/unit/batch_overlap/test_tbo_cuda_graph_num_token_device.py b/test/registered/unit/batch_overlap/test_tbo_cuda_graph_num_token_device.py index 0e01a46cd..d6ee5eeff 100644 --- a/test/registered/unit/batch_overlap/test_tbo_cuda_graph_num_token_device.py +++ b/test/registered/unit/batch_overlap/test_tbo_cuda_graph_num_token_device.py @@ -4,7 +4,7 @@ ``ForwardBatch.num_token_non_padded`` is a scalar tensor on the model device (see ``ForwardBatch.compute``, which does ``.to(device, ...)``). The eager TBO split path already honors this -- ``compute_tbo_children_num_token_non_padded_raw`` -moves the tensor to ``get_global_server_args().device`` -- but +moves the tensor to ``get_server_args().device`` -- but ``TboCudaGraphRunnerPlugin`` preallocated its persistent buffer with a bare ``torch.zeros((2,), dtype=torch.int32)``, leaving it on CPU. @@ -39,7 +39,7 @@ class TestTboCudaGraphNumTokenDevice(CustomTestCase): # Use 'meta' so the configured device differs from the implicit CPU # default; a bare torch.zeros() would leave the buffer on CPU and fail. fake_args = SimpleNamespace(device="meta") - with patch.object(tbo, "get_global_server_args", lambda: fake_args): + with patch.object(tbo, "get_server_args", lambda: fake_args): plugin = TboCudaGraphRunnerPlugin() buf = plugin._tbo_children_num_token_non_padded @@ -51,7 +51,7 @@ class TestTboCudaGraphNumTokenDevice(CustomTestCase): # Both the preallocated cuda-graph buffer and the eager split tensor must # land on the same (model) device, matching ForwardBatch's contract. fake_args = SimpleNamespace(device="meta") - with patch.object(tbo, "get_global_server_args", lambda: fake_args): + with patch.object(tbo, "get_server_args", lambda: fake_args): eager = ( TboForwardBatchPreparer.compute_tbo_children_num_token_non_padded_raw( tbo_split_token_index=3, num_token_non_padded=8 @@ -68,7 +68,7 @@ class TestTboCudaGraphNumTokenDevice(CustomTestCase): # value_a = min(split, n); value_b = max(0, n - split). Computed on CPU # so the values are materializable. fake_args = SimpleNamespace(device="cpu") - with patch.object(tbo, "get_global_server_args", lambda: fake_args): + with patch.object(tbo, "get_server_args", lambda: fake_args): eager = ( TboForwardBatchPreparer.compute_tbo_children_num_token_non_padded_raw( tbo_split_token_index=3, num_token_non_padded=8 diff --git a/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py b/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py index a3cd88203..81ecc00e8 100644 --- a/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py +++ b/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py @@ -39,7 +39,7 @@ def _make_target_verify_batch(bs: int) -> ForwardBatch: def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch: fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3") with get_parallel().override(attn_tp_size=1), patch.object( - tbo, "get_global_server_args", lambda: fake_args + tbo, "get_server_args", lambda: fake_args ): return TboForwardBatchPreparer.filter_batch( batch, diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index 8c84b9069..0920f7a5c 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -1257,7 +1257,7 @@ class TestMlxOverlapScheduler(unittest.TestCase): logits_output = SimpleNamespace(customized_info=None) original_release = batch_result_processor_module.release_kv_cache original_get_indexer = batch_result_processor_module.get_global_indexer_capturer - original_get_server_args = batch_result_processor_module.get_global_server_args + original_get_server_args = batch_result_processor_module.get_server_args def fake_release_kv_cache(release_req, tree_cache, is_insert=False): events.append(("release", release_req.rid)) @@ -1265,7 +1265,7 @@ class TestMlxOverlapScheduler(unittest.TestCase): batch_result_processor_module.release_kv_cache = fake_release_kv_cache batch_result_processor_module.get_global_indexer_capturer = lambda: None - batch_result_processor_module.get_global_server_args = lambda: SimpleNamespace( + batch_result_processor_module.get_server_args = lambda: SimpleNamespace( enable_mamba_extra_buffer_lazy=lambda: False ) try: @@ -1279,9 +1279,7 @@ class TestMlxOverlapScheduler(unittest.TestCase): batch_result_processor_module.get_global_indexer_capturer = ( original_get_indexer ) - batch_result_processor_module.get_global_server_args = ( - original_get_server_args - ) + batch_result_processor_module.get_server_args = original_get_server_args self.assertEqual( events, diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 3c0e7156a..77b2f5716 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -56,10 +56,10 @@ from sglang.srt.mem_cache.unified_radix_cache import ( UnifiedRadixCache, UnifiedTreeNode, ) +from sglang.srt.runtime_context import get_server_args from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ( ServerArgs, - get_global_server_args, set_global_server_args_for_scheduler, ) from sglang.srt.utils import get_device @@ -942,13 +942,13 @@ class UnifiedRadixCacheSuite: req.mamba_last_track_seqlen = kv_len req.reasoning_tokens = 1 - get_global_server_args().strip_thinking_cache = True + get_server_args().strip_thinking_cache = True try: avail_before = allocator.available_size() cache.cache_finished_req(req, is_insert=True) start_p, end_p = req.pop_overallocated_kv_cache() finally: - get_global_server_args().strip_thinking_cache = False + get_server_args().strip_thinking_cache = False if ps > 1: start_p = ((start_p + ps - 1) // ps) * ps if start_p < end_p: @@ -3527,7 +3527,7 @@ class UnifiedRadixCacheSuite: if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1: self.skipTest("requires page_size=1 Full+Mamba") cache, allocator, req_to_token_pool = build_fixture(self.cfg) - chunk_size = get_global_server_args().mamba_cache_chunk_size + chunk_size = get_server_args().mamba_cache_chunk_size tokens = self._make_seq(1, chunk_size + 1) self._insert(cache, allocator, req_to_token_pool, tokens) leaf = cache.match_prefix( diff --git a/test/registered/unit/observability/test_trace.py b/test/registered/unit/observability/test_trace.py index 547beab65..a9bf9f078 100644 --- a/test/registered/unit/observability/test_trace.py +++ b/test/registered/unit/observability/test_trace.py @@ -21,6 +21,7 @@ from sglang.srt.observability.trace import ( TraceThreadContext, TraceThreadInfo, extract_trace_headers, + get_global_trace_level, get_global_tracing_enabled, process_tracing_init, set_global_trace_level, @@ -51,19 +52,29 @@ class TestTraceFunctions(unittest.TestCase): self.assertEqual(extract_trace_headers({}), {}) def test_set_global_trace_level(self): - orig = mod.global_trace_level - set_global_trace_level(5) - self.assertEqual(mod.global_trace_level, 5) - mod.global_trace_level = orig + from sglang.srt.runtime_context import get_resources + + orig = get_resources().trace_level + try: + set_global_trace_level(5) + self.assertEqual(get_global_trace_level(), 5) + finally: + get_resources().trace_level = orig def test_global_trace_level_env_var(self): - import importlib + # The level lives on ctx.resources and is seeded lazily from the env + # on first read after a reset (no module reload involved). + from sglang.srt.runtime_context import get_resources - with patch.dict(os.environ, {"SGLANG_TRACE_LEVEL": "2"}): - importlib.reload(mod) - self.assertEqual(mod.global_trace_level, 2) - importlib.reload(mod) # restore default (SGLANG_TRACE_LEVEL unset → 3) - self.assertEqual(mod.global_trace_level, 3) + orig = get_resources().trace_level + try: + with patch.dict(os.environ, {"SGLANG_TRACE_LEVEL": "2"}): + get_resources().trace_level = None + self.assertEqual(get_global_trace_level(), 2) + get_resources().trace_level = None # SGLANG_TRACE_LEVEL unset → 3 + self.assertEqual(get_global_trace_level(), 3) + finally: + get_resources().trace_level = orig def test_get_global_tracing_enabled(self): self.assertEqual(get_global_tracing_enabled(), mod.opentelemetry_initialized) @@ -244,7 +255,9 @@ class TestTraceReqContextEnabled(unittest.TestCase): self.orig_initialized = mod.opentelemetry_initialized self.orig_tracer = mod.tracer self.orig_threads = mod.threads_info.copy() - self.orig_level = mod.global_trace_level + from sglang.srt.runtime_context import get_resources + + self.orig_level = get_resources().trace_level # Reset OTel global TracerProvider so set_tracer_provider works each test otel_trace._TRACER_PROVIDER_SET_ONCE._done = False @@ -254,14 +267,16 @@ class TestTraceReqContextEnabled(unittest.TestCase): otel_trace.set_tracer_provider(self.provider) mod.opentelemetry_initialized = True mod.tracer = otel_trace.get_tracer("test") - mod.global_trace_level = 3 + set_global_trace_level(3) def tearDown(self): mod.opentelemetry_initialized = self.orig_initialized mod.tracer = self.orig_tracer mod.threads_info.clear() mod.threads_info.update(self.orig_threads) - mod.global_trace_level = self.orig_level + from sglang.srt.runtime_context import get_resources + + get_resources().trace_level = self.orig_level def test_trace_set_thread_info(self): trace_set_thread_info("scheduler", tp_rank=0, dp_rank=0) diff --git a/test/registered/unit/parser/test_template_manager.py b/test/registered/unit/parser/test_template_manager.py index 32f147008..58425807b 100644 --- a/test/registered/unit/parser/test_template_manager.py +++ b/test/registered/unit/parser/test_template_manager.py @@ -602,10 +602,18 @@ class TestResolveAutoParsers(unittest.TestCase): qwen3_template = "{% set enable_thinking = enable_thinking if enable_thinking is defined else true %}" + class _Args(SimpleNamespace): + # Write-through override, per the runtime-context testing idiom: + # production adjusts parsers through override(source, ...), so the + # stand-in needs the method (a bare SimpleNamespace would raise). + def override(self, source, **fields): + for key, value in fields.items(): + setattr(self, key, value) + def _make_server_args( self, reasoning_parser=None, tool_call_parser=None, chat_template=None ): - return SimpleNamespace( + return self._Args( reasoning_parser=reasoning_parser, tool_call_parser=tool_call_parser, model_path="Qwen/Qwen3-0.6B", @@ -650,12 +658,8 @@ class TestResolveAutoParsers(unittest.TestCase): self.assertEqual(args.tool_call_parser, "qwen") def test_nonexistent_model_disables_both_parsers(self): - args = SimpleNamespace( - reasoning_parser="auto", - tool_call_parser="auto", - model_path="nonexistent/model-does-not-exist-xyz", - trust_remote_code=False, - ) + args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") + args.model_path = "nonexistent/model-does-not-exist-xyz" with _patch_hf_transformers_utils( Mock(side_effect=RuntimeError("tokenizer unavailable")), Mock(side_effect=RuntimeError("config unavailable")), diff --git a/test/registered/unit/sampling/test_sampling_batch_info.py b/test/registered/unit/sampling/test_sampling_batch_info.py index 023bf93bb..46ed180b9 100644 --- a/test/registered/unit/sampling/test_sampling_batch_info.py +++ b/test/registered/unit/sampling/test_sampling_batch_info.py @@ -452,7 +452,7 @@ class TestFromScheduleBatch(CustomTestCase): req.tokenizer.eos_token_id = eos_id return req - @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + @patch("sglang.srt.sampling.sampling_batch_info.get_server_args") def test_basic_construction(self, mock_server_args): """Test that from_schedule_batch correctly extracts sampling params from requests.""" mock_server_args.return_value.enable_deterministic_inference = False @@ -469,7 +469,7 @@ class TestFromScheduleBatch(CustomTestCase): self.assertAlmostEqual(info.top_ps[0].item(), 0.9, places=5) self.assertEqual(info.top_ks[0].item(), 50) - @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + @patch("sglang.srt.sampling.sampling_batch_info.get_server_args") def test_greedy_detection(self, mock_server_args): """Test that top_k=1 sets is_all_greedy=True.""" mock_server_args.return_value.enable_deterministic_inference = False @@ -482,7 +482,7 @@ class TestFromScheduleBatch(CustomTestCase): info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) self.assertTrue(info.is_all_greedy) - @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + @patch("sglang.srt.sampling.sampling_batch_info.get_server_args") def test_logit_bias_construction(self, mock_server_args): """Test that logit_bias dict is converted to a tensor with correct values.""" mock_server_args.return_value.enable_deterministic_inference = False @@ -498,7 +498,7 @@ class TestFromScheduleBatch(CustomTestCase): self.assertAlmostEqual(info.logit_bias[0, 10].item(), -1.0) self.assertAlmostEqual(info.logit_bias[0, 0].item(), 0.0) - @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + @patch("sglang.srt.sampling.sampling_batch_info.get_server_args") def test_deterministic_seed(self, mock_server_args): """Test that explicit seed=123 is kept and missing seed defaults to 42.""" mock_server_args.return_value.enable_deterministic_inference = True @@ -513,7 +513,7 @@ class TestFromScheduleBatch(CustomTestCase): self.assertEqual(info.sampling_seed[0].item(), 123) self.assertEqual(info.sampling_seed[1].item(), 42) # default - @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + @patch("sglang.srt.sampling.sampling_batch_info.get_server_args") def test_from_schedule_batch_sampling_flags(self, mock_server_args): """Test that sampling flags (need_top_p/top_k/min_p) are set correctly.""" mock_server_args.return_value.enable_deterministic_inference = False @@ -529,7 +529,7 @@ class TestFromScheduleBatch(CustomTestCase): self.assertTrue(info.need_min_p_sampling) # 0.1 > 0 self.assertFalse(info.is_all_greedy) # top_k=50 > 1 - @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + @patch("sglang.srt.sampling.sampling_batch_info.get_server_args") def test_no_logit_bias_when_all_none(self, mock_server_args): """Test that logit_bias stays None when no request has logit_bias set.""" mock_server_args.return_value.enable_deterministic_inference = False @@ -542,7 +542,7 @@ class TestFromScheduleBatch(CustomTestCase): info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) self.assertIsNone(info.logit_bias) - @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + @patch("sglang.srt.sampling.sampling_batch_info.get_server_args") def test_custom_logit_processor_merging(self, mock_server_args): """Test deserialization and merging of custom logit processors.""" from sglang.srt.sampling.custom_logit_processor import ( diff --git a/test/registered/unit/test_legacy_global_ratchet.py b/test/registered/unit/test_legacy_global_ratchet.py index 684a822bf..a2c40b888 100644 --- a/test/registered/unit/test_legacy_global_ratchet.py +++ b/test/registered/unit/test_legacy_global_ratchet.py @@ -25,7 +25,9 @@ _SRT_ROOT = Path(next(iter(sglang.srt.__path__))) # Baselines counted over python/sglang/srt/**/*.py, including each function's # own def line. Ratchet: decrease-only. _RATCHETS = [ - ("get_global_server_args", r"\bget_global_server_args\s*\(", 279), + # Down to the shim definition itself; every call-site now goes through + # runtime_context.get_server_args(). + ("get_global_server_args", r"\bget_global_server_args\s*\(", 1), ( "set_global_server_args_for_*", r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(", diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 3567b5c79..cdef84111 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -312,12 +312,11 @@ class TestGoldenModelOverrides(_IsolatedPublish): def _publish(self, server_args): from sglang.srt.server_args import ( - get_global_server_args, set_global_server_args_for_scheduler, ) set_global_server_args_for_scheduler(server_args) - return get_global_server_args() + return get_server_args() def test_mistral_large3_forces_bfloat16(self): sa = self._construct("MistralLarge3ForCausalLM", "mistral") diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 491c6db94..220b6f73c 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -5,6 +5,7 @@ from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=5, suite="base-a-test-cpu") import dataclasses +import os import unittest from unittest.mock import patch @@ -207,6 +208,86 @@ class TestServerArgsOwnership(_IsolatedServerArgs): with self.assertRaises(ValueError): get_server_args() + +class TestServerArgsScopedOverride(_IsolatedServerArgs): + """ctx.override_server_args: the config tier's scoped test override — + tests force execution paths by overriding the context, not by + hand-building and publishing config objects.""" + + def test_install_publishes_fresh_config_with_fields(self): + reset_context() + override = get_context().override_server_args( + attention_backend="triton", chunked_prefill_size=-1 + ) + published = override.install() + self.assertIs(get_server_args(), published) + self.assertEqual(published.attention_backend, "triton") + self.assertEqual(published.chunked_prefill_size, -1) + # unnamed fields keep their dataclass defaults + self.assertEqual(published.tp_size, 1) + + def test_fields_carry_provenance(self): + published = get_context().override_server_args(tp_size=4).install() + self.assertIn(("test-override", {"tp_size": 4}), published._runtime_mutations) + + def test_restore_reinstates_previous_publish(self): + previous = object() + get_context().set_server_args(previous) + override = get_context().override_server_args(tp_size=8) + override.install() + self.assertEqual(get_server_args().tp_size, 8) + override.restore() + self.assertIs(get_server_args(), previous) + + def test_restore_reinstates_the_empty_slot(self): + reset_context() + with get_context().override_server_args(): + get_server_args() # published inside the scope + with self.assertRaises(ValueError): + get_server_args() + + def test_nesting_restores_in_order(self): + reset_context() + with get_context().override_server_args(tp_size=2) as outer: + with get_context().override_server_args(tp_size=4): + self.assertEqual(get_server_args().tp_size, 4) + self.assertIs(get_server_args(), outer) + self.assertEqual(get_server_args().tp_size, 2) + + def test_private_attribute_seeding(self): + # Property caches (e.g. _mamba_cache_chunk_size) are seeded through + # the same call; the strict guard exempts underscore names. + published = ( + get_context().override_server_args(_mamba_cache_chunk_size=64).install() + ) + self.assertEqual(published.mamba_cache_chunk_size, 64) + + def test_installed_config_arms_the_strict_guard(self): + # The published dummy must behave like a resolved config: bare writes + # raise under the strict harness; override() stays the entry point. + published = get_context().override_server_args(tp_size=2).install() + with self.assertRaises(AttributeError): + published.tp_size = 4 + published.override(source="test", tp_size=4) + self.assertEqual(published.tp_size, 4) + + def test_restore_resets_the_capture_seed(self): + # install() seeds flags.capture from the published dummy; restore() + # must put back the pre-install runtime state on both restore paths. + reset_context() + self.assertFalse(get_flags().capture.enable_torch_compile) + override = get_context().override_server_args(enable_torch_compile=True) + override.install() + self.assertTrue(get_flags().capture.enable_torch_compile) + override.restore() + self.assertFalse(get_flags().capture.enable_torch_compile) + + def test_double_install_rejected(self): + override = get_context().override_server_args() + override.install() + with self.assertRaises(AssertionError): + override.install() + def test_module_global_removed(self): # The legacy storage must not survive: a stale _global_server_args would # silently fork the config into two objects. @@ -462,6 +543,58 @@ class TestNamedStreams(_IsolatedServerArgs): reset_context() self.assertEqual(get_context().resources.streams, {}) + def test_capturer_slots_roundtrip_and_reset(self): + from sglang.srt.state_capturer.indexer_topk import ( + get_global_indexer_capturer, + set_global_indexer_capturer, + ) + from sglang.srt.state_capturer.routed_experts import ( + get_global_experts_capturer, + set_global_experts_capturer, + ) + + reset_context() + self.assertIsNone(get_global_indexer_capturer()) + self.assertIsNone(get_global_experts_capturer()) + indexer, experts = object(), object() + set_global_indexer_capturer(indexer) + set_global_experts_capturer(experts) + self.assertIs(get_global_indexer_capturer(), indexer) + self.assertIs(get_global_experts_capturer(), experts) + reset_context() + self.assertIsNone(get_global_indexer_capturer()) + self.assertIsNone(get_global_experts_capturer()) + + def test_tcp_store_slot_roundtrip_and_reset(self): + from sglang.srt.distributed.utils import ( + get_global_tcp_store, + set_global_tcp_store, + ) + + reset_context() + self.assertIsNone(get_global_tcp_store()) + store = object() + set_global_tcp_store(store) + self.assertIs(get_global_tcp_store(), store) + reset_context() + self.assertIsNone(get_global_tcp_store()) + + def test_trace_level_env_seeded_lazy_default(self): + from sglang.srt.observability.trace import ( + get_global_trace_level, + set_global_trace_level, + ) + + reset_context() + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("SGLANG_TRACE_LEVEL", None) + self.assertEqual(get_global_trace_level(), 3) + set_global_trace_level(5) + self.assertEqual(get_global_trace_level(), 5) + reset_context() + with patch.dict(os.environ, {"SGLANG_TRACE_LEVEL": "1"}): + self.assertEqual(get_global_trace_level(), 1) + class TestEpBufferState(_IsolatedServerArgs): """EP dispatcher buffer managers: state lives on ctx.resources; the diff --git a/test/registered/unit/test_server_args_mutation_ratchet.py b/test/registered/unit/test_server_args_mutation_ratchet.py index c8711a632..848c2d78c 100644 --- a/test/registered/unit/test_server_args_mutation_ratchet.py +++ b/test/registered/unit/test_server_args_mutation_ratchet.py @@ -38,17 +38,18 @@ _MUTATION_PATTERNS = [ re.compile(r"\bserver_args\.[a-z0-9_]+\s*=(?![=}])"), re.compile(r"\bsa\.[a-z0-9_]+\s*=(?![=}])"), re.compile(r"get_(?:global_)?server_args\(\)\.[a-z0-9_]+\s*=(?![=}])"), + # setattr is the same write with the attribute name behind a variable. + re.compile( + r"setattr\(\s*(?:[\w.]+\.)?(?:server_args|sa|get_(?:global_)?server_args\(\))\s*," + ), ] -# The resolution pipeline itself (mutation is its job); multimodal_gen, whose -# ServerArgs is a different class outside this contract; and the sanctioned -# mock-fixture factory (bare object.__new__ instances never materialize, so -# the strict guard does not apply to their construction). +# The resolution pipeline itself (mutation is its job) and multimodal_gen, +# whose ServerArgs is a different class outside this contract. _EXCLUDED = ( "srt/server_args.py", "srt/arg_groups", "multimodal_gen", - "test/kits/attention_unittest/mock_server_args.py", ) _BASELINE = 0