[refactor] Retire the legacy config accessor and the remaining process singletons (#30493)

This commit is contained in:
Cheng Wan
2026-07-09 02:10:47 -07:00
committed by GitHub
parent 122b3266a2
commit 1f15308dca
158 changed files with 773 additions and 710 deletions
+1 -1
View File
@@ -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,
)
@@ -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
+2 -2
View File
@@ -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
@@ -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
@@ -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
@@ -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)
+8 -14
View File
@@ -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):
+2 -2
View File
@@ -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
@@ -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
@@ -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)
@@ -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 (
@@ -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
@@ -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(
@@ -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] = []
+2 -3
View File
@@ -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
@@ -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:
@@ -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"
)
@@ -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(
@@ -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
@@ -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()
@@ -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(
@@ -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
+5 -8
View File
@@ -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), (
@@ -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
+11 -12
View File
@@ -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(
+4 -4
View File
@@ -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):
+2 -2
View File
@@ -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))
+1 -1
View File
@@ -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
+7 -8
View File
@@ -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]
+2 -2
View File
@@ -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
@@ -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
+6 -10
View File
@@ -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,
+2 -3
View File
@@ -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
+3 -4
View File
@@ -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()
+2 -2
View File
@@ -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
)
@@ -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:
+2 -3
View File
@@ -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
@@ -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()
@@ -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)
@@ -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,
@@ -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 = (
+6 -7
View File
@@ -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
+2 -2
View File
@@ -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)
@@ -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())
@@ -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:
@@ -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):
@@ -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
@@ -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):
+3 -6
View File
@@ -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:
+4 -5
View File
@@ -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()
+7 -7
View File
@@ -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):
+12 -12
View File
@@ -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()
@@ -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)
+9 -10
View File
@@ -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(
@@ -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:
@@ -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()
@@ -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)
+4 -4
View File
@@ -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
@@ -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)
@@ -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,
)
+2 -2
View File
@@ -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,
)
@@ -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
@@ -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
@@ -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:
@@ -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
@@ -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,
)
@@ -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
@@ -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(
@@ -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,
)
+6 -6
View File
@@ -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
+2 -3
View File
@@ -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(
@@ -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,
+2 -5
View File
@@ -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)
@@ -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 (
@@ -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
@@ -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")
+3 -4
View File
@@ -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)
+12 -14
View File
@@ -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
+3 -3
View File
@@ -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 "
+2 -4
View File
@@ -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
+3 -5
View File
@@ -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 "
+2 -2
View File
@@ -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():
+1 -3
View File
@@ -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
+2 -4
View File
@@ -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
@@ -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
+1 -5
View File
@@ -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
+2 -3
View File
@@ -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,
+1 -2
View File
@@ -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
+2 -2
View File
@@ -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,
+1 -2
View File
@@ -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,
+2 -3
View File
@@ -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
+2 -2
View File
@@ -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,
+1 -3
View File
@@ -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,
+2 -3
View File
@@ -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(
+2 -2
View File
@@ -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(
+2 -2
View File
@@ -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
+2 -4
View File
@@ -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
+2 -2
View File
@@ -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
+3 -4
View File
@@ -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,
+2 -2
View File
@@ -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:

Some files were not shown because too many files have changed in this diff Show More