[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 # 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 ( from sglang.srt.layers.attention.dsa.utils import (
aiter_can_use_preshuffle_paged_mqa, 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.model_executor.forward_context import get_attn_backend
from sglang.srt.runtime_context import get_parallel, get_server_args 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.speculative.spec_info import SpecInput
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip 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 cpu_value
if isinstance(cpu_value, torch.Tensor) if isinstance(cpu_value, torch.Tensor)
else torch.tensor(cpu_value, dtype=old_device_value.dtype) 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) setattr(batch, device_field, new_device_value)
if sum_field is not None: if sum_field is not None:
@@ -336,7 +335,7 @@ def compute_split_indices_for_cuda_graph_replay(
class TboCudaGraphRunnerPlugin: class TboCudaGraphRunnerPlugin:
def __init__(self): def __init__(self):
self._tbo_children_num_token_non_padded = torch.zeros( 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): 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` # TODO improve, e.g. unify w/ `init_raw`
if ( 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 and batch.global_dp_buffer_len is not None
): ):
sum_len = end_token_index - start_token_index sum_len = end_token_index - start_token_index
@@ -835,7 +834,7 @@ class TboForwardBatchPreparer:
value_a = min(tbo_split_token_index, num_token_non_padded) value_a = min(tbo_split_token_index, num_token_non_padded)
value_b = max(0, num_token_non_padded - tbo_split_token_index) value_b = max(0, num_token_non_padded - tbo_split_token_index)
return torch.tensor([value_a, value_b], dtype=torch.int32).to( 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 @classmethod
+2 -2
View File
@@ -1783,9 +1783,9 @@ class _SGLangPlugin(_FrameworkPlugin):
return None return None
try: 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: if args is None:
return 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 ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, 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__) logger = logging.getLogger(__name__)
@@ -25,7 +25,7 @@ class PyMscclppCommunicator:
def _is_symm_mem_enabled(self) -> bool: def _is_symm_mem_enabled(self) -> bool:
try: try:
return get_global_server_args().enable_symm_mem return get_server_args().enable_symm_mem
except ValueError: except ValueError:
return False return False
@@ -15,7 +15,7 @@ from torch.cuda.memory import (
from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.environ import envs 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 from sglang.srt.utils.common import torch_release
after_2_8_0 = torch_release >= (2, 8) after_2_8_0 = torch_release >= (2, 8)
@@ -159,7 +159,7 @@ _register_func = None
def is_symmetric_memory_enabled(): def is_symmetric_memory_enabled():
try: try:
return get_global_server_args().enable_symm_mem return get_server_args().enable_symm_mem
except ValueError: except ValueError:
return False 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 NCCL. Covers the spec-decode batch plus, if ``include_prefill``, a prefill
chunk. Returns ``floor`` if server args are unavailable.""" chunk. Returns ``floor`` if server args are unavailable."""
try: 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: def g(name: str) -> int:
v = getattr(sa, name, 0) v = getattr(sa, name, 0)
+8 -14
View File
@@ -19,19 +19,13 @@ from torch.distributed import TCPStore
logger = logging.getLogger(__name__) 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: 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 get_resources().tcp_store = store
the store available to all components that need it.
"""
global _global_tcp_store
_global_tcp_store = store
logger.info("Global TCPStore has been set") logger.info("Global TCPStore has been set")
@@ -45,15 +39,15 @@ def get_global_tcp_store() -> Optional[TCPStore]:
Returns: Returns:
The global TCPStore instance, or None if not initialized yet. 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( logger.warning(
"Global TCPStore not found. Make sure init_distributed_environment " "Global TCPStore not found. Make sure init_distributed_environment "
"was called with a tcp:// init method." "was called with a tcp:// init method."
) )
return store
return _global_tcp_store
def ensure_divisibility(numerator, denominator): def ensure_divisibility(numerator, denominator):
+2 -2
View File
@@ -166,7 +166,7 @@ def init_tokenizer_manager(
if getattr(server_args, attr) != "auto": if getattr(server_args, attr) != "auto":
continue continue
if suggested is not None: if suggested is not None:
setattr(server_args, attr, suggested) server_args.override(source="template-detection", **{attr: suggested})
logger.info( logger.info(
f"Auto-detected --{attr.replace('_', '-')} as '{suggested}' from chat template" 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"--{attr.replace('_', '-')}=auto specified but could not detect "
f"{label} from chat template. Disabling {label}." 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 return tokenizer_manager, template_manager
@@ -18,7 +18,7 @@ from typing import Literal, Optional
import torch import torch
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata 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 from sglang.srt.utils import is_hip
_is_hip = is_hip() _is_hip = is_hip()
@@ -37,7 +37,7 @@ class ExpertLocationDispatchInfo:
@classmethod @classmethod
def init_new(cls, layer_id: int): 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() expert_location_metadata = get_global_expert_location_metadata()
assert expert_location_metadata is not None assert expert_location_metadata is not None
@@ -25,7 +25,7 @@ from sglang.srt.eplb.expert_location import (
ExpertLocationMetadata, ExpertLocationMetadata,
get_global_expert_location_metadata, 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 from sglang.srt.utils import get_bool_env_var, get_int_env_var, is_hip
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -108,7 +108,7 @@ def _update_expert_weights_with_canary(
canary_tensor = ( canary_tensor = (
_get_canary_value(old_expert_location_metadata, layer_id) _get_canary_value(old_expert_location_metadata, layer_id)
.clone() .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) 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, uses_sliding_window_attention,
) )
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool 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__) logger = logging.getLogger(__name__)
@@ -281,7 +281,7 @@ class MlxModelRunner:
): ):
return None 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 track_len = prefix_len + (new_token_count // chunk_size) * chunk_size
branching_len = getattr(req, "mamba_branching_seqlen", None) branching_len = getattr(req, "mamba_branching_seqlen", None)
if ( if (
@@ -23,7 +23,7 @@ from sglang.srt.layers.utils.cp_utils import (
cp_allgather_and_save_kv_cache, cp_allgather_and_save_kv_cache,
) )
from sglang.srt.mem_cache.memory_pool import KVWriteLoc 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: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention 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() and not forward_batch.forward_mode.is_draft_extend_v2()
): ):
if forward_batch.attn_attend_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
assert forward_batch.prefix_chunk_idx is not None 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_cu_seq_lens is not None
assert forward_batch.prefix_chunk_max_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() or forward_batch.forward_mode.is_draft_extend_v2()
): ):
B = forward_batch.batch_size 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( actual_q = torch.arange(
n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device 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() forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2() 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: else:
max_seqlen_q = 1 max_seqlen_q = 1
return self._kernel_metadata_from_parts( 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.layers.attention.vision import VisionAttention
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner 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): class ViTNpuGraphRunner(ViTCudaGraphRunner):
@@ -70,7 +70,7 @@ class ViTNpuGraphRunner(ViTCudaGraphRunner):
graph = torch_npu.npu.NPUGraph() graph = torch_npu.npu.NPUGraph()
vit = self.vit 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): with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool):
y = None y = None
deepstack_outs: List[torch.Tensor] = [] deepstack_outs: List[torch.Tensor] = []
+2 -3
View File
@@ -33,8 +33,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.runtime_context import get_parallel 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 ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -90,7 +89,7 @@ logger = logging.getLogger(__name__)
class SiluAndMul(MultiPlatformOp): class SiluAndMul(MultiPlatformOp):
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
super().__init__(*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 self._forward_method = self.forward_native
elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get(): elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get():
self._forward_method = self.forward_aiter self._forward_method = self.forward_aiter
@@ -112,7 +112,7 @@ from sglang.srt.model_executor.forward_context import (
get_token_to_kv_pool, get_token_to_kv_pool,
) )
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.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() _use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
if TYPE_CHECKING: 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: def _uses_dsa_attention_backend(forward_batch: ForwardBatch) -> bool:
attn_backend = get_attn_backend() 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, decode_backend = server_args.get_attention_backends()
prefill_backend = ( prefill_backend = (
getattr(attn_backend, "prefill_attention_backend_str", None) or prefill_backend getattr(attn_backend, "prefill_attention_backend_str", None) or prefill_backend
@@ -394,7 +394,7 @@ class Indexer(MultiPlatformOp):
if _is_cuda: if _is_cuda:
self.sm_count = deep_gemm.get_num_sms() self.sm_count = deep_gemm.get_num_sms()
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) 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 self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank
else: else:
self.logits_with_pp_recv = False self.logits_with_pp_recv = False
@@ -446,7 +446,7 @@ class Indexer(MultiPlatformOp):
base=rope_theta, # type: ignore base=rope_theta, # type: ignore
rope_scaling=rope_scaling, rope_scaling=rope_scaling,
is_neox_style=is_neox_style, is_neox_style=is_neox_style,
device=get_global_server_args().device, device=get_server_args().device,
) )
self.block_size = block_size self.block_size = block_size
self.scale_fmt = scale_fmt 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 = torch.cuda.get_device_properties(device_index).total_memory
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION) 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: if mem_fraction_static is None:
static_budget = total_mem_budget static_budget = total_mem_budget
else: 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 ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_parallel 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 get_bool_env_var, is_cuda, is_hip from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip
from sglang.srt.utils.common import ceil_align, ceil_div 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(): 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(): def is_dsa_prefill_cp_in_seq_split():
return ( return (
is_dsa_enable_prefill_cp() 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(): def is_dsa_prefill_cp_round_robin_split():
return ( return (
is_dsa_enable_prefill_cp() 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.rotary_emb = rotary_emb
self.freqs_cis = freqs_cis self.freqs_cis = freqs_cis
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5 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 self.alt_streams = alt_streams
def compute_q( 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.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode 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.speculative.spec_info import SpecInput, SpeculativeAlgorithm
from sglang.srt.utils import get_compiler_backend from sglang.srt.utils import get_compiler_backend
@@ -1342,7 +1342,7 @@ class FlashAttentionBackend(AttentionBackend):
): ):
# Do multi-head attention with chunked prefix cache # Do multi-head attention with chunked prefix cache
if forward_batch.attn_attend_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 # 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_idx is not None
assert forward_batch.prefix_chunk_cu_seq_lens 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 ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_buffer from sglang.srt.runtime_context import get_buffer, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width, draft_kv_indices_buffer_width,
@@ -226,9 +225,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.enable_chunk_kv = ( self.enable_chunk_kv = (
not skip_prefill not skip_prefill
and get_global_server_args().disaggregation_mode != "decode" and get_server_args().disaggregation_mode != "decode"
and not get_global_server_args().disable_chunked_prefix_cache and not get_server_args().disable_chunked_prefix_cache
and not get_global_server_args().flashinfer_mla_disable_ragged and not get_server_args().flashinfer_mla_disable_ragged
) )
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
@@ -404,7 +403,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
prefix_lens = forward_batch.extend_prefix_lens prefix_lens = forward_batch.extend_prefix_lens
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
use_ragged = ( use_ragged = (
not get_global_server_args().flashinfer_mla_disable_ragged not get_server_args().flashinfer_mla_disable_ragged
and extend_no_prefix and extend_no_prefix
# Piecewise cuda graph should use paged prefill to be compatible with prefix cache # Piecewise cuda graph should use paged prefill to be compatible with prefix cache
and not is_in_tc_piecewise_cuda_graph() 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.mem_cache.memory_pool import HybridReqToTokenPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.model_runner import ModelRunner 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.eagle_info import EagleDraftInput, EagleVerifyInput
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
@@ -246,7 +246,7 @@ class MambaAttnBackendBase(AttentionBackend):
lens_to_track = ( lens_to_track = (
forward_batch.mamba_track_seqlens - forward_batch.extend_prefix_lens 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 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 = query_start_loc[:-1] + aligned_len - conv_state_len
start_indices = start_indices[forward_batch.mamba_track_mask] 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 """src/dst indices to track SSM states for prefix caching: aligned seqs
cache last_recurrent_state, unaligned cache intermediate `h` at the last cache last_recurrent_state, unaligned cache intermediate `h` at the last
chunk boundary.""" 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 # CPU to avoid kernel launches for the masking ops
mamba_track_mask = forward_batch.mamba_track_mask.cpu() mamba_track_mask = forward_batch.mamba_track_mask.cpu()
extend_seq_lens = forward_batch.extend_seq_lens.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 % """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 mamba_track_interval == 0, so force-flush and snapshot fire on the same
steps (no off-by-one).""" 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: if seq_lens_cpu is None:
# Should not happen for the supported config; stay safe and never flush. # Should not happen for the supported config; stay safe and never flush.
return torch.zeros((bs,), dtype=torch.bool) 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 # 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. # reads it (CUDA causal_conv1d garbles it). A model may also force Triton.
use_triton_causal_conv = ( use_triton_causal_conv = (
use_triton_causal_conv use_triton_causal_conv or get_server_args().enable_page_major_kv_layout
or get_global_server_args().enable_page_major_kv_layout
) )
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
mixer_out, intermediate_states = mixer.forward( 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 ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_buffer, get_parallel from sglang.srt.runtime_context import get_buffer, get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2 from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2
if is_flashinfer_available(): if is_flashinfer_available():
@@ -199,7 +198,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None
self.disable_chunked_prefix_cache = ( 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 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.quantization import QuantizationConfig
from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb 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 from sglang.srt.utils import add_prefix, get_bool_env_var
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _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 # Select attention backend via a unified method
_passed_backend = qkv_backend _passed_backend = qkv_backend
qkv_backend = self._determine_attention_backend(_passed_backend) qkv_backend = self._determine_attention_backend(_passed_backend)
if ( if get_server_args().mm_attention_backend is None and _passed_backend is None:
get_global_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"Multimodal attention backend not set. Use {qkv_backend}.")
print_info_once(f"Using {qkv_backend} as multimodal attention backend.") print_info_once(f"Using {qkv_backend} as multimodal attention backend.")
@@ -1047,7 +1044,7 @@ class VisionAttention(nn.Module):
weight_dtype=torch.float32, weight_dtype=torch.float32,
cast_x_before_out_mul=True, 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 {} else {}
) )
q_norm = RMSNorm( q_norm = RMSNorm(
@@ -1075,7 +1072,7 @@ class VisionAttention(nn.Module):
- CUDA (other): "triton_attn" - CUDA (other): "triton_attn"
- Non-CUDA: "sdpa" - 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: if override_backend is not None:
backend = override_backend backend = override_backend
elif passed_backend is not None: elif passed_backend is not None:
@@ -1179,7 +1176,7 @@ class VisionAttention(nn.Module):
x = x.unsqueeze(0) x = x.unsqueeze(0)
assert x.dim() == 3, x.shape assert x.dim() == 3, x.shape
if ( 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 and position_embeddings is not None
): ):
assert isinstance(position_embeddings, tuple), ( assert isinstance(position_embeddings, tuple), (
@@ -12,10 +12,10 @@ from sglang.srt.layers.attention.flashattention_backend import (
merge_state_v2_wrapper, merge_state_v2_wrapper,
prepare_swa_spec_page_table_triton, 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.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -638,7 +638,7 @@ class XPUAttentionBackend(AttentionBackend):
): ):
# Do multi-head attention with chunked prefix cache # Do multi-head attention with chunked prefix cache
if forward_batch.attn_attend_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 # 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_idx is not None
assert forward_batch.prefix_chunk_cu_seq_lens 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, check_cuda_graph_backend,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_forward, get_parallel from sglang.srt.runtime_context import get_forward, 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.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
@@ -171,7 +170,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
and batch_size > 0 and batch_size > 0
and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE
and not is_dp_attention_enabled() 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() 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 total_bytes <= 8 * 1024 * 8192
and get_parallel().tp_size != 6 and get_parallel().tp_size != 6
and not is_dp_attention_enabled() 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): def init_context(self, q_lora_rank, is_dsa):
self.is_dsa = is_dsa self.is_dsa = is_dsa
self.allow_input_scattered = ( self.allow_input_scattered = (
get_global_server_args().enable_attn_tp_input_scattered get_server_args().enable_attn_tp_input_scattered
and (_is_cuda or _is_npu) and (_is_cuda or _is_npu)
and q_lora_rank is not None and q_lora_rank is not None
and not is_dsa and not is_dsa
@@ -275,9 +274,9 @@ class AttnTpContext:
and get_moe_a2a_backend().is_none() and get_moe_a2a_backend().is_none()
and not enable_moe_dense_fully_dp() and not enable_moe_dense_fully_dp()
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) 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: if not self.allow_input_scattered:
logging.info( logging.info(
"attn_tp_input_scattered is not enabled while other conditions are not met" "attn_tp_input_scattered is not enabled while other conditions are not met"
@@ -407,7 +406,7 @@ class LayerScatterModes:
not context.is_layer_sparse not context.is_layer_sparse
and context.is_next_layer_sparse and context.is_next_layer_sparse
and enable_moe_dense_fully_dp() and enable_moe_dense_fully_dp()
and get_global_server_args().enable_two_batch_overlap and get_server_args().enable_two_batch_overlap
) )
@classmethod @classmethod
@@ -434,7 +433,7 @@ class LayerScatterModes:
def enable_moe_dense_fully_dp(): 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: class LayerCommunicator:
@@ -463,7 +462,7 @@ class LayerCommunicator:
) )
self._post_init_communicate() self._post_init_communicate()
self._speculative_algo = SpeculativeAlgorithm.from_string( self._speculative_algo = SpeculativeAlgorithm.from_string(
get_global_server_args().speculative_algorithm get_server_args().speculative_algorithm
) )
def _post_init_communicate(self): def _post_init_communicate(self):
@@ -811,7 +810,7 @@ class LayerCommunicator:
and get_parallel().tp_size != 6 and get_parallel().tp_size != 6
and not is_dp_attention_enabled() and not is_dp_attention_enabled()
and get_moe_a2a_backend().is_none() 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) and (not self.is_last_layer)
@@ -1116,7 +1115,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
if not handled: if not handled:
quantize_communications = ( quantize_communications = (
not forward_batch.forward_mode.is_decode_or_idle() 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: if quantize_communications:
hidden_states = attention_tensor_model_parallel_quant_all_reduce( 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: 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( return bool(
getattr(sa, "enable_prefill_cp", False) getattr(sa, "enable_prefill_cp", False)
and getattr(sa, "_is_dsa_model_arch", False) and getattr(sa, "_is_dsa_model_arch", False)
@@ -247,10 +247,10 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]:
global _STRATEGY global _STRATEGY
if _STRATEGY is None: if _STRATEGY is None:
from sglang.srt.server_args import get_global_server_args from sglang.srt.runtime_context import get_server_args
try: try:
server_args = get_global_server_args() server_args = get_server_args()
except ValueError: except ValueError:
return None return None
if server_args is not None and getattr(server_args, "enable_prefill_cp", False): 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_prev_list.append(block_sizes[cp_rank])
actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1]) 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: try:
device = torch.device(get_global_server_args().device) device = torch.device(get_server_args().device)
except Exception: except Exception:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
cu_prev = [0] + list(accumulate(actual_seq_q_prev_list)) 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 from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
# NOTE: planner.py is intentionally NOT imported here. It depends on server_args # 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, # 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, # 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 # 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.layout import update_local_kv_lens_for_dcp
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
def prepare_decode_context_parallel_metadata( def prepare_decode_context_parallel_metadata(
@@ -54,12 +53,12 @@ def prepare_decode_context_parallel_metadata(
extend_prefix_starts = torch.zeros( extend_prefix_starts = torch.zeros(
len(seq_lens), len(seq_lens),
dtype=torch.int32, dtype=torch.int32,
device=get_global_server_args().device, device=get_server_args().device,
) )
extend_cu_prefix_lens = torch.zeros( extend_cu_prefix_lens = torch.zeros(
len(seq_lens) + 1, len(seq_lens) + 1,
dtype=torch.int32, 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[1:] = torch.cumsum(extend_prefix_lens, dim=0)
extend_cu_prefix_lens = extend_cu_prefix_lens[:-1] 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( dcp_prefix_kv_indices = torch.empty(
sum(extend_prefix_lens_cpu), sum(extend_prefix_lens_cpu),
dtype=torch.int32, dtype=torch.int32,
device=get_global_server_args().device, device=get_server_args().device,
) )
create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)]( create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)](
req_to_token, req_to_token,
@@ -82,20 +81,20 @@ def prepare_decode_context_parallel_metadata(
dcp_kv_indptr = torch.zeros( dcp_kv_indptr = torch.zeros(
len(seq_lens) + 1, len(seq_lens) + 1,
dtype=torch.int32, 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[1:] = seq_lens.cumsum(dim=0)
dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)] dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)]
dcp_kv_indices = torch.zeros( dcp_kv_indices = torch.zeros(
seq_lens_sum, seq_lens_sum,
dtype=torch.int32, dtype=torch.int32,
device=get_global_server_args().device, device=get_server_args().device,
) )
extend_cu_lens = torch.zeros( extend_cu_lens = torch.zeros(
len(seq_lens) + 1, len(seq_lens) + 1,
dtype=torch.int32, 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[1:] = torch.cumsum(extend_seq_lens, dim=0)
extend_cu_lens = extend_cu_lens[:-1] 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: def is_enable_moe_cp_allgather() -> bool:
"""True when moe_dp_size < attn_cp_size, requiring allgather across CP ranks before MoE.""" """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 return sa.attn_cp_size > sa.moe_dp_size
@@ -13,8 +13,7 @@ from sglang.srt.distributed import (
get_tp_group, get_tp_group,
) )
from sglang.srt.distributed.parallel_state import in_the_same_node_as from sglang.srt.distributed.parallel_state import in_the_same_node_as
from sglang.srt.runtime_context import get_parallel 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 ( from sglang.srt.utils import (
ceil_align, ceil_align,
get_cuda_driver_bindings, get_cuda_driver_bindings,
@@ -673,7 +672,7 @@ def ensure_workspace_initialized(
token_num = token_num or max_token_num token_num = token_num or max_token_num
group_key = (device_group, cpu_group) group_key = (device_group, cpu_group)
effective_dtype = dtype or torch.bfloat16 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) backend = resolve_flashinfer_allreduce_fusion_backend(server_args)
if backend is None: if backend is None:
return False return False
+6 -10
View File
@@ -31,8 +31,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.runtime_context import get_parallel 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 ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -193,7 +192,7 @@ def _forward_with_allreduce_fusion(
return fused_result return fused_result
# For AITER route, preserve correctness when fused path is unavailable. # 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) x = tensor_model_parallel_all_reduce(x)
return norm_module.forward(x, residual, None) return norm_module.forward(x, residual, None)
@@ -271,7 +270,7 @@ class RMSNorm(MultiPlatformOp):
if ( if (
residual is not None residual is not None
or self.cast_x_before_out_mul 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 self.forward_native(x, residual, post_residual_addition)
return rms_norm_batch_invariant( return rms_norm_batch_invariant(
@@ -371,7 +370,7 @@ class RMSNorm(MultiPlatformOp):
if ( if (
residual is not None residual is not None
or self.cast_x_before_out_mul 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) or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0)
): ):
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
@@ -432,7 +431,7 @@ class RMSNorm(MultiPlatformOp):
if ( if (
residual is not None residual is not None
or self.cast_x_before_out_mul 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 self.forward_native(x, residual, post_residual_addition)
return rms_norm_batch_invariant( return rms_norm_batch_invariant(
@@ -559,10 +558,7 @@ class RMSNorm(MultiPlatformOp):
if self.variance_size_override is not None: if self.variance_size_override is not None:
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
if is_batch_invariant_mode_enabled(): if is_batch_invariant_mode_enabled():
if ( if residual is not None or get_server_args().rl_on_policy_target == "fsdp":
residual is not None
or get_global_server_args().rl_on_policy_target == "fsdp"
):
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
return rms_norm_batch_invariant( return rms_norm_batch_invariant(
x, x,
+2 -3
View File
@@ -37,8 +37,7 @@ from sglang.srt.layers.parameter import (
_ColumnvLLMParameter, _ColumnvLLMParameter,
) )
from sglang.srt.layers.utils import pad_or_narrow_weight from sglang.srt.layers.utils import pad_or_narrow_weight
from sglang.srt.runtime_context import get_parallel 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 get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -1545,7 +1544,7 @@ class RowParallelLinear(LinearBase):
quantize_communications = ( quantize_communications = (
( (
not forward_batch.forward_mode.is_decode_or_idle() 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 if forward_batch is not None
else False else False
+3 -4
View File
@@ -48,7 +48,6 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args 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 ( from sglang.srt.utils.common import (
is_cpu, is_cpu,
is_npu, is_npu,
@@ -336,7 +335,7 @@ class LogitsProcessor(nn.Module):
self.vocab_size = config.vocab_size self.vocab_size = config.vocab_size
self.logit_scale = logit_scale self.logit_scale = logit_scale
self.use_attn_tp_group = get_server_args().enable_dp_lm_head self.use_attn_tp_group = get_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: if self.use_attn_tp_group:
self.attn_tp_size = get_parallel().attn_tp_size self.attn_tp_size = get_parallel().attn_tp_size
self.do_tensor_parallel_all_gather = ( self.do_tensor_parallel_all_gather = (
@@ -360,7 +359,7 @@ class LogitsProcessor(nn.Module):
self.final_logit_softcapping = None self.final_logit_softcapping = None
self.return_full_logits = return_full_logits 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( self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer(
max_tokens=triton_symm_mem_ag.recommended_max_tokens( max_tokens=triton_symm_mem_ag.recommended_max_tokens(
@@ -970,7 +969,7 @@ class LogitsProcessor(nn.Module):
None, # bias None, # bias
True, # is_vnni 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 # Due to tie-weight, we may not be able to change lm_head's weight dtype
logits = torch.matmul( logits = torch.matmul(
hidden_states.bfloat16(), lm_head.weight.T.bfloat16() 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 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. 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] 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( buckets = get_mhc_pre_token_count_representatives(
max_num_tokens, hc_mult * hidden_size 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, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight 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.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -276,7 +275,7 @@ class FusedMoE(torch.nn.Module):
) )
self.quant_method: Optional[FusedMoEMethodBase] = None 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) kt_config = create_kt_config_from_server_args(server_args, layer_id)
if kt_config is not None: if kt_config is not None:
if quant_config is not None: if quant_config is not None:
+2 -3
View File
@@ -44,11 +44,10 @@ class HashTopK(nn.Module):
): ):
super().__init__() super().__init__()
self.layer_id = layer_id 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 = ( self.enable_deepep_waterfill = (
num_fused_shared_experts > 0 num_fused_shared_experts > 0 and get_server_args().enable_deepep_waterfill
and get_global_server_args().enable_deepep_waterfill
) )
self.deepep_waterfill_balancer = None self.deepep_waterfill_balancer = None
@@ -244,14 +244,14 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None:
"Install with: pip install flashinfer" "Install with: pip install flashinfer"
) from e ) 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, ( assert layer.intermediate_size_per_partition > 0, (
f"CuteDSL MoE: intermediate_size_per_partition must be > 0, " f"CuteDSL MoE: intermediate_size_per_partition must be > 0, "
f"got {layer.intermediate_size_per_partition}. Check EP/TP configuration." 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 # CuteDSL wrapper preallocates CG buffers used by any captured graph
# that routes through this MoE — decode and prefill alike. # that routes through this MoE — decode and prefill alike.
use_cuda_graph = not cuda_graph_fully_disabled() 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.environ import envs
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.utils import get_moe_padding_size 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 ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -488,7 +488,7 @@ def _fused_moe_kernel_sequence(
out_hidden_states = torch.empty_like(hidden_states) out_hidden_states = torch.empty_like(hidden_states)
use_fused_moe_sum_all_reduce = ( 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 (not no_combine)
and (topk > 2) and (topk > 2)
and (not use_int8_w8a16) and (not use_int8_w8a16)
@@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional, Tuple
import torch import torch
import triton 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 from sglang.srt.utils import get_device_name, is_hip
logger = logging.getLogger(__name__) 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 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. 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( logger.warning(
"Deterministic inference is enabled, using default MoE kernel config." "Deterministic inference is enabled, using default MoE kernel config."
) )
@@ -170,7 +170,7 @@ def get_default_config(
is_marlin: bool, is_marlin: bool,
block_shape: Optional[List[int]] = None, block_shape: Optional[List[int]] = None,
) -> Dict[str, int]: ) -> Dict[str, int]:
if get_global_server_args().enable_deterministic_inference: if get_server_args().enable_deterministic_inference:
config = { config = {
"BLOCK_SIZE_M": 64, "BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 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.topk import StandardTopKOutput, TopKOutput
from sglang.srt.layers.moe.utils import get_moe_runner_backend 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.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import get_int_env_var 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, # 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 # 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). # (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) default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096)
self.max_num_tokens = get_int_env_var( self.max_num_tokens = get_int_env_var(
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", "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. # Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized.
speculative_algo = SpeculativeAlgorithm.from_string( 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(): if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle():
total_dispatch_payload_size_per_token = ( 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 assert num_expert_group is not None and topk_group is not None
self.layer_id = layer_id 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 = ( self.enable_deepep_waterfill = (
num_fused_shared_experts > 0 num_fused_shared_experts > 0 and get_server_args().enable_deepep_waterfill
and get_global_server_args().enable_deepep_waterfill
) )
self.deepep_waterfill_balancer = None self.deepep_waterfill_balancer = None
@@ -475,9 +474,9 @@ class TopK(MultiPlatformOp):
# ===== TO BE REFACTORED ==== # ===== TO BE REFACTORED ====
elif get_moe_runner_backend().is_experimental_sgl_trtllm(): elif get_moe_runner_backend().is_experimental_sgl_trtllm():
try: 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: except ValueError:
use_standard_for_lora = False use_standard_for_lora = False
output_format = ( output_format = (
@@ -1256,10 +1255,10 @@ def _eplb_remap_enabled() -> bool:
# initial expert placement is non-trivial, or there are redundant physical # 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 # 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). # 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: try:
server_args = get_global_server_args() server_args = get_server_args()
except ValueError: except ValueError:
# Global server args are not initialized outside the server runtime # Global server args are not initialized outside the server runtime
# (e.g. in unit tests that call select_experts directly). In that case # (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: if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs 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__) logger = logging.getLogger(__name__)
@@ -206,7 +206,7 @@ def get_deepep_output_dtype(self) -> DeepEPOutputDtype:
""" """
# 0. Parse server argument. # 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": if server_args and server_args.deepep_dispatcher_output_dtype != "auto":
return DeepEPOutputDtype(server_args.deepep_dispatcher_output_dtype) 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_deepgemm,
w8a8_block_fp8_matmul_triton, 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 ( from sglang.srt.utils import (
ceil_align, ceil_align,
ceil_div, ceil_div,
@@ -1691,8 +1691,7 @@ def apply_fp8_linear(
if ( if (
input_scale is not None input_scale is not None
and input_scale.numel() == 1 and input_scale.numel() == 1
and get_global_server_args().cuda_graph_config.prefill.tc_compiler and get_server_args().cuda_graph_config.prefill.tc_compiler == "inductor"
== "inductor"
): ):
qinput = ( qinput = (
(input_2d * input_scale.reciprocal()) (input_2d * input_scale.reciprocal())
@@ -48,7 +48,7 @@ from sglang.srt.layers.quantization.base_config import (
QuantizeMethodBase, QuantizeMethodBase,
) )
from sglang.srt.layers.quantization.utils import is_layer_skipped 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 ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
is_cpu, is_cpu,
@@ -334,7 +334,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4() self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4()
self.use_marlin = get_moe_runner_backend().is_marlin() self.use_marlin = get_moe_runner_backend().is_marlin()
self.flashinfer_mxfp4_moe_precision = ( 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 # When `flashinfer_mxfp4` is enabled, dispatch to one of two FlashInfer
# entry points depending on the GPU: # 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.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe.utils import RoutingMethodType 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 ( from sglang.srt.utils import (
is_flashinfer_available, is_flashinfer_available,
log_info_on_rank0, log_info_on_rank0,
@@ -128,7 +128,7 @@ class Mxfp4FlashinferTrtllmMoEMethod:
self._fp8 = fp8_method self._fp8 = fp8_method
self.prefix = prefix self.prefix = prefix
self.flashinfer_mxfp4_moe_precision = ( 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): 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.rotary_embedding.utils import apply_rotary_emb
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.platforms import current_platform 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 ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -127,7 +127,7 @@ class RotaryEmbedding(MultiPlatformOp):
self._apply_rotary_emb_wrapped = apply_rotary_emb self._apply_rotary_emb_wrapped = apply_rotary_emb
# XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend # 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._forward_method = self.forward_native
self._apply_rotary_emb_wrapped = torch.compile( self._apply_rotary_emb_wrapped = torch.compile(
dynamic=True, dynamic=True,
@@ -151,7 +151,7 @@ class RotaryEmbedding(MultiPlatformOp):
# create the cache on GPU for faster initialization. This may cause # create the cache on GPU for faster initialization. This may cause
# a slight numerical difference between the HF implementation and ours. # a slight numerical difference between the HF implementation and ours.
init_device = ( 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 / ( inv_freq = 1.0 / (
base base
@@ -162,7 +162,7 @@ class RotaryEmbedding(MultiPlatformOp):
/ self.rotary_dim / 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() inv_freq = inv_freq.cuda()
return inv_freq return inv_freq
@@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import (
yarn_get_mscale_simple, yarn_get_mscale_simple,
yarn_linear_ramp_mask, 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 ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
is_cuda, is_cuda,
@@ -216,7 +216,7 @@ class MRotaryEmbedding(RotaryEmbedding):
self.register_buffer("axis_map", axis_map, persistent=False) self.register_buffer("axis_map", axis_map, persistent=False)
else: else:
self.axis_map = None 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 self._forward_method = self.forward_native
def get_cos_sin_with_position(self, positions): 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.runtime_context import get_parallel, get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL 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.async_probe import sanitize_nan_logits
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
get_bool_env_var, get_bool_env_var,
@@ -72,11 +71,9 @@ class Sampler(nn.Module):
if is_dp_attention_enabled(): if is_dp_attention_enabled():
self.tp_sync_group = get_parallel().attn_tp_group.device_group 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. # In RL on-policy mode, deterministic inference is automatically enabled.
self.enable_deterministic = ( self.enable_deterministic = get_server_args().enable_deterministic_inference
get_global_server_args().enable_deterministic_inference
)
# In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. # 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_log_softmax_logprob = self.rl_on_policy_target is not None
self.use_ascend_backend = get_server_args().sampling_backend == "ascend" 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": def create_sampler(backend: Optional[str] = None) -> "Sampler":
"""Create a sampler honoring custom backend registrations.""" """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) backend = backend or (server_args.sampling_backend if server_args else None)
if backend in _CUSTOM_SAMPLER_FACTORIES: 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.layers.moe import get_moe_a2a_backend
from sglang.srt.mem_cache.memory_pool import KVWriteLoc 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.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
@dataclass @dataclass
@@ -59,13 +58,13 @@ class ContextParallelMetadata:
def is_prefill_context_parallel_enabled(): 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(): def is_prefill_cp_in_seq_split():
return ( return (
is_prefill_context_parallel_enabled() 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: 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() 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.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.multimodal.evs import EVSEmbeddingResult 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 import flatten_nested_list, is_npu, print_warning_once
from sglang.srt.utils.stale_shm_cleanup import make_shm_name from sglang.srt.utils.stale_shm_cleanup import make_shm_name
from sglang.utils import logger from sglang.utils import logger
@@ -714,7 +714,7 @@ def _adjust_embedding_length(
f"tokens from multimodal embeddings." f"tokens from multimodal embeddings."
) )
if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding: 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: if chunked_prefill_size != -1:
logger.warning( logger.warning(
"You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill" "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) for i, seq_len in enumerate(forward_batch.extend_seq_lens_cpu)
if forward_batch.mm_inputs[i] is not None 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: 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 # Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches
input_embeds, other_info = _embed_mm_inputs_with_split( input_embeds, other_info = _embed_mm_inputs_with_split(
@@ -1119,7 +1119,7 @@ def general_mm_embed_routine(
feature = getattr(mm_item, "feature", None) feature = getattr(mm_item, "feature", None)
if isinstance(feature, torch.Tensor) and feature.is_cuda: if isinstance(feature, torch.Tensor) and feature.is_cuda:
mm_item.feature = feature.to("cpu", non_blocking=True) 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( precomputed_embeddings = getattr(
mm_item, "precomputed_embeddings", None mm_item, "precomputed_embeddings", None
) )
@@ -1756,7 +1756,7 @@ def _get_is_default_transport():
) )
_is_default_tensor_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 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. 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 return obj
if obj.mm_inputs: if obj.mm_inputs:
@@ -1859,7 +1859,7 @@ def unwrap_shm_features(obj):
Restore ShmPointerMMData wrappers back into standard torch.Tensors. Restore ShmPointerMMData wrappers back into standard torch.Tensors.
Handles both single requests and batch requests. 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 return obj
# Handle batch requests # Handle batch requests
if isinstance(obj, BaseBatchReq): if isinstance(obj, BaseBatchReq):
+12 -12
View File
@@ -106,10 +106,10 @@ from sglang.srt.observability.req_time_stats import (
DPControllerReqTimeStats, DPControllerReqTimeStats,
SchedulerReqTimeStats, 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_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import SamplingParams 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 import flatten_nested_list
from sglang.srt.utils.cuda_ipc_transport_utils import CudaIpcTensorTransportProxy 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).""" """Check if this request is prefill-only (no token generation needed)."""
# NOTE: when spec is enabled, prefill_only optimizations are disabled # 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 return self.sampling_params.max_new_tokens == 0 and spec_alg is None
@property @property
@@ -1056,7 +1056,7 @@ class Req(ReqDllmMixin):
def _cache_commit_len(self) -> int: def _cache_commit_len(self) -> int:
# Report only the prompt prefix so thinking + answer fall into the # Report only the prompt prefix so thinking + answer fall into the
# overallocated range and are reclaimed by release_kv_cache. #22373. # 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 min(self.kv_committed_len, len(self.origin_input_ids))
return self.kv_committed_len return self.kv_committed_len
@@ -2205,7 +2205,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
req.already_computed = seq_len req.already_computed = seq_len
req.is_retracted = False 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) track_entry = self._mamba_radix_cache_v2_req_prepare_for_extend(req)
mamba_track_mask_cpu.append(track_entry.track_mask) mamba_track_mask_cpu.append(track_entry.track_mask)
mamba_track_indices_cpu.append(track_entry.track_index) 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_logprob_start_lens = extend_logprob_start_lens
self.extend_input_logprob_token_ids = extend_input_logprob_token_ids 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( self.mamba_track_indices = torch.tensor(
mamba_track_indices_cpu, mamba_track_indices_cpu,
dtype=torch.int64, dtype=torch.int64,
@@ -2344,7 +2344,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self, self,
req: Req, req: Req,
) -> _MambaRadixCacheV2TrackEntry: ) -> _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: def _force_track_h(i: int) -> int:
assert i % mamba_cache_chunk_size == 0 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 # 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 # allocated yet; it will be allocated on demand at the track boundary
# in mamba_lazy_prealloc_at_boundary during prepare_for_decode. # 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 = ( req.mamba_next_track_idx = (
self.req_to_token_pool.get_mamba_ping_pong_other_idx( self.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx req.mamba_next_track_idx
@@ -2736,15 +2736,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.req_pool_indices_cpu, self.req_pool_indices_cpu,
) )
if get_global_server_args().enable_mamba_extra_buffer(): if get_server_args().enable_mamba_extra_buffer():
mamba_track_interval = get_global_server_args().mamba_track_interval mamba_track_interval = get_server_args().mamba_track_interval
if len(self.reqs) == 0: if len(self.reqs) == 0:
self.mamba_track_indices = torch.empty( self.mamba_track_indices = torch.empty(
(0,), dtype=torch.int64, device=self.device (0,), dtype=torch.int64, device=self.device
) )
else: 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) self.mamba_lazy_prealloc_at_boundary(mamba_track_interval)
set_mamba_track_indices_from_reqs(self) set_mamba_track_indices_from_reqs(self)
@@ -2932,7 +2932,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
def maybe_evict_swa(self): def maybe_evict_swa(self):
if self.tree_cache.supports_swa(): if self.tree_cache.supports_swa():
sliding_window_size = self.tree_cache.sliding_window_size sliding_window_size = self.tree_cache.sliding_window_size
server_args = get_global_server_args() server_args = get_server_args()
release_leaf_lock = ( release_leaf_lock = (
envs.SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW.get() envs.SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW.get()
@@ -56,7 +56,8 @@ from sglang.srt.mem_cache.multi_ended_allocator import (
UnifiedMambaTokenToKVPoolAllocator, UnifiedMambaTokenToKVPoolAllocator,
) )
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode 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: if TYPE_CHECKING:
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
@@ -185,7 +186,7 @@ class SchedulePolicy:
if ( if (
not isinstance(policy, CacheAwarePolicy) not isinstance(policy, CacheAwarePolicy)
and self.tree_cache.supports_fast_match_prefix() 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: for r in waiting_queue:
match_prefix_for_req(self.tree_cache, r, include_req=True) 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.parser.reasoning_parser import ReasoningParser
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.plugins import load_plugins 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.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.session.session_controller import SessionController
from sglang.srt.speculative.dflash_utils import validate_dflash_request from sglang.srt.speculative.dflash_utils import validate_dflash_request
from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec 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( self.min_free_slots_delayer = MinFreeSlotsDelayer(
min_free_slots=min_free_slots min_free_slots=min_free_slots
) )
if not get_global_server_args().pp_max_micro_batch_size: if not get_server_args().pp_max_micro_batch_size:
get_global_server_args().override( get_server_args().override(
"scheduler.pp_max_micro_batch_size_default", "scheduler.pp_max_micro_batch_size_default",
pp_max_micro_batch_size=max( pp_max_micro_batch_size=max(
self.max_running_requests // self.ps.pp_size, 1 self.max_running_requests // self.ps.pp_size, 1
@@ -2730,7 +2730,7 @@ class Scheduler(
return ret return ret
def get_num_allocatable_reqs(self, running_bs): 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()) res = min(res, self.req_to_token_pool.available_size())
return res return res
@@ -3766,7 +3766,7 @@ class Scheduler(
return success return success
def get_internal_state(self, recv_req: GetInternalStateReq): 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["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput
ret["memory_usage"] = { ret["memory_usage"] = {
"weight": round(self.tp_worker.model_runner.weight_load_mem_usage, 2), "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_accept_tokens = (
self.metrics_reporter.spec_total_num_forward_ct self.metrics_reporter.spec_total_num_forward_ct
) = 0 ) = 0
for k, v in server_args_dict.items(): get_server_args().override(source="update_server_args", **server_args_dict)
setattr(get_global_server_args(), k, v) logger.info(f"Global server args updated! {get_server_args()=}")
logger.info(f"Global server args updated! {get_global_server_args()=}")
server_args = dict(vars(get_global_server_args())) server_args = dict(vars(get_server_args()))
# This field is not serializable. # This field is not serializable.
server_args.pop("model_config", None) server_args.pop("model_config", None)
return SetInternalStateReqOutput( return SetInternalStateReqOutput(
@@ -25,7 +25,7 @@ from sglang.srt.mem_cache.common import (
maybe_cache_unfinished_req, maybe_cache_unfinished_req,
release_kv_cache, 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.indexer_topk import get_global_indexer_capturer
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
@@ -848,7 +848,7 @@ class SchedulerBatchResultProcessor:
prepare_release(req) prepare_release(req)
is_insert = ( is_insert = (
req.mamba_lazy_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 else True
) )
release_kv_cache(req, self.tree_cache, is_insert=is_insert) 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: if req.mamba_ping_pong_track_buffer is None:
return 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( at_boundary, track_seqlen = self._mamba_check_track_boundary(
req, batch, result, i req, batch, result, i
) )
@@ -915,7 +915,7 @@ class SchedulerBatchResultProcessor:
For spec decode, the boundary is detected by comparing the For spec decode, the boundary is detected by comparing the
accepted seq_len range against interval boundaries. 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 batch.spec_algorithm.is_none():
if req.kv_committed_len % interval == 0: if req.kv_committed_len % interval == 0:
@@ -18,7 +18,7 @@ import torch
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
from sglang.srt.model_executor.forward_batch_info import ForwardMode 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 import is_mps, is_npu
from sglang.srt.utils.profile_merger import ProfileMerger from sglang.srt.utils.profile_merger import ProfileMerger
from sglang.srt.utils.profile_utils import ProfileManager from sglang.srt.utils.profile_utils import ProfileManager
@@ -245,7 +245,7 @@ class SchedulerProfilerManager:
self.profile_in_progress = True self.profile_in_progress = True
if "CUDA_PROFILER" in activities: 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() torch.cuda.cudart().cudaProfilerStart()
self.profile_in_progress = True self.profile_in_progress = True
@@ -355,7 +355,7 @@ class SchedulerProfilerManager:
torch.cuda.memory._record_memory_history(enabled=None) torch.cuda.memory._record_memory_history(enabled=None)
if "CUDA_PROFILER" in self.profiler_activities: 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() torch.cuda.cudart().cudaProfilerStop()
merge_message = self._merge_profile_traces() merge_message = self._merge_profile_traces()
@@ -216,9 +216,9 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
) )
def init_metrics_collector(self): 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__} labels = {"cache_type": self.__class__.__name__}
if server_args.extra_metric_labels: if server_args.extra_metric_labels:
labels.update(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, write_req_to_token_pool_triton,
) )
from sglang.srt.runtime_context import get_server_args 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 import is_cuda, is_hip, is_npu, support_triton
from sglang.srt.utils.common import ceil_align, is_pin_memory_available 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: def get_alloc_len_per_decode(server_args: Optional[ServerArgs] = None) -> int:
if server_args is None: if server_args is None:
server_args = get_global_server_args() server_args = get_server_args()
if server_args.speculative_algorithm is None: if server_args.speculative_algorithm is None:
return 1 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 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 # 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. # 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.token_to_kv_pool_allocator.page_size
return batch.tree_cache.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() 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 page_size = global_server_args.page_size
spec_algo = global_server_args.speculative_algorithm 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.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
from sglang.srt.mem_cache.memory_pool import KVCache 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 from sglang.srt.utils import ceil_div, is_hip
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -276,7 +276,7 @@ class DeepSeekV4IndexerPool(KVCache):
end_layer, end_layer,
) )
self.index_head_dim = index_head_dim 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() self._create_buffer()
@@ -569,7 +569,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.swa_kv_pool = None self.swa_kv_pool = None
self.c4_kv_pool = None self.c4_kv_pool = None
self.c128_kv_pool = None self.c128_kv_pool = None
server_args = get_global_server_args() server_args = get_server_args()
spec_extra = ( spec_extra = (
(server_args.speculative_num_draft_tokens - 1) (server_args.speculative_num_draft_tokens - 1)
if server_args.speculative_algorithm is not None 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 self.full_to_swa_index_mapping = full_to_swa_index_mapping
def get_ring_size(self, compress_ratio: int) -> int: 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 is_speculative = server_args.speculative_algorithm is not None
return get_compress_state_ring_size(compress_ratio, is_speculative) return get_compress_state_ring_size(compress_ratio, is_speculative)
@@ -1030,9 +1030,9 @@ class HiMambaRadixCache(MambaRadixCache):
node_update = node_update.parent node_update = node_update.parent
if len(value) > best_value_len: 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 = ( mamba_cache_chunk_aligned_seqlen = (
sum(len(v) for v in value) // mamba_cache_chunk_size sum(len(v) for v in value) // mamba_cache_chunk_size
) * mamba_cache_chunk_size ) * mamba_cache_chunk_size
@@ -1272,10 +1272,10 @@ class HiMambaRadixCache(MambaRadixCache):
} }
if extra_metric_labels: if extra_metric_labels:
labels.update(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( storage_cls = resolve_collector_class(
get_global_server_args(), get_server_args(),
STAT_LOGGER_ROLE_STORAGE, STAT_LOGGER_ROLE_STORAGE,
StorageMetricsCollector, StorageMetricsCollector,
) )
+2 -2
View File
@@ -338,10 +338,10 @@ class HiRadixCache(RadixCache):
labels.update(extra_metric_labels) labels.update(extra_metric_labels)
existing_collector = getattr(self, "storage_metrics_collector", None) existing_collector = getattr(self, "storage_metrics_collector", None)
if existing_collector is 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( storage_cls = resolve_collector_class(
get_global_server_args(), get_server_args(),
STAT_LOGGER_ROLE_STORAGE, STAT_LOGGER_ROLE_STORAGE,
StorageMetricsCollector, StorageMetricsCollector,
) )
@@ -313,10 +313,10 @@ def maybe_init_int8_mamba_checkpoint_pool(
allocating, so an oversized ``--int8-mamba-ckpt-size`` fails with an actionable allocating, so an oversized ``--int8-mamba-ckpt-size`` fails with an actionable
message instead of a cryptic mid-allocation CUDA OOM. 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: try:
_sa = get_global_server_args() _sa = get_server_args()
except ValueError: except ValueError:
# Some unit-test / mock runners construct HybridReqToTokenPool directly # Some unit-test / mock runners construct HybridReqToTokenPool directly
# without a global server-args context. The int8 checkpoint pool is opt-in # 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.radix_cache import RadixKey
from sglang.srt.mem_cache.utils import split_node_hash_value 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: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req 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.req_to_token_pool: HybridReqToTokenPool = params.req_to_token_pool
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator 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.page_size = params.page_size
self.disable = params.disable self.disable = params.disable
@@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchResult, MatchResult,
) )
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode 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: try:
from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector
@@ -108,7 +108,7 @@ class LMCRadixCache(RadixCache):
): ):
super().__init__(params) 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() kvcache = self.token_to_kv_pool_allocator.get_kvcache()
connector_kwargs = dict( connector_kwargs = dict(
@@ -438,7 +438,7 @@ class LMCRadixCache(RadixCache):
self.lmcache_connector.end_session(req.rid) self.lmcache_connector.end_session(req.rid)
return return
global_server_args = get_global_server_args() global_server_args = get_server_args()
topk = global_server_args.speculative_eagle_topk topk = global_server_args.speculative_eagle_topk
enable_kv_committed_len = topk is None or topk == 1 enable_kv_committed_len = topk is None or topk == 1
if enable_kv_committed_len: if enable_kv_committed_len:
@@ -26,7 +26,7 @@ from sglang.srt.mem_cache.unified_cache_components.tree_component import (
TreeComponent, TreeComponent,
get_and_increase_time_counter, 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: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req 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 # states. We temporarily skip branching-state fill in that mode and can
# add a HiCache-aware branching policy later. # add a HiCache-aware branching policy later.
if self.cache.cache_controller is None and len(value_chunks) > best_value_len: 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 = ( aligned_seqlen = (
sum(len(v) for v in value_chunks) // chunk_size sum(len(v) for v in value_chunks) // chunk_size
) * chunk_size ) * chunk_size
@@ -2280,10 +2280,10 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
labels.update(extra_metric_labels) labels.update(extra_metric_labels)
existing_collector = self.storage_metrics_collector existing_collector = self.storage_metrics_collector
if existing_collector is 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( storage_cls = resolve_collector_class(
get_global_server_args(), get_server_args(),
STAT_LOGGER_ROLE_STORAGE, STAT_LOGGER_ROLE_STORAGE,
StorageMetricsCollector, 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 Module-level imports are pure stdlib — no torch / sglang.srt deps — so
ServerArgs can import everything here without pulling in backend 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. 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 """True if cuda_graph_config[phase].backend == backend on the
global server args. Returns False if the global server args have not global server args. Returns False if the global server args have not
been initialized yet (e.g. unit tests, early startup).""" 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: try:
server_args = get_global_server_args() server_args = get_server_args()
except ValueError: except ValueError:
return False return False
cfg = server_args.cuda_graph_config cfg = server_args.cuda_graph_config
@@ -49,8 +49,7 @@ from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
ForwardBatchDeepSeekMHAMixin, ForwardBatchDeepSeekMHAMixin,
) )
from sglang.srt.model_executor.triton_ops.position import compute_position_triton from sglang.srt.model_executor.triton_ops.position import compute_position_triton
from sglang.srt.runtime_context import get_parallel 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 ( from sglang.srt.utils import (
is_cuda, is_cuda,
is_hip, is_hip,
@@ -910,7 +909,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# --enable-mis: every request must carry delimiter indices (the score # --enable-mis: every request must carry delimiter indices (the score
# endpoint always produces MIS-structured requests; consumers index # endpoint always produces MIS-structured requests; consumers index
# without None-checking). # 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 r.multi_item_delimiter_indices is not None for r in batch.reqs
): ):
assert all( assert all(
@@ -1073,7 +1072,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# 3 * N # 3 * N
if ( if (
mm_input is None 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( mrope_positions_list[batch_idx] = torch.full(
(3, 1), (3, 1),
@@ -1092,7 +1091,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
) )
if ( if (
mm_input is None 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 # text only
mrope_positions = torch.tensor( mrope_positions = torch.tensor(
@@ -780,9 +780,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
torchao_applied = getattr(self.model, "torchao_applied", False) torchao_applied = getattr(self.model, "torchao_applied", False)
# In layered loading, torchao may have been applied # In layered loading, torchao may have been applied
if not torchao_applied: if not torchao_applied:
apply_torchao_config_to_model( apply_torchao_config_to_model(self.model, get_server_args().torchao_config)
self.model, get_global_server_args().torchao_config
)
# Apply torch TP if the model supports it # Apply torch TP if the model supports it
supports_torch_tp = getattr(self.model, "supports_torch_tp", False) supports_torch_tp = getattr(self.model, "supports_torch_tp", False)
@@ -1007,7 +1005,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
set_global_experts_capturer( set_global_experts_capturer(
RoutedExpertsCapturer.create( RoutedExpertsCapturer.create(
enable=get_global_server_args().enable_return_routed_experts, enable=get_server_args().enable_return_routed_experts,
model_config=self.model_config, model_config=self.model_config,
num_fused_shared_experts=num_fused_shared_experts, num_fused_shared_experts=num_fused_shared_experts,
num_tokens=self.max_total_num_tokens + self.page_size, num_tokens=self.max_total_num_tokens + self.page_size,
@@ -1017,7 +1015,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
) )
def init_indexer_capturer(self): 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 # Producer wiring is CUDA-only (Indexer.forward_cuda + MLA skip_topk
# path); other backends would create a capturer but never feed it. # path); other backends would create a capturer but never feed it.
if enable and self.device != "cuda": if enable and self.device != "cuda":
@@ -1709,8 +1707,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
else: else:
# Load the missing weights from disk # Load the missing weights from disk
self.update_weights_from_disk( self.update_weights_from_disk(
get_global_server_args().model_path, get_server_args().model_path,
get_global_server_args().load_format, get_server_args().load_format,
weight_name_filter=weight_name_filter, 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, get_remote_instance_transfer_engine_info_per_rank,
register_memory_region, 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 from sglang.srt.utils import get_available_gpu_memory
# Try to import accelerate (optional dependency) # Try to import accelerate (optional dependency)
@@ -481,7 +481,7 @@ class DefaultModelLoader(BaseModelLoader):
else: else:
hf_folder = model_name_or_path 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: if server_args and server_args.model_checksum is not None:
from sglang.srt.utils.model_file_verifier import verify from sglang.srt.utils.model_file_verifier import verify
@@ -567,7 +567,7 @@ class DefaultModelLoader(BaseModelLoader):
hf_weights_files, hf_weights_files,
) )
elif use_safetensors: 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_disable_mmap = server_args.weight_loader_disable_mmap
weight_loader_prefetch = server_args.weight_loader_prefetch_checkpoints weight_loader_prefetch = server_args.weight_loader_prefetch_checkpoints
prefetch_num_threads = server_args.weight_loader_prefetch_num_threads prefetch_num_threads = server_args.weight_loader_prefetch_num_threads
@@ -866,9 +866,9 @@ class LayeredModelLoader(DefaultModelLoader):
device_config: DeviceConfig, device_config: DeviceConfig,
) -> nn.Module: ) -> nn.Module:
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model 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) target_device = torch.device(device_config.device)
quant_config = _get_quantization_config(model_config, self.load_config) 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: if server_args and server_args.model_checksum is not None:
from sglang.srt.utils.model_file_verifier import verify 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, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream 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 from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
LoraConfig = None LoraConfig = None
@@ -209,7 +208,7 @@ class BailingMoESparseMoeBlock(nn.Module):
self.router_dtype = torch.bfloat16 self.router_dtype = torch.bfloat16
# TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now # 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 # check group topk
self.num_expert_group = getattr(config, "n_group", 0) self.num_expert_group = getattr(config, "n_group", 0)
self.topk_group = getattr(config, "topk_group", 0) self.topk_group = getattr(config, "topk_group", 0)
@@ -224,7 +223,7 @@ class BailingMoESparseMoeBlock(nn.Module):
self.use_grouped_topk = False self.use_grouped_topk = False
self.num_experts = ( 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( 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.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream 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 ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
add_prefix, add_prefix,
@@ -534,7 +533,7 @@ class BailingMoELinearAttention(nn.Module):
base=self.rope_theta, base=self.rope_theta,
rope_scaling=config.rope_scaling, rope_scaling=config.rope_scaling,
is_neox_style=True, is_neox_style=True,
device=get_global_server_args().device, device=get_server_args().device,
dtype=torch.float32, dtype=torch.float32,
) )
@@ -695,7 +694,7 @@ class BailingMoEAttention(nn.Module):
max_position=self.max_position_embeddings, max_position=self.max_position_embeddings,
base=self.rope_theta, base=self.rope_theta,
rope_scaling=config.rope_scaling, rope_scaling=config.rope_scaling,
device=get_global_server_args().device, device=get_server_args().device,
) )
self.attn = RadixAttention( self.attn = RadixAttention(
self.num_heads, 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.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel 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 from sglang.srt.utils import add_prefix
BertConfig = None BertConfig = None
@@ -367,9 +366,7 @@ class BertModel(nn.Module):
prefix=add_prefix("encoder", prefix), prefix=add_prefix("encoder", prefix),
) )
pooling_type = ( pooling_type = (
PoolingType.CLS PoolingType.CLS if get_server_args().is_embedding else PoolingType.LAST
if get_global_server_args().is_embedding
else PoolingType.LAST
) )
self.pooler = ( self.pooler = (
BertPooler(config) BertPooler(config)
@@ -11,7 +11,7 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods
AttnForwardMethod, AttnForwardMethod,
) )
from sglang.srt.models.deepseek_common.utils import _is_hip 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 from sglang.srt.utils import use_intel_amx_backend
MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"] 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): def handle_attention_fa3(attn, forward_batch):
# when deterministic inference is enabled, use 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) return _dispatch_mla_subtype(attn, forward_batch)
else: else:
return _handle_attention_backend(attn, forward_batch, "fa3") return _handle_attention_backend(attn, forward_batch, "fa3")
@@ -183,7 +183,7 @@ def handle_attention_triton(attn, forward_batch):
return AttnForwardMethod.MLA return AttnForwardMethod.MLA
# when deterministic inference is enabled, use 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) return _dispatch_mla_subtype(attn, forward_batch)
if ( if (
@@ -31,7 +31,7 @@ from sglang.srt.models.deepseek_common.utils import (
_use_aiter_bpreshuffle_gfx95, _use_aiter_bpreshuffle_gfx95,
_use_aiter_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 from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2
_use_fp8_prefill_attn = ( _use_fp8_prefill_attn = (
@@ -115,7 +115,7 @@ class DeepseekMHAForwardMixin:
def init_mha_forward(self: DeepseekV2AttentionMLA): def init_mha_forward(self: DeepseekV2AttentionMLA):
self.disable_chunked_prefix_cache = ( 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 # TODO: Design a finer way to determine the threshold
@@ -279,8 +279,8 @@ class DeepseekMHAForwardMixin:
self.use_dsa self.use_dsa
and self.kv_cache_dtype == "fp8_e4m3" and self.kv_cache_dtype == "fp8_e4m3"
and ( and (
not get_global_server_args().dsa_decode_backend == "trtllm" not get_server_args().dsa_decode_backend == "trtllm"
or not get_global_server_args().dsa_prefill_backend == "trtllm" or not get_server_args().dsa_prefill_backend == "trtllm"
) )
): ):
# FP8 path: dequantize DSA-specific FP8 format to BF16 # 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_bpreshuffle_gfx95,
_use_aiter_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 ( from sglang.srt.state_capturer.indexer_topk import (
maybe_capture_indexer_topk, maybe_capture_indexer_topk,
) )
@@ -178,7 +178,7 @@ def _should_defer_dsa_cp_kv_gather(
class DeepseekMLAForwardMixin: class DeepseekMLAForwardMixin:
def init_mla_forward(self: DeepseekV2AttentionMLA): def init_mla_forward(self: DeepseekV2AttentionMLA):
self.flashinfer_mla_disable_ragged = ( self.flashinfer_mla_disable_ragged = (
get_global_server_args().flashinfer_mla_disable_ragged get_server_args().flashinfer_mla_disable_ragged
) )
def should_run_indexer( def should_run_indexer(
@@ -1002,8 +1002,8 @@ class DeepseekMLAForwardMixin:
""" """
if self.current_attention_backend in ("dsa", "nsa"): if self.current_attention_backend in ("dsa", "nsa"):
return ( return (
get_global_server_args().dsa_decode_backend == "trtllm" get_server_args().dsa_decode_backend == "trtllm"
or get_global_server_args().dsa_prefill_backend == "trtllm" or get_server_args().dsa_prefill_backend == "trtllm"
) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn ) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn
return ( return (
@@ -1020,7 +1020,7 @@ class DeepseekMLAForwardMixin:
""" """
Check if we should skip rope and use fused rope+cache path for TileLang DSA on gfx95. 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 ( return (
_use_aiter_gfx95 _use_aiter_gfx95
and self.current_attention_backend in ("dsa", "nsa") 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.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_parallel, get_server_args 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 from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -119,7 +118,7 @@ class DeepseekModelNextN(nn.Module):
self.rot_weight = None self.rot_weight = None
if _is_npu: 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): if os.path.isfile(rot_weight_path):
self.rot_weight = load_file(rot_weight_path) self.rot_weight = load_file(rot_weight_path)
self.rot_weight = self.rot_weight["rot.weight"].npu() self.rot_weight = self.rot_weight["rot.weight"].npu()
@@ -132,8 +131,8 @@ class DeepseekModelNextN(nn.Module):
layer_name = "decoder" layer_name = "decoder"
if _is_npu and ( if _is_npu and (
get_global_server_args().speculative_draft_model_path get_server_args().speculative_draft_model_path
== get_global_server_args().model_path == get_server_args().model_path
): ):
layer_name = "layers." + str(config.num_hidden_layers) 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, _use_aiter_gfx95,
) )
from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args 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.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -484,7 +483,7 @@ class MoEGate(nn.Module):
True, # is_vnni 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) return F.linear(hidden_states, self.weight, None)
if ( if (
@@ -621,7 +620,7 @@ class DeepseekV2MoE(nn.Module):
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=num_experts_for_moe 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, num_fused_shared_experts=self.num_fused_shared_experts,
top_k=top_k_for_moe, top_k=top_k_for_moe,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
@@ -794,8 +793,7 @@ class DeepseekV2MoE(nn.Module):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts config.n_routed_experts + get_server_args().ep_num_redundant_experts
+ get_global_server_args().ep_num_redundant_experts
) )
self.renormalize = config.norm_topk_prob self.renormalize = config.norm_topk_prob
self.topk_group = config.topk_group self.topk_group = config.topk_group
@@ -836,7 +834,7 @@ class DeepseekV2MoE(nn.Module):
self, hidden_states: torch.Tensor, server_args=None self, hidden_states: torch.Tensor, server_args=None
) -> bool: ) -> bool:
if server_args is None: if server_args is None:
server_args = get_global_server_args() server_args = get_server_args()
return ( return (
_enable_pcg_dsv2_dual_stream _enable_pcg_dsv2_dual_stream
and (is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph()) 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: 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): if self._can_dual_stream_graph(hidden_states, server_args):
return dsv2_flashinfer_moe_dual_stream_graph( return dsv2_flashinfer_moe_dual_stream_graph(
hidden_states, hidden_states,
@@ -935,7 +933,7 @@ class DeepseekV2MoE(nn.Module):
# into the decode CUDA graph and replays from null. # into the decode CUDA graph and replays from null.
current_stream = torch.cuda.current_stream() current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream) self.alt_stream.wait_stream(current_stream)
server_args = get_global_server_args() server_args = get_server_args()
dispatch_info = ( dispatch_info = (
ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id) ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id)
if server_args.enable_eplb if server_args.enable_eplb
@@ -1032,7 +1030,7 @@ class DeepseekV2MoE(nn.Module):
self.shared_experts.gate_up_proj self.shared_experts.gate_up_proj
): ):
return self.forward_cpu(hidden_states, should_allreduce_fusion) return self.forward_cpu(hidden_states, should_allreduce_fusion)
server_args = get_global_server_args() server_args = get_server_args()
dispatch_info = ( dispatch_info = (
ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id) ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id)
if server_args.enable_eplb if server_args.enable_eplb
@@ -1599,7 +1597,7 @@ class DeepseekV2AttentionMLA(
self.scaling = self.qk_head_dim**-0.5 self.scaling = self.qk_head_dim**-0.5
self.rope_theta = rope_theta self.rope_theta = rope_theta
self.max_position_embeddings = max_position_embeddings 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 # NOTE modification to rope_scaling must be done early enough, b/c e.g. Indexer needs it
if rope_scaling: if rope_scaling:
@@ -1711,7 +1709,7 @@ class DeepseekV2AttentionMLA(
base=rope_theta, base=rope_theta,
rope_scaling=rope_scaling, rope_scaling=rope_scaling,
is_neox_style=is_neox_style, 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): 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 # Determine attention backend name for current forward batch: prefer the
# name stamped per-runner on the backend object, else resolve from server args. # name stamped per-runner on the backend object, else resolve from server args.
backend = get_attn_backend() 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() default_prefill_str, default_decode_str = server_args.get_attention_backends()
prefill_backend_str = ( prefill_backend_str = (
backend.prefill_attention_backend_str or default_prefill_str backend.prefill_attention_backend_str or default_prefill_str
@@ -2071,7 +2069,7 @@ class DeepseekV2DecoderLayer(nn.Module):
rope_scaling = config.rope_scaling rope_scaling = config.rope_scaling
max_position_embeddings = config.max_position_embeddings max_position_embeddings = config.max_position_embeddings
self.speculative_algorithm = SpeculativeAlgorithm.from_string( 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.dsa_enable_prefill_cp = dsa_enable_prefill_cp
self.mla_enable_prefill_cp = mla_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, architecture: str = "DeepseekV3ForCausalLM"
): ):
self.num_fused_shared_experts = 0 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: if get_server_args().disable_shared_experts_fusion:
return return
+3 -3
View File
@@ -135,7 +135,7 @@ if _is_xpu:
else: else:
from sglang.srt.layers.mhc import hc_split_sinkhorn, mhc_fused_post_pre, npu_hc_pre 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 ( from sglang.srt.utils import (
LazyValue, LazyValue,
add_prefix, add_prefix,
@@ -346,7 +346,7 @@ class MQALayer(nn.Module):
base=rope_base, base=rope_base,
rope_scaling=rope_scaling, rope_scaling=rope_scaling,
is_neox_style=False, 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 from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis
@@ -2210,7 +2210,7 @@ class DeepseekV4ForCausalLM(nn.Module):
return return
disable_reason = None 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: if self.config.n_shared_experts != 1:
raise ValueError( raise ValueError(
"DeepSeek V4 shared-experts fusion expects exactly one shared " "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_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream 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 from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -171,8 +170,7 @@ class ExaoneMoESparseMoEBlock(nn.Module):
) )
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=config.num_experts num_experts=config.num_experts + get_server_args().ep_num_redundant_experts,
+ get_global_server_args().ep_num_redundant_experts,
top_k=config.num_experts_per_tok, top_k=config.num_experts_per_tok,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
intermediate_size=config.moe_intermediate_size, intermediate_size=config.moe_intermediate_size,
@@ -213,7 +211,7 @@ class ExaoneMoESparseMoEBlock(nn.Module):
if get_moe_a2a_backend().is_deepep(): if get_moe_a2a_backend().is_deepep():
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.num_experts + get_global_server_args().ep_num_redundant_experts config.num_experts + get_server_args().ep_num_redundant_experts
) )
self.top_k = config.num_experts_per_tok self.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 ( from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
) )
from sglang.srt.runtime_context import get_parallel 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, make_layers from sglang.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -255,8 +254,7 @@ class Gemma4MoE(nn.Module):
experts_type = get_moe_impl_class(quant_config) experts_type = get_moe_impl_class(quant_config)
self.experts = experts_type( self.experts = experts_type(
num_experts=config.num_experts num_experts=config.num_experts + get_server_args().ep_num_redundant_experts,
+ get_global_server_args().ep_num_redundant_experts,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
intermediate_size=config.moe_intermediate_size, intermediate_size=config.moe_intermediate_size,
layer_id=layer_id, layer_id=layer_id,
@@ -789,7 +787,7 @@ class Gemma4TextModel(PreTrainedModel):
# combination until the runner becomes schema-aware; users can run # combination until the runner becomes schema-aware; users can run
# PP + PLE eagerly with --disable-cuda-graph. # PP + PLE eagerly with --disable-cuda-graph.
if self.pp_group.world_size > 1 and self.hidden_size_per_layer_input > 0: 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: if sa is not None and not sa.disable_cuda_graph:
raise ValueError( raise ValueError(
"Pipeline parallelism is currently incompatible with " "Pipeline parallelism is currently incompatible with "
+2 -2
View File
@@ -181,9 +181,9 @@ class Gemma4VisionAttention(nn.Module):
@staticmethod @staticmethod
def _select_backend() -> str: def _select_backend() -> str:
"""Mirror VisionAttention._determine_attention_backend for consistency.""" """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: if override is not None:
return override return override
if is_cuda(): if is_cuda():
+1 -3
View File
@@ -88,7 +88,6 @@ from sglang.srt.runtime_context import (
get_server_args, get_server_args,
get_stream, get_stream,
) )
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
cpu_has_amx_support, cpu_has_amx_support,
@@ -530,8 +529,7 @@ class Glm4MoeSparseMoeBlock(nn.Module):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts config.n_routed_experts + get_server_args().ep_num_redundant_experts
+ get_global_server_args().ep_num_redundant_experts
) )
self.renormalize = config.norm_topk_prob self.renormalize = config.norm_topk_prob
self.topk_group = config.topk_group 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_common.utils import _is_cuda, _use_aiter
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream 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 ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
LazyValue, LazyValue,
@@ -216,7 +215,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=config.n_routed_experts num_experts=config.n_routed_experts
+ self.num_fused_shared_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, num_fused_shared_experts=self.num_fused_shared_experts,
top_k=config.num_experts_per_tok + self.num_fused_shared_experts, top_k=config.num_experts_per_tok + self.num_fused_shared_experts,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
@@ -284,8 +283,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts config.n_routed_experts + get_server_args().ep_num_redundant_experts
+ get_global_server_args().ep_num_redundant_experts
) )
self.renormalize = config.norm_topk_prob self.renormalize = config.norm_topk_prob
self.topk_group = config.topk_group self.topk_group = config.topk_group
@@ -36,7 +36,6 @@ from sglang.srt.models.glm4_moe_lite import (
Glm4MoeLiteForCausalLM, Glm4MoeLiteForCausalLM,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args 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 from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -140,10 +139,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_parallel().tp_size self.tp_size = get_parallel().tp_size
if ( if is_npu() and get_server_args().speculative_draft_model_quantization is None:
is_npu()
and get_global_server_args().speculative_draft_model_quantization is None
):
quant_config = None quant_config = None
self.quant_config = quant_config 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.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
from sglang.srt.runtime_context import get_parallel, get_server_args 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 from sglang.srt.utils import add_prefix, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -126,10 +125,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_parallel().tp_size self.tp_size = get_parallel().tp_size
if ( if is_npu() and get_server_args().speculative_draft_model_quantization is None:
is_npu()
and get_global_server_args().speculative_draft_model_quantization is None
):
quant_config = None quant_config = None
self.quant_config = quant_config 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.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.glm4 import Glm4Model from sglang.srt.models.glm4 import Glm4Model
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.runtime_context import get_parallel 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 from sglang.srt.utils import add_prefix, is_npu
from sglang.srt.utils.hf_transformers_utils import get_processor 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.pp_group = get_pp_group()
self.config = config 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) vision_utils.update_vit_attn_dummy_heads_config(self.config)
self.visual = Glm4vVisionModel( self.visual = Glm4vVisionModel(
config.vision_config, 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.glm4_moe import Glm4MoeModel
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel
from sglang.srt.runtime_context import get_parallel, get_server_args 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 import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -42,7 +41,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config 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) vision_utils.update_vit_attn_dummy_heads_config(self.config)
self.tp_size = get_parallel().tp_size self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
+2 -2
View File
@@ -50,7 +50,7 @@ from sglang.srt.models.glm4v import (
Glm4vVisionModel, Glm4vVisionModel,
Glm4vVisionPatchEmbed, 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 import add_prefix
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -272,7 +272,7 @@ class GlmOcrForConditionalGeneration(Glm4vForConditionalGeneration):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config 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( self.visual = GlmOcrVisionModel(
vision_config=config.vision_config, vision_config=config.vision_config,
text_config=config.text_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, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args 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 ( from sglang.srt.utils import (
LazyValue, LazyValue,
add_prefix, add_prefix,
@@ -228,7 +227,7 @@ class GptOssSparseMoeBlock(nn.Module):
self.experts = experts_type( self.experts = experts_type(
num_experts=config.num_local_experts 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, top_k=config.num_experts_per_tok,
layer_id=layer_id, layer_id=layer_id,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
+2 -3
View File
@@ -40,8 +40,7 @@ from sglang.srt.multimodal.internvl_vit_cuda_graph_runner import (
InternViTCudaGraphRunner, InternViTCudaGraphRunner,
) )
from sglang.srt.multimodal.mm_utils import run_dp_sharded_vision_model from sglang.srt.multimodal.mm_utils import run_dp_sharded_vision_model
from sglang.srt.runtime_context import get_parallel 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 is_cuda from sglang.srt.utils import is_cuda
from sglang.utils import logger from sglang.utils import logger
@@ -498,7 +497,7 @@ class InternVLChatModel(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.config = config 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 self.quant_config = quant_config
vision_utils.update_vit_attn_dummy_heads_config(self.config) vision_utils.update_vit_attn_dummy_heads_config(self.config)
image_size = config.force_image_size or config.vision_config.image_size 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.kimi_vl_moonvit import MLP2
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model 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 from sglang.srt.utils import add_prefix, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -633,7 +633,7 @@ class KimiK25ForConditionalGeneration(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
self.quant_config = quant_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 # Create vision tower
self.vision_tower = MoonViT3dPretrainedModel( self.vision_tower = MoonViT3dPretrainedModel(
config.vision_config, 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.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_parallel, get_server_args 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 from sglang.srt.utils import LazyValue, add_prefix, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -161,8 +160,7 @@ class LagunaMoE(nn.Module):
self.gate = LagunaMoEGate(config, prefix=add_prefix("gate", prefix)) self.gate = LagunaMoEGate(config, prefix=add_prefix("gate", prefix))
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=config.num_experts num_experts=config.num_experts + get_server_args().ep_num_redundant_experts,
+ get_global_server_args().ep_num_redundant_experts,
top_k=config.num_experts_per_tok, top_k=config.num_experts_per_tok,
layer_id=layer_id, layer_id=layer_id,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
+2 -3
View File
@@ -77,7 +77,6 @@ from sglang.srt.models.utils import (
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream 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 ( from sglang.srt.utils import (
add_prefix, add_prefix,
is_cuda, is_cuda,
@@ -220,7 +219,7 @@ class LLaDA2MoeSparseMoeBlock(nn.Module):
self.router_dtype = torch.bfloat16 self.router_dtype = torch.bfloat16
# TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now # 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 # check group topk
self.num_expert_group = getattr(config, "n_group", 0) self.num_expert_group = getattr(config, "n_group", 0)
self.topk_group = getattr(config, "topk_group", 0) self.topk_group = getattr(config, "topk_group", 0)
@@ -235,7 +234,7 @@ class LLaDA2MoeSparseMoeBlock(nn.Module):
self.use_grouped_topk = False self.use_grouped_topk = False
self.num_experts = ( 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( 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_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.llama import LlamaDecoderLayer, LlamaForCausalLM, LlamaMLP 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): class LlamaDecoderLayer(LlamaDecoderLayer):
@@ -258,7 +258,7 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM):
# Cache draft SWA size from server args once; consumed both by the post-init # 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. # attention patch below and by `get_attention_sliding_window_size` later.
self._draft_window_size: Optional[int] = ( self._draft_window_size: Optional[int] = (
get_global_server_args().speculative_draft_window_size get_server_args().speculative_draft_window_size
) )
self.model = LlamaModel( 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 transformers.models.qwen2.modeling_qwen2 import Qwen2Model
from sglang.srt.layers.quantization.base_config import QuantizationConfig 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 from sglang.srt.utils import is_cuda
if is_cuda(): if is_cuda():
@@ -1223,7 +1223,7 @@ class AudioEncoderMixin:
else: else:
raise ValueError(f"Invalid projection layers: {config.projection_layers}") 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): if not os.path.isdir(model_path):
from huggingface_hub import snapshot_download 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_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig
from sglang.srt.runtime_context import get_parallel, get_server_args 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 ( from sglang.srt.utils import (
LazyValue, LazyValue,
add_prefix, add_prefix,
@@ -253,7 +252,7 @@ class MiMoV2MoE(nn.Module):
experts_type = get_moe_impl_class(quant_config) experts_type = get_moe_impl_class(quant_config)
self.experts = experts_type( self.experts = experts_type(
num_experts=config.n_routed_experts 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, top_k=config.num_experts_per_tok,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
intermediate_size=config.moe_intermediate_size, intermediate_size=config.moe_intermediate_size,
@@ -288,8 +287,7 @@ class MiMoV2MoE(nn.Module):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts config.n_routed_experts + get_server_args().ep_num_redundant_experts
+ get_global_server_args().ep_num_redundant_experts
) )
self.renormalize = config.norm_topk_prob self.renormalize = config.norm_topk_prob
self.topk_group = config.topk_group 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.layernorm import RMSNorm
from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.models.qwen2_5_vl import Qwen2_5_VisionPatchMerger, Qwen2_5_VLMLP 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 from sglang.srt.utils import add_prefix
@@ -232,7 +232,7 @@ class MiMoVisionTransformer(nn.Module):
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
super().__init__() 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 self.vit_window_attn_types = vision_config.vit_window_attn_types
patch_size: int = vision_config.patch_size patch_size: int = vision_config.patch_size
temporal_patch_size: int = vision_config.temporal_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, maybe_remap_kv_scale_name,
narrow_padded_param_and_loaded_weight, narrow_padded_param_and_loaded_weight,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
# get_bool_env_var is defined in sglang.srt.utils.common, not sglang.srt.distributed. # 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 # 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) props = torch.cuda.get_device_properties(device)
# probe the maximum tokens for one prefill # 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 max_tokens = server_args.chunked_prefill_size
if max_tokens is None: if max_tokens is None:
max_tokens = server_args.model_config.context_len max_tokens = server_args.model_config.context_len
@@ -514,7 +513,7 @@ class MiniMaxM2MoE(nn.Module):
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=config.num_local_experts 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, top_k=config.num_experts_per_tok,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
intermediate_size=config.intermediate_size, intermediate_size=config.intermediate_size,
+2 -2
View File
@@ -33,7 +33,7 @@ from sglang.srt.managers.schedule_batch import (
MultimodalInputs, MultimodalInputs,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch 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 from sglang.srt.utils import is_cpu
_is_cpu = is_cpu() _is_cpu = is_cpu()
@@ -477,7 +477,7 @@ class Llama4ForConditionalGeneration(nn.Module):
) )
self.has_vision = ( 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: if self.has_vision:

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