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