diff --git a/STACK_REVIEW_PLACEHOLDER.md b/STACK_REVIEW_PLACEHOLDER.md new file mode 100644 index 000000000..6146f0e57 --- /dev/null +++ b/STACK_REVIEW_PLACEHOLDER.md @@ -0,0 +1,5 @@ +# Full-stack review placeholder + +This file exists only to distinguish the full-stack review PR from the last +stack member. This PR is for review and full CI only — merging happens +through the individual stack PRs. Do not merge this PR. diff --git a/python/sglang/srt/arg_groups/deepseek_v4_hook.py b/python/sglang/srt/arg_groups/deepseek_v4_hook.py index 53f5b0b88..c58fc94b7 100644 --- a/python/sglang/srt/arg_groups/deepseek_v4_hook.py +++ b/python/sglang/srt/arg_groups/deepseek_v4_hook.py @@ -12,8 +12,16 @@ logger = logging.getLogger(__name__) def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None: - """Apply DeepSeek V4 model-specific server arg defaults and constraints.""" - from sglang.srt.server_args import ServerArgs + """Residual imperative arm of the DeepSeek V4 defaults. + + The attention/page/window/MoE-runner declarations moved to the override + registry (arg_groups/overrides.py: _deepseek_v4_overrides) and the + kv-cache dtype default to the resolution pipeline + (_deepseek_v4_kv_cache_dtype, invoked below at its legacy slot). This + keeps, at the legacy slot: the ROCm env fill (env-write policy), the + max_running_requests fill (the speculative hook is a later writer of + that field) and the validations. + """ from sglang.srt.utils import is_hip # FlashMLA sparse prefill (SGLANG_OPT_FLASHMLA_SPARSE_PREFILL, default on) @@ -28,31 +36,15 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None ) envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.set(False) - server_args.attention_backend = "dsv4" - server_args.page_size = 256 - if server_args.kv_cache_dtype == "auto": - server_args.kv_cache_dtype = "fp8_e4m3" - logger.warning( - f"Setting KV cache dtype to {server_args.kv_cache_dtype} for {model_arch}." - ) - - if server_args.device == "npu": - # NPU keeps the device-aware "dsv4" backend (the registry routes it to - # the Ascend V4 subclass); only the pool geometry / dtype differ. - # set_default_server_args() pins all three backends to "ascend" for - # generic NPU models; undo that here so V4 stays consistently on dsv4. - server_args.prefill_attention_backend = "dsv4" - server_args.decode_attention_backend = "dsv4" - server_args.page_size = 128 - server_args.kv_cache_dtype = "bfloat16" - - logger.info( - f"Use dsv4 attention backend for {model_arch}, setting page_size to {server_args.page_size}." + # The kv-cache dtype default moved to the resolution pipeline + # (arg_groups/overrides.py: _deepseek_v4_kv_cache_dtype), invoked here at + # its legacy slot. + from sglang.srt.arg_groups.overrides import ( + _deepseek_v4_kv_cache_dtype, + run_post_process_pass, ) - assert server_args.kv_cache_dtype in [ - "fp8_e4m3", - "bfloat16", - ], f"{server_args.kv_cache_dtype} is not supported for {model_arch}" + + run_post_process_pass(server_args, _deepseek_v4_kv_cache_dtype) if server_args.max_running_requests is None: server_args.max_running_requests = 256 @@ -68,23 +60,6 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None server_args.speculative_eagle_topk == 1 ), f"Only EAGLE speculative algorithm with topk == 1 is supported for {model_arch}" - if server_args.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio: - server_args.swa_full_tokens_ratio = 0.1 - logger.info( - f"Setting swa_full_tokens_ratio to {server_args.swa_full_tokens_ratio} for {model_arch}." - ) - - # nvidia/DeepSeek-V4-Pro-NVFP4 uses flashinfer_trtllm_routed MoE runner backend. - if ( - server_args.moe_runner_backend == "auto" - and server_args.get_model_config().nvfp4_moe_meta is not None - ): - server_args.moe_runner_backend = "flashinfer_trtllm_routed" - logger.info( - "Use flashinfer_trtllm_routed as MoE runner backend for " - f"{model_arch} hybrid FP8+NVFP4 checkpoint." - ) - def validate_deepseek_v4_cp(server_args: ServerArgs) -> None: """Validate DeepSeek V4 context-parallel configuration.""" diff --git a/python/sglang/srt/arg_groups/hisparse_hook.py b/python/sglang/srt/arg_groups/hisparse_hook.py index 856f35527..4ff8caf46 100644 --- a/python/sglang/srt/arg_groups/hisparse_hook.py +++ b/python/sglang/srt/arg_groups/hisparse_hook.py @@ -36,30 +36,8 @@ def _hisparse_allowed_backends(kv_cache_dtype: str) -> set[str]: ) -def apply_hisparse_dsa_backend_defaults( - server_args: ServerArgs, - user_set_prefill: bool, - user_set_decode: bool, - kv_cache_dtype: str, -) -> bool: - """Pick DSA backends for --enable-hisparse based on KV dtype. - - CUDA uses dtype-specific FlashMLA backends; ROCm uses TileLang. Returns - True if hisparse handled backend selection. - """ - if not server_args.enable_hisparse: - return False - - backend = _hisparse_default_backend(kv_cache_dtype) - if not user_set_prefill: - server_args.dsa_prefill_backend = backend - if not user_set_decode: - server_args.dsa_decode_backend = backend - logger.warning( - f"HiSparse enabled ({kv_cache_dtype}): using DSA backends " - f"prefill={server_args.dsa_prefill_backend}, decode={server_args.dsa_decode_backend}." - ) - return True +# The hisparse DSA backend defaults moved to the resolution pipeline +# (arg_groups/overrides.py: _dsa_split_backend_resolution, hisparse arm). def validate_hisparse_dsa_backend( diff --git a/python/sglang/srt/arg_groups/nemotron_h_hook.py b/python/sglang/srt/arg_groups/nemotron_h_hook.py deleted file mode 100644 index 5e636e0e9..000000000 --- a/python/sglang/srt/arg_groups/nemotron_h_hook.py +++ /dev/null @@ -1,66 +0,0 @@ -from __future__ import annotations - -import logging -from typing import TYPE_CHECKING - -from sglang.srt.utils.common import get_device_capability, is_cuda, is_sm100_supported - -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs - -logger = logging.getLogger(__name__) - - -def apply_nemotron_h_defaults(server_args: ServerArgs, model_arch: str) -> None: - """Apply NemotronH model-specific server arg defaults and constraints.""" - model_config = server_args.get_model_config() - is_modelopt = model_config.quantization in [ - "modelopt", - "modelopt_fp8", - "modelopt_fp4", - "modelopt_mixed", - ] - if is_modelopt: - assert model_config.hf_config.mlp_hidden_act == "relu2" - if model_config.quantization == "modelopt": - quant_algo = model_config.hf_config.quantization_config["quant_algo"] - if quant_algo == "MIXED_PRECISION": - server_args.quantization = "modelopt_mixed" - else: - server_args.quantization = ( - "modelopt_fp4" if quant_algo == "NVFP4" else "modelopt_fp8" - ) - else: - server_args.quantization = model_config.quantization - - if (is_modelopt or model_config.quantization is None) and ( - server_args.moe_runner_backend == "auto" - ): - if is_sm100_supported() and server_args.moe_a2a_backend == "none": - server_args.moe_runner_backend = "flashinfer_trtllm" - logger.info( - f"Use flashinfer_trtllm as MoE runner backend on sm100 for {model_arch}" - ) - elif ( - ( - model_config.quantization in ("modelopt_fp4", "modelopt_mixed") - or server_args.quantization == "modelopt_fp4" - ) - and is_cuda() - and (8, 0) <= get_device_capability() < (10, 0) - ): - server_args.moe_runner_backend = "marlin" - logger.info( - "Use marlin as MoE runner backend on SM80-SM90 for " - f"{model_arch} {model_config.quantization}" - ) - else: - server_args.moe_runner_backend = "flashinfer_cutlass" - - if is_sm100_supported() and server_args.attention_backend is None: - server_args.attention_backend = "flashinfer" - server_args._handle_mamba_radix_cache(model_arch=model_arch) - assert server_args.attention_backend != "triton", ( - "NemotronHForCausalLM does not support triton attention backend," - "as the first layer might not be an attention layer" - ) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 5bf8106f5..e94ac8ee0 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -37,6 +37,7 @@ from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.runtime_context import resolve_flag_leaf from sglang.srt.utils.common import ( cpu_has_amx_support, + get_device_capability, get_device_sm, get_nvidia_driver_version, get_quantization_config, @@ -176,10 +177,31 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None: ) if declared: entry = (fn.__qualname__, dict(declared)) - server_args._resolved_overrides.append(entry) + stash = getattr(server_args, "_resolved_overrides", None) + if stash is None: + # Handlers hosting pass slots may be invoked directly on fixtures + # that never ran the monolith dispatch (which owns the stash); + # create it lazily. Real publishes always pass through the + # dispatch first — the dispatch ASSIGNS the stash, so pass slots + # must sit at or after it in __post_init__ order. + stash = server_args._resolved_overrides = [] + stash.append(entry) apply_declarations_to_server_args(server_args, [entry]) +def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None: + """Transition helper for load-time resolved fields (model-file config + overrides, weight-resolved dtypes): dual-apply the declaration onto the + published ``server_args`` — byte-identical to the imperative write this + replaces — and record it into the flags tier through the runtime gate.""" + from sglang.srt.runtime_context import get_context + + ctx = get_context() + entry = (source, dict(declared)) + apply_declarations_to_server_args(ctx.server_args, [entry]) + ctx.record_runtime_overrides([entry]) + + def collect_model_override_declarations( architecture: str, server_args: Any, hf_config: Any ) -> List[Tuple[str, Dict[str, Any]]]: @@ -548,6 +570,23 @@ def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict: return overrides +@_register_for("MossVLForConditionalGeneration") +def _moss_vl_overrides(server_args: Any, hf_config: Any) -> dict: + overrides: Dict[str, Any] = {} + if server_args.is_attention_backend_not_set(): + overrides["prefill_attention_backend"] = "flashinfer" + logger.info("Use flashinfer as default prefill attention backend for Moss-VL") + prefill_backend = ( + overrides.get("prefill_attention_backend") + or server_args.get_attention_backends()[0] + ) + assert prefill_backend == "flashinfer", ( + "MossVLForConditionalGeneration requires flashinfer prefill " + "attention backend for cross-attention custom mask support." + ) + return overrides + + @_register_for("MiniCPMV4_6ForConditionalGeneration") def _minicpm_v4_6_overrides(server_args: Any, hf_config: Any) -> dict: if is_sm100_supported() and server_args.attention_backend is None: @@ -581,6 +620,108 @@ def _lfm2_overrides(server_args: Any, hf_config: Any) -> dict: return {} +@_register_for("DeepseekV4ForCausalLM") +def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict: + """DeepSeek V4 attention/page/window/MoE-runner defaults (from + arg_groups/deepseek_v4_hook.py). The kv-cache dtype and NPU split-backend + writes, the max_running_requests fill and the validations stay in the + hook at its legacy slot.""" + from sglang.srt.server_args import ServerArgs + + model_arch = hf_config.architectures[0] + overrides: Dict[str, Any] = {"attention_backend": "dsv4"} + + page_size = 256 + if server_args.device == "npu": + # NPU keeps the device-aware "dsv4" backend (the registry routes it to + # the Ascend V4 subclass); only the pool geometry / dtype differ. + # set_default_server_args() pins all three backends to "ascend" for + # generic NPU models; override that here so V4 stays consistently on + # dsv4. + page_size = 128 + overrides["prefill_attention_backend"] = "dsv4" + overrides["decode_attention_backend"] = "dsv4" + overrides["page_size"] = page_size + logger.info( + f"Use dsv4 attention backend for {model_arch}, setting page_size to {page_size}." + ) + + if server_args.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio: + overrides["swa_full_tokens_ratio"] = 0.1 + logger.info(f"Setting swa_full_tokens_ratio to 0.1 for {model_arch}.") + + # nvidia/DeepSeek-V4-Pro-NVFP4 uses flashinfer_trtllm_routed MoE runner backend. + if ( + server_args.moe_runner_backend == "auto" + and server_args.get_model_config().nvfp4_moe_meta is not None + ): + overrides["moe_runner_backend"] = "flashinfer_trtllm_routed" + logger.info( + "Use flashinfer_trtllm_routed as MoE runner backend for " + f"{model_arch} hybrid FP8+NVFP4 checkpoint." + ) + return overrides + + +@_register_for("NemotronHForCausalLM", "NemotronHPuzzleForCausalLM") +def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict: + """NemotronH quantization / MoE runner / attention backend defaults + (absorbed from the retired arg_groups/nemotron_h_hook.py; the mamba radix + cache handling and the triton-backend assert stay in the arch branch).""" + model_arch = hf_config.architectures[0] + model_config = server_args.get_model_config() + overrides: Dict[str, Any] = {} + + is_modelopt = model_config.quantization in [ + "modelopt", + "modelopt_fp8", + "modelopt_fp4", + "modelopt_mixed", + ] + quantization = server_args.quantization + if is_modelopt: + assert model_config.hf_config.mlp_hidden_act == "relu2" + if model_config.quantization == "modelopt": + quant_algo = model_config.hf_config.quantization_config["quant_algo"] + if quant_algo == "MIXED_PRECISION": + quantization = "modelopt_mixed" + else: + quantization = ( + "modelopt_fp4" if quant_algo == "NVFP4" else "modelopt_fp8" + ) + else: + quantization = model_config.quantization + overrides["quantization"] = quantization + + if (is_modelopt or model_config.quantization is None) and ( + server_args.moe_runner_backend == "auto" + ): + if is_sm100_supported() and server_args.moe_a2a_backend == "none": + overrides["moe_runner_backend"] = "flashinfer_trtllm" + logger.info( + f"Use flashinfer_trtllm as MoE runner backend on sm100 for {model_arch}" + ) + elif ( + ( + model_config.quantization in ("modelopt_fp4", "modelopt_mixed") + or quantization == "modelopt_fp4" + ) + and is_cuda() + and (8, 0) <= get_device_capability() < (10, 0) + ): + overrides["moe_runner_backend"] = "marlin" + logger.info( + "Use marlin as MoE runner backend on SM80-SM90 for " + f"{model_arch} {model_config.quantization}" + ) + else: + overrides["moe_runner_backend"] = "flashinfer_cutlass" + + if is_sm100_supported() and server_args.attention_backend is None: + overrides["attention_backend"] = "flashinfer" + return overrides + + @_register_for( "Qwen3NextForCausalLM", "Qwen3_5MoeForConditionalGeneration", @@ -750,6 +891,213 @@ def _step3p_overrides(server_args: Any, hf_config: Any) -> dict: # --------------------------------------------------------------------------- +# Architectures whose monolith branch routes through the mamba radix cache +# handling (hybrid linear-attention models). Keep in sync with the branch +# guards in _handle_model_specific_adjustments. +_MAMBA_RADIX_CACHE_ARCHS = frozenset( + { + "KimiLinearForCausalLM", + "BailingMoeV2_5ForCausalLM", + "Qwen3NextForCausalLM", + "Qwen3_5MoeForConditionalGeneration", + "InternS2PreviewForConditionalGeneration", + "Qwen3_5ForConditionalGeneration", + "MiniCPMV4_6ForConditionalGeneration", + "NemotronHForCausalLM", + "NemotronHPuzzleForCausalLM", + "FalconH1ForCausalLM", + "JetNemotronForCausalLM", + "JetVLMForConditionalGeneration", + "Lfm2ForCausalLM", + "ZayaForCausalLM", + } +) + +# Architectures that support the extra_buffer mamba radix cache strategy. +# Single source of truth: ServerArgs._support_mamba_cache_extra_buffer +# delegates here. +_MAMBA_EXTRA_BUFFER_ARCHS = frozenset( + { + "Qwen3_5ForConditionalGeneration", + "Qwen3_5MoeForConditionalGeneration", + "Qwen3NextForCausalLM", + "InternS2PreviewForConditionalGeneration", + "MiniCPMV4_6ForConditionalGeneration", + "BailingMoeV2_5ForCausalLM", + "FalconH1ForCausalLM", + "GraniteMoeHybridForCausalLM", + "NemotronHForCausalLM", + "NemotronHPuzzleForCausalLM", + } +) + + +def supports_mamba_cache_extra_buffer(view: Any, model_arch: str) -> bool: + """Whether ``model_arch`` supports the extra_buffer strategy on the + configured linear-attention backend (pure read).""" + if model_arch in _MAMBA_EXTRA_BUFFER_ARCHS: + return view.linear_attn_backend == "triton" + return False + + +@register_post_process +def _mamba_radix_cache_resolution(view: Any) -> dict: + """Resolve the hybrid-mamba radix cache fields (pure). + + Slot pass: invoked at each legacy ``_handle_mamba_radix_cache`` slot — + the hybrid-spec call at the head of the monolith and the per-arch branch + calls — where it reads the mid-resolution ``page_size`` / + ``disable_overlap_schedule`` exactly as the legacy helper did. The arch + guard replicates the union of the legacy call-site guards so the pass is + self-sufficient in the end-state pass list. + """ + from sglang.srt.configs.linear_attn_model_registry import ( + get_linear_attn_spec_by_arch, + ) + + hf_config = view.get_model_config().hf_config + model_arch = hf_config.architectures[0] + + in_branch = model_arch in _MAMBA_RADIX_CACHE_ARCHS + if model_arch == "GraniteMoeHybridForCausalLM": + in_branch = any( + layer_type == "mamba" + for layer_type in getattr(hf_config, "layer_types", []) + ) + spec = get_linear_attn_spec_by_arch(model_arch) + if not ((spec is not None and spec.uses_mamba_radix_cache) or in_branch): + return {} + + if view.disable_radix_cache: + return {} + + declared: Dict[str, Any] = {"uses_mamba_radix_cache": True} + if view.mamba_radix_cache_strategy == "auto": + wants_overlap = not view.disable_overlap_schedule + wants_paging = view.page_size is not None and view.page_size > 1 + if (wants_overlap or wants_paging) and supports_mamba_cache_extra_buffer( + view, model_arch + ): + declared["mamba_radix_cache_strategy"] = "extra_buffer" + else: + declared["mamba_radix_cache_strategy"] = "no_buffer" + declared["disable_overlap_schedule"] = True + return declared + + +@register_post_process +def _dsa_kv_cache_dtype_default(view: Any) -> dict: + """Slot pass in the DSA arm, ordered before the split-backend + resolution: default the kv-cache dtype from the device capability + (Blackwell FP8, Hopper bf16) and normalize the bf16 alias. Reads the + PRISTINE dsa split backends (their resolution runs after this pass).""" + from sglang.srt.configs.model_config import is_deepseek_dsa + + hf_config = view.get_model_config().hf_config + if hf_config.architectures[0] not in _DEEPSEEK_FAMILY_ARCHS: + return {} + if not is_deepseek_dsa(hf_config): + return {} + if is_npu() or is_xpu(): + return {} + + import torch + + major, _ = torch.cuda.get_device_capability() + + # If user specified a backend but didn't explicitly set kv_cache_dtype, + # suggest them to be explicit about kv_cache_dtype to avoid surprises + if ( + view.dsa_prefill_backend is not None or view.dsa_decode_backend is not None + ) and view.kv_cache_dtype == "auto": + logger.warning( + "When specifying --dsa-prefill-backend or --dsa-decode-backend, " + "you should also explicitly set --kv-cache-dtype (e.g., 'fp8_e4m3' or 'bfloat16'). " + "DeepSeek V3.2 defaults to FP8 KV cache which may not be compatible with all backends." + ) + + kv_cache_dtype = view.kv_cache_dtype + if kv_cache_dtype == "auto": + kv_cache_dtype = "fp8_e4m3" if major >= 10 else "bfloat16" + logger.warning( + f"Setting KV cache dtype to {kv_cache_dtype} for DeepSeek DSA on SM{major} device." + ) + if kv_cache_dtype == "bf16": + kv_cache_dtype = "bfloat16" + assert kv_cache_dtype in [ + "bfloat16", + "fp8_e4m3", + ], "DeepSeek DSA only supports bf16/bfloat16 or fp8_e4m3 kv_cache_dtype" + if kv_cache_dtype != view.kv_cache_dtype: + return {"kv_cache_dtype": kv_cache_dtype} + return {} + + +@register_post_process +def _dsa_split_backend_resolution(view: Any) -> dict: + """Slot pass in the DSA arm: default the DSA prefill/decode split + backends from the mid-resolution kv-cache dtype and the device + capability. The hisparse arm takes precedence under --enable-hisparse.""" + from sglang.srt.configs.model_config import is_deepseek_dsa + + hf_config = view.get_model_config().hf_config + if hf_config.architectures[0] not in _DEEPSEEK_FAMILY_ARCHS: + return {} + if not is_deepseek_dsa(hf_config): + return {} + if is_npu() or is_xpu(): + return {} + + import torch + + major, _ = torch.cuda.get_device_capability() + kv_cache_dtype = view.kv_cache_dtype + user_set_prefill = view.dsa_prefill_backend is not None + user_set_decode = view.dsa_decode_backend is not None + declared: Dict[str, Any] = {} + + if view.enable_hisparse: + from sglang.srt.arg_groups.hisparse_hook import _hisparse_default_backend + + backend = _hisparse_default_backend(kv_cache_dtype) + if not user_set_prefill: + declared["dsa_prefill_backend"] = backend + if not user_set_decode: + declared["dsa_decode_backend"] = backend + prefill = declared.get("dsa_prefill_backend", view.dsa_prefill_backend) + decode = declared.get("dsa_decode_backend", view.dsa_decode_backend) + logger.warning( + f"HiSparse enabled ({kv_cache_dtype}): using DSA backends " + f"prefill={prefill}, decode={decode}." + ) + return declared + + if not user_set_prefill and not user_set_decode and is_hip(): + declared["dsa_prefill_backend"] = "tilelang" + declared["dsa_decode_backend"] = "tilelang" + elif kv_cache_dtype == "fp8_e4m3": + # Blackwell FP8 defaults to trtllm; Hopper FP8 to flashmla_kv. + default = "trtllm" if major >= 10 else "flashmla_kv" + if not user_set_prefill: + declared["dsa_prefill_backend"] = default + if not user_set_decode: + declared["dsa_decode_backend"] = default + else: + # Set prefill/decode backends based on hardware architecture. + if not user_set_prefill: + declared["dsa_prefill_backend"] = "flashmla_sparse" + if not user_set_decode: + declared["dsa_decode_backend"] = "trtllm" if major >= 10 else "fa3" + + prefill = declared.get("dsa_prefill_backend", view.dsa_prefill_backend) + decode = declared.get("dsa_decode_backend", view.dsa_decode_backend) + logger.warning( + f"Set DSA backends for {kv_cache_dtype} KV Cache: " + f"prefill={prefill}, decode={decode}." + ) + return declared + + # Keep in sync with the DeepSeek family list on _deepseek_family_overrides. _DEEPSEEK_FAMILY_ARCHS = frozenset( { @@ -835,6 +1183,169 @@ def _deepseek_moe_quant_resolution(view: Any) -> dict: return overrides +@register_post_process +def _deepseek_spec_moe_resolution(view: Any) -> dict: + """Slot pass at the DeepSeek branch's HIP arm: draft (nextn) spec-MoE + backends for the DeepSeek fp4 checkpoint. Reads the mid-resolution + quantization (after _deepseek_moe_quant_resolution) and the pre-a2a + ep_size, exactly like the legacy in-branch writes.""" + from sglang.srt.environ import envs + + hf_config = view.get_model_config().hf_config + model_arch = hf_config.architectures[0] + if model_arch not in _DEEPSEEK_FAMILY_ARCHS: + return {} + if not is_hip(): + return {} + if not ( + view.quantization == "modelopt_fp4" + and view.speculative_algorithm == "EAGLE" + and ( + view.speculative_moe_runner_backend is None + or view.speculative_moe_a2a_backend is None + ) + ): + return {} + if envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE.get(): + logger.info( + "Use deep_gemm moe runner and deepep a2a backend for bf16 nextn layer in deepseek fp4 checkpoint." + ) + # Validate usage of ep + if view.ep_size == 1: + raise ValueError( + "Invalid configuration: 'deep_gemm' speculative MoE runner backend with " + "'deepep' a2a backend requires expert parallelism (ep_size > 1). " + f"Current ep_size is {view.ep_size}. " + "Please set --ep-size > 1 (e.g., --ep-size 8) to use this configuration, " + "or change --speculative-moe-a2a-backend to 'none' if expert parallelism is not available." + ) + return { + "speculative_moe_runner_backend": "deep_gemm", + "speculative_moe_a2a_backend": "deepep", + } + logger.info( + "Use triton fused moe by default for bf16 nextn layer in deepseek fp4 checkpoint." + ) + return { + "speculative_moe_runner_backend": "triton", + "speculative_moe_a2a_backend": "none", + } + + +@register_post_process +def _deepseek_v4_kv_cache_dtype(view: Any) -> dict: + """Slot pass in the DeepSeek V4 hook: default the kv-cache dtype to FP8 + (bfloat16 on NPU, where the pool geometry differs) and validate the + result. The NPU split-backend writes stay in the hook.""" + hf_config = view.get_model_config().hf_config + model_arch = hf_config.architectures[0] + if model_arch != "DeepseekV4ForCausalLM": + return {} + + kv_cache_dtype = view.kv_cache_dtype + if kv_cache_dtype == "auto": + kv_cache_dtype = "fp8_e4m3" + logger.warning(f"Setting KV cache dtype to {kv_cache_dtype} for {model_arch}.") + if view.device == "npu": + kv_cache_dtype = "bfloat16" + assert kv_cache_dtype in [ + "fp8_e4m3", + "bfloat16", + ], f"{kv_cache_dtype} is not supported for {model_arch}" + if kv_cache_dtype != view.kv_cache_dtype: + return {"kv_cache_dtype": kv_cache_dtype} + return {} + + +@register_post_process +def _deepseek_v4_sm120_moe(view: Any) -> dict: + """Slot pass in the DeepSeek V4 validation branch: SM120 lacks + tcgen05/TMEM, fall back to the marlin MoE runner (reads the + mid-resolution moe_runner_backend, after the dispatch-time nvfp4 + default).""" + hf_config = view.get_model_config().hf_config + if hf_config.architectures[0] != "DeepseekV4ForCausalLM": + return {} + if is_sm120_supported() and view.moe_runner_backend == "auto": + logger.info("Use marlin as MoE runner backend on SM120 for DeepseekV4") + return {"moe_runner_backend": "marlin"} + return {} + + +@register_post_process +def _sparse_head_overlap_disable(view: Any) -> dict: + from sglang.srt.environ import envs + + if envs.SGLANG_EMBEDDINGS_SPARSE_HEAD.is_set(): + logger.warning( + "Overlap scheduler is disabled when using sparse head for embedding model." + ) + return {"disable_overlap_schedule": True} + return {} + + +# Architectures with explicit FlashInfer AllReduce Fusion support. Keep in +# sync with the model-side fusion implementations. +_FLASHINFER_ALLREDUCE_FUSION_ARCHS = frozenset( + { + "DeepseekV3ForCausalLM", + "DeepseekV32ForCausalLM", + "GptOssForCausalLM", + "GlmMoeDsaForCausalLM", + "Glm4MoeForCausalLM", + "Glm4MoeLiteForCausalLM", + "MistralLarge3ForCausalLM", + "Qwen3MoeForCausalLM", + "Qwen3VLMoeForConditionalGeneration", + "Qwen3NextForCausalLM", + "KimiK25ForConditionalGeneration", + "Qwen3_5MoeForConditionalGeneration", + "InternS2PreviewForConditionalGeneration", + "Qwen3_5ForConditionalGeneration", + "NemotronHForCausalLM", + "NemotronHPuzzleForCausalLM", + } +) + + +@register_post_process +def _flashinfer_allreduce_fusion_auto_enable(view: Any) -> dict: + """Slot pass at the monolith tail: auto-enable FlashInfer AllReduce + Fusion on SM90/SM100 for models with explicit support. auto resolves to + mnnvl on Blackwell (single- and multi-node) and trtllm on SM90 + single-node systems. Reads the mid-resolution enable_dp_attention / + moe_a2a_backend (after the DeepSeek CP and a2a declarations), exactly + like the legacy tail block.""" + model_arch = view.get_model_config().hf_config.architectures[0] + if ( + view.flashinfer_allreduce_fusion_backend is None + and model_arch in _FLASHINFER_ALLREDUCE_FUSION_ARCHS + and (is_sm90_supported() or is_sm100_supported()) + and view.tp_size > 1 + and not view.enable_dp_attention + and (view.nnodes == 1 or is_sm100_supported()) + and view.moe_a2a_backend == "none" + ): + logger.info( + f"Auto-enabling FlashInfer AllReduce Fusion on SM90/SM10X for {model_arch}" + ) + return {"flashinfer_allreduce_fusion_backend": "auto"} + return {} + + +@register_post_process +def _enforce_disable_allreduce_fusion(view: Any) -> dict: + """Slot pass right after the auto-enable: the user's enforce-disable + switch wins over every model-specific adjustment.""" + if view.enforce_disable_flashinfer_allreduce_fusion: + logger.info( + "FlashInfer allreduce fusion is forcibly disabled " + "via --enforce-disable-flashinfer-allreduce-fusion." + ) + return {"flashinfer_allreduce_fusion_backend": None} + return {} + + @register_post_process def _sampling_backend_default(view: Any) -> dict: if view.sampling_backend is None: @@ -878,6 +1389,19 @@ def _deterministic_is_deepseek_model(view: Any) -> bool: return False +@register_post_process +def _deterministic_allreduce_fusion_disable(view: Any) -> dict: + if ( + view.enable_deterministic_inference + and view.flashinfer_allreduce_fusion_backend is not None + ): + logger.warning( + "Disable --flashinfer-allreduce-fusion-backend because deterministic inference is enabled." + ) + return {"flashinfer_allreduce_fusion_backend": None} + return {} + + @register_post_process def _deterministic_attention_backend(view: Any) -> dict: if not view.enable_deterministic_inference: @@ -995,6 +1519,39 @@ def _mla_backend_page_constraints(view: Any) -> dict: return {} +@register_post_process +def _cutedsl_prefill_backend_fill(view: Any) -> dict: + """Slot pass in the attention-backend compatibility handler: CuteDSL MLA + is decode-only, so validate the combination and default the prefill side + to trtllm_mla. The trtllm_mha check that follows at the legacy slot reads + the dual-applied value.""" + if not ( + view.attention_backend == "cutedsl_mla" + or view.decode_attention_backend == "cutedsl_mla" + or view.prefill_attention_backend == "cutedsl_mla" + ): + return {} + assert ( + view.prefill_attention_backend != "cutedsl_mla" + ), "CuteDSL MLA only supports decoding for now" + if not is_sm100_supported(): + raise ValueError( + "CuteDSL MLA backend is only supported on Blackwell GPUs (SM100). Please use a different backend." + ) + if view.kv_cache_dtype not in [ + "fp8_e4m3", + "bf16", + "bfloat16", + "auto", + ]: + raise ValueError( + "CuteDSL MLA backend only supports kv-cache-dtype of fp8_e4m3, bf16, or auto." + ) + if view.prefill_attention_backend is None: + return {"prefill_attention_backend": "trtllm_mla"} + return {} + + @register_post_process def _attention_backend_fa3_fp8_fallback(view: Any) -> dict: if view.attention_backend == "fa3" and view.kv_cache_dtype == "fp8_e5m2": @@ -1221,6 +1778,24 @@ def _a2a_ep_size(view: Any) -> dict: return {} +@register_post_process +def _pipeline_parallel_overlap_disable(view: Any) -> dict: + if view.pp_size > 1: + logger.warning("Pipeline parallelism is incompatible with overlap schedule.") + return {"disable_overlap_schedule": True} + return {} + + +@register_post_process +def _speculative_moe_runner_default(view: Any) -> dict: + """Default the speculative (draft) MoE runner backend to the resolved + target-model backend. Invoked at the head of the speculative-decoding + hook, after the MoE kernel chain has resolved.""" + if view.speculative_moe_runner_backend is None: + return {"speculative_moe_runner_backend": view.moe_runner_backend} + return {} + + @register_post_process def _gguf_quantization(view: Any) -> dict: from sglang.srt.utils.hf_transformers_utils import check_gguf_file @@ -1257,6 +1832,18 @@ def _dllm_attention_backend(view: Any) -> dict: return {} +@register_post_process +def _dllm_overlap_disable(view: Any) -> dict: + if view.dllm_algorithm is None: + return {} + if view.disable_overlap_schedule: + return {} + logger.warning( + "Overlap schedule is disabled because of using diffusion LLM inference" + ) + return {"disable_overlap_schedule": True} + + @register_post_process def _dllm_page_size(view: Any) -> dict: if view.dllm_algorithm is None or view.disable_radix_cache: diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index b1b107215..bd2cc9dd4 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -58,8 +58,14 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None: ): server_args.speculative_draft_model_revision = "main" - if server_args.speculative_moe_runner_backend is None: - server_args.speculative_moe_runner_backend = server_args.moe_runner_backend + # Moved to the resolution pipeline (arg_groups/overrides.py: + # _speculative_moe_runner_default), invoked here at its legacy slot. + from sglang.srt.arg_groups.overrides import ( + _speculative_moe_runner_default, + run_post_process_pass, + ) + + run_post_process_pass(server_args, _speculative_moe_runner_default) if server_args.speculative_algorithm is not None: server_args.speculative_algorithm = server_args.speculative_algorithm.upper() diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 145b47a63..7d67a71ec 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -39,7 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import ( compute_position, ) from sglang.srt.model_executor.forward_context import get_attn_backend -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel 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 @@ -615,7 +615,7 @@ class TboForwardBatchPreparer: sum_field=None, ) _, child_b.extend_start_loc = compute_position( - get_global_server_args().attention_backend, + get_flags().attn.backend, child_b.extend_prefix_lens, child_b.extend_seq_lens, child_b.extend_num_tokens, diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index be07fbe08..71c1fc828 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -27,6 +27,7 @@ from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache 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_flags from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast @@ -319,7 +320,7 @@ class AscendAttnBackend(AttentionBackend): self.graph_mode = False self.use_fa = get_bool_env_var("ASCEND_USE_FA", "False") self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False") - self.enable_torch_compile = model_runner.server_args.enable_torch_compile + self.enable_torch_compile = get_flags().capture.enable_torch_compile self.speculative_num_draft_tokens = ( model_runner.server_args.speculative_num_draft_tokens ) diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 142e90543..661c3c6c8 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -47,7 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, ForwardMode, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils.common import ( is_cpu, @@ -267,7 +267,7 @@ class LogitsProcessor(nn.Module): self.config = config self.vocab_size = config.vocab_size self.logit_scale = logit_scale - self.use_attn_tp_group = get_global_server_args().enable_dp_lm_head + self.use_attn_tp_group = get_flags().enable_dp_lm_head self.use_fp32_lm_head = get_global_server_args().enable_fp32_lm_head if self.use_attn_tp_group: self.attn_tp_size = get_parallel().attn_tp_size diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index 10b0d99c6..4ddeebeff 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -18,6 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import ( yarn_get_mscale_simple, yarn_linear_ramp_mask, ) +from sglang.srt.runtime_context import get_flags from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( cpu_has_amx_support, @@ -226,7 +227,7 @@ class MRotaryEmbedding(RotaryEmbedding): last_dim = cos_sin.size()[-1] cos, sin = cos_sin.chunk(2, dim=-1) if self.mrope_interleaved: - if support_triton(get_global_server_args().attention_backend): + if support_triton(get_flags().attn.backend): cos = apply_interleaved_rope_triton(cos, self.mrope_section) sin = apply_interleaved_rope_triton(sin, self.mrope_section) else: diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index 941cdb353..e83e4157f 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -13,6 +13,7 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.utils.hash import murmur_hash32 from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs +from sglang.srt.runtime_context import get_flags 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 @@ -79,7 +80,7 @@ class Sampler(nn.Module): ) # 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_global_server_args().sampling_backend == "ascend" + self.use_ascend_backend = get_flags().sampling_backend == "ascend" def _preprocess_logits( self, logits: torch.Tensor, sampling_info: SamplingBatchInfo @@ -230,7 +231,7 @@ class Sampler(nn.Module): positions=positions, ) else: - backend = get_global_server_args().sampling_backend + backend = get_flags().sampling_backend if backend == "flashinfer": assert ( sampling_info.sampling_seed is None diff --git a/python/sglang/srt/layers/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index 74ddc07db..4c74eeb6f 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -87,7 +87,7 @@ def get_cp_padding_align_size() -> int: def is_mla_prefill_cp_enabled() -> bool: sa = get_global_server_args() - return sa.enable_prefill_context_parallel and sa.use_mla_backend + return sa.enable_prefill_context_parallel and sa.use_mla_backend() def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None): diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 1987ad172..808662a3f 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -553,6 +553,14 @@ class Scheduler( self.init_batch_result_processor() + # The config-resolution lifecycle of this scheduler process ends + # here: every load-time stage has run (target and draft model init, + # weight-resolved kv-cache dtype), so lock the static flag groups. + # flags.capture stays writable; late resolution writes now raise. + from sglang.srt.runtime_context import get_context + + get_context().freeze_flags() + self.is_initializing = False def init_zbal_on_npu(self): diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 036f412a2..233a6ef29 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -25,6 +25,7 @@ from sglang.srt.mem_cache.triton_ops.common import ( get_last_loc_triton_safe, write_req_to_token_pool_triton, ) +from sglang.srt.runtime_context import get_flags from sglang.srt.server_args import ServerArgs, get_global_server_args 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 @@ -133,7 +134,7 @@ def write_cache_indices( prefix_tensors: list[torch.Tensor], req_to_token_pool: ReqToTokenPool, ): - if support_triton(get_global_server_args().attention_backend): + if support_triton(get_flags().attn.backend): prefix_pointers = torch.tensor( [t.data_ptr() for t in prefix_tensors], dtype=torch.uint64, @@ -174,7 +175,7 @@ def get_last_loc( req_pool_indices_tensor: torch.Tensor, prefix_lens_tensor: torch.Tensor, ) -> torch.Tensor: - attn_backend = get_global_server_args().attention_backend + attn_backend = get_flags().attn.backend uses_triton_dispatch = attn_backend not in ("ascend", "torch_native") if _is_hip and uses_triton_dispatch: diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 23f2e9596..a02ccedf6 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -38,6 +38,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode +from sglang.srt.runtime_context import get_flags from sglang.srt.utils import ( empty_context, log_info_on_rank0, @@ -558,7 +559,7 @@ class CPUGraphRunner: # bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only) self.graphs_cross = {} self.output_buffers = {} - self.enable_torch_compile = model_runner.server_args.enable_torch_compile + self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 27b2452ae..ea17a2558 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -180,6 +180,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( from sglang.srt.model_loader.utils import set_default_torch_dtype from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.platforms import current_platform +from sglang.srt.runtime_context import get_flags from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.server_args import ( ServerArgs, @@ -532,15 +533,13 @@ class ModelRunner(ModelRunnerKVCacheMixin): # only, so a draft init cannot clobber target-derived global state). if not self.is_draft_worker: set_global_server_args_for_scheduler(server_args) - # FIXME: hacky set `use_mla_backend` - get_global_server_args().use_mla_backend = self.use_mla_backend # Init OpenMP threads binding for CPU if self.device == "cpu": self.init_threads_binding() # Set float32 matmul precision - if server_args.enable_tf32_matmul: + if get_flags().enable_tf32_matmul: torch.set_float32_matmul_precision("high") # Get available memory before model loading. @@ -2427,6 +2426,23 @@ class ModelRunner(ModelRunnerKVCacheMixin): result = self._get_linear_attn_registry_result() return result[1] if result else None + def _record_kv_cache_dtype(self, resolved: str) -> None: + # Load-time resolution transition: the weight-resolved kv-cache dtype + # is declared into the flags tier; the dual-apply inside the helper + # replaces the legacy in-place write. Mock runners whose server_args + # is not the published object keep the plain write. + from sglang.srt.runtime_context import get_context + + if get_context()._server_args is self.server_args: + from sglang.srt.arg_groups.overrides import declare_load_time_override + + declare_load_time_override( + "ModelRunner.configure_kv_cache_dtype", + {"kv_cache_dtype": resolved}, + ) + else: + self.server_args.kv_cache_dtype = resolved + def configure_kv_cache_dtype(self): if self.server_args.kv_cache_dtype == "auto": quant_config = getattr(self.model, "quant_config", None) @@ -2435,16 +2451,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): isinstance(kv_cache_quant_algo, str) and kv_cache_quant_algo.upper() == "FP8" ): - if _is_hip: - self.kv_cache_dtype = fp8_dtype - self.server_args.kv_cache_dtype = TORCH_DTYPE_TO_KV_CACHE_STR[ - self.kv_cache_dtype - ] - else: - self.kv_cache_dtype = torch.float8_e4m3fn - self.server_args.kv_cache_dtype = TORCH_DTYPE_TO_KV_CACHE_STR[ - self.kv_cache_dtype - ] + self.kv_cache_dtype = fp8_dtype if _is_hip else torch.float8_e4m3fn + self._record_kv_cache_dtype( + TORCH_DTYPE_TO_KV_CACHE_STR[self.kv_cache_dtype] + ) else: self.kv_cache_dtype = self.dtype elif self.server_args.kv_cache_dtype == "fp8_e5m2": @@ -2624,7 +2634,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): ): return - if self.device == "cpu" and not self.server_args.enable_torch_compile: + if self.device == "cpu" and not get_flags().capture.enable_torch_compile: return tic = time.perf_counter() diff --git a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py index d71c93b9a..dee348d3e 100644 --- a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py @@ -23,7 +23,7 @@ from contextlib import contextmanager from typing import TYPE_CHECKING, Any, List, Sequence, Tuple from sglang.srt.model_executor.runner.base_runner import BaseRunner -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import require_gathered_buffer if TYPE_CHECKING: @@ -94,7 +94,7 @@ def get_batch_sizes_to_capture( assert len(capture_bs) > 0 and capture_bs[0] > 0, f"{capture_bs=}" compile_bs = ( [bs for bs in capture_bs if bs <= server_args.torch_compile_max_bs] - if server_args.enable_torch_compile + if get_flags().capture.enable_torch_compile else [] ) return capture_bs, compile_bs diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 12417c16d..78458c9f6 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -44,7 +44,7 @@ from sglang.srt.model_executor.runner.flashinfer_autotune import ( run_flashinfer_autotune_forward, should_run_flashinfer_autotune, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.speculative.spec_info import create_dummy_verify_input from sglang.srt.utils import ( empty_context, @@ -370,7 +370,7 @@ class BaseRunner(ABC): seq_len_fill_value = mr.attn_backend.get_cuda_graph_seq_len_fill_value() - if mr.server_args.enable_torch_compile: + if get_flags().capture.enable_torch_compile: set_torch_compile_config() should_disable_torch_compile = not getattr( mr.model, "_can_torch_compile", True @@ -381,7 +381,7 @@ class BaseRunner(ABC): "Transformers backend model reports it is not torch.compile " "compatible (e.g. dynamic rope scaling). Disabling torch.compile.", ) - mr.server_args.enable_torch_compile = False + get_flags().capture.enable_torch_compile = False # NOTE: aux hidden state capture (eagle3/dflash) is already # configured by init_aux_hidden_state_capture() in initialize(). diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index eceb932ea..9d6e0e1e2 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -93,6 +93,7 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import ( DeepEPCudaGraphRunnerAdapter, ) from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups +from sglang.srt.runtime_context import get_flags from sglang.srt.utils import ( empty_context, get_available_gpu_memory, @@ -187,7 +188,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): ): super().__init__(model_runner) # --- core state ------------------------------------------------ - self.enable_torch_compile = model_runner.server_args.enable_torch_compile + self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) @@ -668,9 +669,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): self.warmup() # warmup() may disable torch.compile for a model whose _can_torch_compile # is False; recompute the compile bucket so capture matches. - if self.enable_torch_compile and not ( - self.model_runner.server_args.enable_torch_compile - ): + if self.enable_torch_compile and not (get_flags().capture.enable_torch_compile): self.enable_torch_compile = False _, self.compile_bs = get_batch_sizes_to_capture( self.model_runner, self.num_tokens_per_bs diff --git a/python/sglang/srt/models/apertus.py b/python/sglang/srt/models/apertus.py index f20380724..6d8daf25e 100644 --- a/python/sglang/srt/models/apertus.py +++ b/python/sglang/srt/models/apertus.py @@ -52,8 +52,7 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, maybe_remap_kv_scale_name, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix, make_layers logger = logging.getLogger(__name__) @@ -447,7 +446,7 @@ class ApertusForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) diff --git a/python/sglang/srt/models/arcee.py b/python/sglang/srt/models/arcee.py index 7934b4418..0bc11476a 100644 --- a/python/sglang/srt/models/arcee.py +++ b/python/sglang/srt/models/arcee.py @@ -46,8 +46,7 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, maybe_remap_kv_scale_name, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix, make_layers logger = logging.getLogger(__name__) @@ -406,7 +405,7 @@ class ArceeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 093bccd60..b86fd5eef 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -77,7 +77,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel 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 @@ -832,7 +832,7 @@ class BailingMoEForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/bailing_moe_linear.py b/python/sglang/srt/models/bailing_moe_linear.py index 1e38b2916..4e2dbdfbe 100644 --- a/python/sglang/srt/models/bailing_moe_linear.py +++ b/python/sglang/srt/models/bailing_moe_linear.py @@ -58,7 +58,7 @@ 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.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( BumpAllocator, @@ -1089,7 +1089,7 @@ class BailingMoELinearForCausalLM(nn.Module): config.hidden_size, params_dtype=torch.float32, quant_config=quant_config, - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/bailing_moe_nextn.py b/python/sglang/srt/models/bailing_moe_nextn.py index 648b304be..489fa58c1 100644 --- a/python/sglang/srt/models/bailing_moe_nextn.py +++ b/python/sglang/srt/models/bailing_moe_nextn.py @@ -42,8 +42,7 @@ from sglang.srt.models.bailing_moe_linear import ( BailingMoeV2_5ForCausalLM, ) from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import BumpAllocator, add_prefix LoraConfig = None @@ -209,7 +208,7 @@ class BailingMoeForCausalLMNextN(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("model.shared_head.head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid": diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 6847cc180..1a902303e 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -57,7 +57,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8 from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu @@ -172,7 +172,7 @@ class DeepseekModelNextN(nn.Module): if ( _is_npu and self.quant_config is None - and get_global_server_args().quantization is not None + and get_flags().quantization is not None ): # ascend mtp unquant exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) @@ -321,7 +321,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("model.shared_head.head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 395a7a023..82d421a96 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -175,7 +175,7 @@ from sglang.srt.models.deepseek_common.utils import ( _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( @@ -548,7 +548,7 @@ class DeepseekV2MoE(nn.Module): n_shared_experts = ( 0 if config.n_shared_experts is None else int(config.n_shared_experts) ) - _fusion_disabled = get_global_server_args().disable_shared_experts_fusion + _fusion_disabled = get_flags().disable_shared_experts_fusion # num_fused_shared_experts drives weight remapping in deepseek_weight_loader: # mlp.shared_experts → mlp.experts.256 when > 0. @@ -886,7 +886,7 @@ class DeepseekV2MoE(nn.Module): and hidden_states.shape[0] > 0 and get_is_capture_mode() and not ( - server_args.enable_torch_compile + get_flags().capture.enable_torch_compile and hidden_states.shape[0] <= server_args.torch_compile_max_bs * (server_args.speculative_num_draft_tokens or 1) @@ -2684,7 +2684,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) else: # ranks other than the last rank will have a placeholder layer @@ -2723,7 +2723,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): self.num_fused_shared_experts = 0 server_args = get_global_server_args() - if server_args.disable_shared_experts_fusion: + if get_flags().disable_shared_experts_fusion: return disable_reason = None @@ -2774,7 +2774,12 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): disable_reason = "Deepseek V3/R1 W4AFP8 model uses different quant method for routed experts and shared experts." if disable_reason is not None: - server_args.disable_shared_experts_fusion = True + from sglang.srt.arg_groups.overrides import declare_load_time_override + + declare_load_time_override( + "DeepseekV2ForCausalLM.determine_num_fused_shared_experts", + {"disable_shared_experts_fusion": True}, + ) self.num_fused_shared_experts = 0 log_info_on_rank0( logger, diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 452cba6af..636ab7e14 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -119,7 +119,7 @@ from sglang.srt.models.deepseek_v2 import ( _is_npu, _is_xpu, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel if not _is_hip: from sglang.srt.layers.utils.cp_utils import ( @@ -1824,7 +1824,7 @@ class DeepseekV4ForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) else: self.lm_head = PPMissingLayer() @@ -1862,7 +1862,7 @@ class DeepseekV4ForCausalLM(nn.Module): def determine_num_fused_shared_experts(self): self.num_fused_shared_experts = 0 - if get_global_server_args().disable_shared_experts_fusion: + if get_flags().disable_shared_experts_fusion: return disable_reason = None @@ -1876,7 +1876,12 @@ class DeepseekV4ForCausalLM(nn.Module): disable_reason = "Config does not support fused shared expert(s)." if disable_reason is not None: - get_global_server_args().disable_shared_experts_fusion = True + from sglang.srt.arg_groups.overrides import declare_load_time_override + + declare_load_time_override( + "DeepseekV4ForCausalLM.determine_num_fused_shared_experts", + {"disable_shared_experts_fusion": True}, + ) log_info_on_rank0( logger, f"{disable_reason} Shared experts fusion optimization is disabled.", diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index 46ea9e52f..e80c1ec6e 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -38,8 +38,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -234,7 +233,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("model.shared_head.head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/exaone4.py b/python/sglang/srt/models/exaone4.py index e14bf6e2b..8465a5d6e 100644 --- a/python/sglang/srt/models/exaone4.py +++ b/python/sglang/srt/models/exaone4.py @@ -31,8 +31,7 @@ from sglang.srt.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix, make_layers from sglang.utils import get_exception_traceback, logger @@ -444,7 +443,7 @@ class Exaone4ForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/exaone_moe.py b/python/sglang/srt/models/exaone_moe.py index 1ee8feeb3..fef7487c7 100755 --- a/python/sglang/srt/models/exaone_moe.py +++ b/python/sglang/srt/models/exaone_moe.py @@ -62,7 +62,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers @@ -652,7 +652,7 @@ class ExaoneMoEForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) # For EAGLE3 support diff --git a/python/sglang/srt/models/exaone_moe_mtp.py b/python/sglang/srt/models/exaone_moe_mtp.py index ed7125e7b..5975087ee 100644 --- a/python/sglang/srt/models/exaone_moe_mtp.py +++ b/python/sglang/srt/models/exaone_moe_mtp.py @@ -30,8 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -64,7 +63,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/falcon_h1.py b/python/sglang/srt/models/falcon_h1.py index 1bc2d8103..780631250 100644 --- a/python/sglang/srt/models/falcon_h1.py +++ b/python/sglang/srt/models/falcon_h1.py @@ -33,8 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.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_flags, get_parallel from sglang.srt.utils import add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -478,7 +477,7 @@ class FalconH1ForCausalLM(nn.Module): quant_config=quant_config, org_num_embeddings=config.vocab_size, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.lm_head = self.lm_head.float() self.lm_head_multiplier = config.lm_head_multiplier diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 5bc6346c5..558a83b38 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -82,7 +82,7 @@ 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.models.deepseek_v2 import DeepseekV2ForCausalLM from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, @@ -404,9 +404,7 @@ class Glm4MoeSparseMoeBlock(nn.Module): self.routed_scaling_factor = config.routed_scaling_factor self.n_shared_experts = config.n_shared_experts self.num_fused_shared_experts = ( - 0 - if get_global_server_args().disable_shared_experts_fusion - else config.n_shared_experts + 0 if get_flags().disable_shared_experts_fusion else config.n_shared_experts ) self.config = config @@ -1186,7 +1184,7 @@ class Glm4MoeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) @@ -1194,7 +1192,7 @@ class Glm4MoeForCausalLM(nn.Module): self.capture_aux_hidden_states = False def determine_num_fused_shared_experts(self): - if get_global_server_args().disable_shared_experts_fusion: + if get_flags().disable_shared_experts_fusion: return disable_reason = None @@ -1217,7 +1215,12 @@ class Glm4MoeForCausalLM(nn.Module): disable_reason = "GLM-4.5 W4AFP8 model uses different quant method for routed experts and shared experts." if disable_reason is not None: - get_global_server_args().disable_shared_experts_fusion = True + from sglang.srt.arg_groups.overrides import declare_load_time_override + + declare_load_time_override( + "Glm4MoeForCausalLM.determine_num_fused_shared_experts", + {"disable_shared_experts_fusion": True}, + ) self.num_fused_shared_experts = 0 log_info_on_rank0( logger, diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index e821fe89f..f97bd3e48 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -74,7 +74,7 @@ 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 +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( BumpAllocator, @@ -188,9 +188,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module): self.routed_scaling_factor = config.routed_scaling_factor self.n_shared_experts = config.n_shared_experts self.num_fused_shared_experts = ( - 0 - if get_global_server_args().disable_shared_experts_fusion - else config.n_shared_experts + 0 if get_flags().disable_shared_experts_fusion else config.n_shared_experts ) self.config = config self.layer_id = layer_id @@ -920,7 +918,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) @@ -941,7 +939,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): self, architecture: str = "Glm4MoeLiteForCausalLM" ): self.num_fused_shared_experts = 0 - if get_global_server_args().disable_shared_experts_fusion: + if get_flags().disable_shared_experts_fusion: return disable_reason = None @@ -956,7 +954,12 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): disable_reason = "GLM-4.5 or GLM-4.6 cannot use shared experts fusion optimization under expert parallelism." if disable_reason is not None: - get_global_server_args().disable_shared_experts_fusion = True + from sglang.srt.arg_groups.overrides import declare_load_time_override + + declare_load_time_override( + "Glm4MoeLiteForCausalLM.determine_num_fused_shared_experts", + {"disable_shared_experts_fusion": True}, + ) self.num_fused_shared_experts = 0 log_info_on_rank0( logger, diff --git a/python/sglang/srt/models/glm4_moe_lite_nextn.py b/python/sglang/srt/models/glm4_moe_lite_nextn.py index 7682afadb..50a644e38 100644 --- a/python/sglang/srt/models/glm4_moe_lite_nextn.py +++ b/python/sglang/srt/models/glm4_moe_lite_nextn.py @@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import ( Glm4MoeLiteDecoderLayer, Glm4MoeLiteForCausalLM, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import BumpAllocator, add_prefix, is_npu @@ -155,12 +155,12 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("model.shared_head.head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.num_fused_shared_experts = ( - 0 if get_global_server_args().disable_shared_experts_fusion else 1 + 0 if get_flags().disable_shared_experts_fusion else 1 ) @torch.no_grad() diff --git a/python/sglang/srt/models/glm4_moe_nextn.py b/python/sglang/srt/models/glm4_moe_nextn.py index 3eeeecf6a..4fa88ba09 100644 --- a/python/sglang/srt/models/glm4_moe_nextn.py +++ b/python/sglang/srt/models/glm4_moe_nextn.py @@ -32,7 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_npu @@ -141,12 +141,12 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("model.shared_head.head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.num_fused_shared_experts = ( - 0 if get_global_server_args().disable_shared_experts_fusion else 1 + 0 if get_flags().disable_shared_experts_fusion else 1 ) @torch.no_grad() diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py index b76ee11c1..83d55414e 100644 --- a/python/sglang/srt/models/glm4v_moe.py +++ b/python/sglang/srt/models/glm4v_moe.py @@ -18,7 +18,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.glm4_moe import Glm4MoeModel from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel 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 @@ -70,7 +70,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) else: # ranks other than the last rank will have a placeholder layer @@ -84,7 +84,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): self.capture_aux_hidden_states = False def determine_num_fused_shared_experts(self): - if get_global_server_args().disable_shared_experts_fusion: + if get_flags().disable_shared_experts_fusion: return disable_reason = None @@ -100,7 +100,12 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): disable_reason = "Shared experts fusion is not supported when Deepep MoE backend is enabled." if disable_reason is not None: - get_global_server_args().disable_shared_experts_fusion = True + from sglang.srt.arg_groups.overrides import declare_load_time_override + + declare_load_time_override( + "Glm4vMoeForConditionalGeneration.determine_num_fused_shared_experts", + {"disable_shared_experts_fusion": True}, + ) log_info_on_rank0( logger, f"{disable_reason} Shared experts fusion optimization is disabled.", diff --git a/python/sglang/srt/models/glm_ocr_nextn.py b/python/sglang/srt/models/glm_ocr_nextn.py index af09d1af7..674509412 100644 --- a/python/sglang/srt/models/glm_ocr_nextn.py +++ b/python/sglang/srt/models/glm_ocr_nextn.py @@ -33,8 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.glm4 import Glm4DecoderLayer from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -135,12 +134,12 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration): config.hidden_size, quant_config=quant_config, prefix=add_prefix("model.shared_head.head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.num_fused_shared_experts = ( - 0 if get_global_server_args().disable_shared_experts_fusion else 1 + 0 if get_flags().disable_shared_experts_fusion else 1 ) @torch.no_grad() diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 7a9be4fc1..c7dadfc43 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -68,7 +68,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, @@ -390,7 +390,7 @@ class GptOssAttention(nn.Module): # Choose dtype of sinks based on attention backend: trtllm_mha requires float32, # others can use bfloat16 - attn_backend = get_global_server_args().attention_backend + attn_backend = get_flags().attn.backend sinks_dtype = torch.float32 if attn_backend == "trtllm_mha" else torch.bfloat16 self.sinks = nn.Parameter( torch.empty(self.num_heads, dtype=sinks_dtype), requires_grad=False @@ -745,7 +745,7 @@ class GptOssForCausalLM(nn.Module): config.hidden_size, # quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.capture_aux_hidden_states = False diff --git a/python/sglang/srt/models/laguna.py b/python/sglang/srt/models/laguna.py index b2f6b0e63..81b5e29af 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -53,7 +53,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.utils import apply_qk_norm -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import LazyValue, add_prefix, make_layers @@ -649,7 +649,7 @@ class LagunaForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) else: self.lm_head = PPMissingLayer() diff --git a/python/sglang/srt/models/llada2.py b/python/sglang/srt/models/llada2.py index 2c0601dac..c795061c8 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -76,7 +76,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, @@ -796,7 +796,7 @@ class LLaDA2MoeModelLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config, return_full_logits=True) diff --git a/python/sglang/srt/models/llama.py b/python/sglang/srt/models/llama.py index d88789099..12721be02 100644 --- a/python/sglang/srt/models/llama.py +++ b/python/sglang/srt/models/llama.py @@ -52,8 +52,7 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, maybe_remap_kv_scale_name, ) -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix, is_cuda, is_npu, is_xpu, make_layers from sglang.utils import get_exception_traceback @@ -502,7 +501,7 @@ class LlamaForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) diff --git a/python/sglang/srt/models/longcat_flash.py b/python/sglang/srt/models/longcat_flash.py index 29d3258c4..2a4d5f384 100644 --- a/python/sglang/srt/models/longcat_flash.py +++ b/python/sglang/srt/models/longcat_flash.py @@ -86,8 +86,7 @@ from sglang.srt.model_loader.utils import ( ) from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import ( BumpAllocator, add_prefix, @@ -616,7 +615,7 @@ class LongcatFlashForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.capture_aux_hidden_states = False diff --git a/python/sglang/srt/models/mimo_v2.py b/python/sglang/srt/models/mimo_v2.py index 87e004a63..174f6fd69 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -76,7 +76,7 @@ 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 +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, @@ -1041,7 +1041,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) else: self.lm_head = PPMissingLayer() diff --git a/python/sglang/srt/models/mimo_v2_nextn.py b/python/sglang/srt/models/mimo_v2_nextn.py index efeaddd2f..7de4e20d3 100644 --- a/python/sglang/srt/models/mimo_v2_nextn.py +++ b/python/sglang/srt/models/mimo_v2_nextn.py @@ -44,8 +44,7 @@ from sglang.srt.models.mimo_v2 import ( MiMoV2MLP, load_mimo_v2_qkv_proj_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_flags, get_parallel from sglang.srt.utils import add_prefix MiMoV2Config = None @@ -260,7 +259,7 @@ class MiMoV2MTP(MiMoV2ForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index e02a5d8ed..3eb70777f 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -87,7 +87,7 @@ from sglang.srt.models.nemotron_h_utils import ( pad_to_original_num_tokens, ) from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, @@ -880,7 +880,7 @@ class NemotronHForCausalLM(nn.Module): else lora_config.lora_vocab_padding_size ), quant_config=quant_config, - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, prefix=add_prefix("lm_head", prefix), ) else: diff --git a/python/sglang/srt/models/nemotron_h_mtp.py b/python/sglang/srt/models/nemotron_h_mtp.py index 71759bcb0..189e364b9 100644 --- a/python/sglang/srt/models/nemotron_h_mtp.py +++ b/python/sglang/srt/models/nemotron_h_mtp.py @@ -39,8 +39,7 @@ from sglang.srt.models.nemotron_h import ( NemotronHMoEDecoderLayer, ) from sglang.srt.models.nemotron_h_utils import is_attn_layer -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix @@ -340,7 +339,7 @@ class NemotronHForCausalLMMTP(NemotronHForCausalLM): self.config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index c79bbf7f1..d5aac381d 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -93,7 +93,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, @@ -148,7 +148,7 @@ def can_fuse_shared_expert( Caller must still gate on the model/backend support flag. """ if ( - get_global_server_args().disable_shared_experts_fusion is True + get_flags().disable_shared_experts_fusion is True or getattr(config, "shared_expert_intermediate_size", 0) <= 0 or config.shared_expert_intermediate_size != config.moe_intermediate_size or get_moe_a2a_backend().is_deepep() @@ -1003,7 +1003,7 @@ class Qwen2MoeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) # For EAGLE3 support diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index dfa440621..2dde70c36 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -33,7 +33,7 @@ from sglang.srt.model_loader.weight_utils import ( from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP from sglang.srt.models.qwen2 import Qwen2Model from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu @@ -493,7 +493,7 @@ class Qwen3ForCausalLM(nn.Module): config.vocab_size, config.hidden_size, quant_config=quant_config, - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, prefix=add_prefix("lm_head", prefix), ) else: diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index db88f9c12..217cece66 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -91,8 +91,7 @@ from sglang.srt.models.utils import ( fused_qk_gemma_rmsnorm, fused_qk_gemma_rmsnorm_with_gate, ) -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_flags, get_parallel # Utils from sglang.srt.utils import ( @@ -134,7 +133,7 @@ cached_get_processor = lru_cache(get_processor) def _disable_shared_experts_fusion() -> bool: # Resolved lazily: the global server args is not set at module import time # (e.g. when this module is imported by unit tests). - return get_global_server_args().disable_shared_experts_fusion + return get_flags().disable_shared_experts_fusion if _is_cuda: @@ -1172,13 +1171,17 @@ class Qwen3_5ForCausalLM(nn.Module): def _maybe_autodisable_shared_experts_fusion(self, config, quant_config): # Auto-disable fusion when the checkpoint can't fuse (e.g. MXFP4 Qwen3.5) # so the model still gets the #25885 multi-streaming path. ROCm-only. - server_args = get_global_server_args() if ( config.model_type == "qwen3_5_moe_text" - and not server_args.disable_shared_experts_fusion + and not get_flags().disable_shared_experts_fusion and not can_fuse_shared_expert(config, quant_config) ): - server_args.disable_shared_experts_fusion = True + from sglang.srt.arg_groups.overrides import declare_load_time_override + + declare_load_time_override( + "Qwen3_5ForCausalLM._maybe_autodisable_shared_experts_fusion", + {"disable_shared_experts_fusion": True}, + ) logger.info( "Qwen3.5: shared-expert fusion not supported for this checkpoint; " "auto-disabling (multi-streaming #25885 still applies)." diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py index 8bc637e4f..feb2bd1d6 100644 --- a/python/sglang/srt/models/qwen3_5_mtp.py +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -34,7 +34,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_npu @@ -148,7 +148,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module): if ( is_npu() and self.quant_config is None - and get_global_server_args().quantization is not None + and get_flags().quantization is not None ): # ascend mtp unquant exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 6b749191e..a7d0ae45e 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -72,7 +72,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, @@ -960,7 +960,7 @@ class Qwen3MoeForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) self.capture_aux_hidden_states = False diff --git a/python/sglang/srt/models/qwen3_moe_mtp.py b/python/sglang/srt/models/qwen3_moe_mtp.py index 973d1adde..48d9e4888 100644 --- a/python/sglang/srt/models/qwen3_moe_mtp.py +++ b/python/sglang/srt/models/qwen3_moe_mtp.py @@ -30,8 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM, Qwen3MoeModel -from sglang.srt.runtime_context import get_parallel -from sglang.srt.server_args import get_global_server_args +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -64,7 +63,7 @@ class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index f085fb43d..e86855cce 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -47,8 +47,7 @@ from sglang.srt.model_loader.weight_utils import ( sharded_weight_loader, ) from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock -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_flags, get_parallel from sglang.srt.utils import ( LazyValue, add_prefix, @@ -1028,7 +1027,7 @@ class Qwen3NextForCausalLM(nn.Module): quant_config=quant_config, org_num_embeddings=config.vocab_size, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) # For EAGLE3 support diff --git a/python/sglang/srt/models/qwen3_next_mtp.py b/python/sglang/srt/models/qwen3_next_mtp.py index 69adf12a7..bc2525e34 100644 --- a/python/sglang/srt/models/qwen3_next_mtp.py +++ b/python/sglang/srt/models/qwen3_next_mtp.py @@ -32,7 +32,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_npu @@ -84,7 +84,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): config.hidden_size, quant_config=quant_config, prefix=add_prefix("model.shared_head.head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) # Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so @@ -114,7 +114,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): if ( is_npu() and self.quant_config is None - and get_global_server_args().quantization is not None + and get_flags().quantization is not None ): # ascend mtp unquant exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index 6a89e803a..75c807a5d 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -68,7 +68,7 @@ from sglang.srt.models.utils import ( ) from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( add_prefix, @@ -1278,7 +1278,7 @@ class Qwen3VLForConditionalGeneration(nn.Module): self.config.vocab_size, self.config.hidden_size, quant_config=quant_config, - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, prefix=add_prefix("lm_head", prefix), ) else: diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index 1e9a8c82f..a56f78438 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -60,7 +60,7 @@ from sglang.srt.models.bailing_moe import BailingMoEForCausalLM from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import ( DeepseekMHAForwardMixin, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( BumpAllocator, @@ -1241,7 +1241,7 @@ class SarvamMLAForCausalLM(nn.Module): config.hidden_size, quant_config=quant_config, prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) diff --git a/python/sglang/srt/models/sdar.py b/python/sglang/srt/models/sdar.py index ddd21950a..2ac243443 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -41,7 +41,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, is_cuda, make_layers @@ -471,7 +471,7 @@ class SDARForCausalLM(nn.Module): config.vocab_size, config.hidden_size, quant_config=quant_config, - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, prefix=add_prefix("lm_head", prefix), ) else: diff --git a/python/sglang/srt/models/sdar_moe.py b/python/sglang/srt/models/sdar_moe.py index 02a3ad6c5..0bb16250a 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -57,7 +57,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers @@ -566,7 +566,7 @@ class SDARMoeForCausalLM(nn.Module): config.vocab_size, config.hidden_size, quant_config=quant_config, - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, prefix=add_prefix("lm_head", prefix), ) else: diff --git a/python/sglang/srt/models/step3p5.py b/python/sglang/srt/models/step3p5.py index cdc8944b9..a2a5a53f3 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -46,7 +46,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.runtime_context import get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel 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 @@ -826,7 +826,7 @@ class Step3p5ForCausalLM(nn.Module): config.vocab_size, config.hidden_size, quant_config=quant_config, - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + use_attn_tp_group=get_flags().enable_dp_lm_head, prefix=add_prefix("lm_head", prefix), ) else: diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 38656810d..b68bbe4b7 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -284,6 +284,8 @@ class AttnFlags(_StaticFlags): # Resolved attention backend; the pristine user request stays on # server_args.attention_backend. backend: str | None = None + prefill_backend: str | None = None + decode_backend: str | None = None @dataclasses.dataclass @@ -299,6 +301,10 @@ class MoeFlags(_StaticFlags): class CaptureFlags(_FlagGroupBase): """Capture-time flags; never frozen (written during cuda-graph capture).""" + # Seeded from server_args at publish; a model whose _can_torch_compile is + # False clears it during warmup (the only post-publish writer). + enable_torch_compile: bool = False + @dataclasses.dataclass class Flags(_StaticFlags): @@ -325,6 +331,16 @@ class Flags(_StaticFlags): sampling_backend: str | None = None page_size: int | None = None quantization: str | None = None + disable_overlap_schedule: bool = False + uses_mamba_radix_cache: bool = False + mamba_radix_cache_strategy: str = "auto" + speculative_moe_runner_backend: str | None = None + speculative_moe_a2a_backend: str | None = None + disable_shared_experts_fusion: bool = False + kv_cache_dtype: str = "auto" + dsa_prefill_backend: str | None = None + dsa_decode_backend: str | None = None + flashinfer_allreduce_fusion_backend: str | None = None # Parallel-request fields: flat transitional home, to be re-homed by the # Parallel Parameters Clarification module. enable_dp_attention: bool = False @@ -348,6 +364,8 @@ class Flags(_StaticFlags): # family as readers migrate. FLAG_LEAF_MAP: dict[str, str] = { "attention_backend": "attn.backend", + "prefill_attention_backend": "attn.prefill_backend", + "decode_attention_backend": "attn.decode_backend", "moe_runner_backend": "moe.runner_backend", } @@ -368,12 +386,16 @@ class RuntimeContext: """Container for the structured runtime accessors; exposes ``parallel``, ``server_args``, and ``flags``.""" - __slots__ = ("parallel", "_server_args", "flags") + __slots__ = ("parallel", "_server_args", "flags", "_runtime_overrides") def __init__(self, parallel: ParallelContext): self.parallel = parallel self._server_args: ServerArgs | None = None self.flags = Flags() + # Post-publish resolution declarations (runner- and load-time + # resolved fields), replayed after the publish-time stash on + # every re-resolve. Cleared on (re-)publish and reset. + self._runtime_overrides: list[tuple[str, dict]] = [] @property def server_args(self) -> ServerArgs: @@ -395,14 +417,90 @@ class RuntimeContext: into the flags tier (skipped for objects without the stash — dummy / "none" fixture ServerArgs and test-kit mocks never compute it). Resolution runs first: if it fails, the previous publish stays intact. + A publish after ``freeze_flags()`` is an ordering violation and raises. """ - self._resolve_flags(server_args) + if self.flags.frozen: + raise RuntimeError( + "set_server_args() after freeze_flags(): the flags tier is " + "frozen for this process; use reset_context() in tests." + ) + # A (re-)publish starts a fresh resolution lifecycle; a failed + # resolve keeps the previous lifecycle (including its recorded + # runtime overrides) intact. + saved_runtime_overrides = self._runtime_overrides + self._runtime_overrides = [] + try: + self._resolve_flags(server_args) + except BaseException: + self._runtime_overrides = saved_runtime_overrides + raise + # Seed the capture tier for the new lifecycle (defaults for sentinel + # and mock publishes, which carry no config). + self.flags.capture.enable_torch_compile = getattr( + server_args, "enable_torch_compile", False + ) self._server_args = server_args + def record_runtime_overrides( + self, entries: list[tuple[str, dict]] + ) -> list[tuple[str, dict]]: + """Append post-publish resolution declarations (the runner- and + load-time stages) and + atomically re-resolve the flags tier. + + Target-worker only, and only before ``freeze_flags()``. During the + dual-apply transition the call sites keep their imperative + ``server_args`` writes; the recorded declarations must match them — + parity is re-asserted on every declared field. On failure the + recorded entries are rolled back and the previous flags stay + installed. + """ + server_args = self._server_args + if server_args is None: + raise ValueError("Global server args is not set yet!") + if self.flags.frozen: + raise RuntimeError( + "record_runtime_overrides() after freeze_flags(): runtime " + "resolution stages must complete before the flags tier " + "freezes." + ) + entries = [(source, dict(declared)) for source, declared in entries] + self._runtime_overrides.extend(entries) + try: + self._resolve_flags(server_args) + except BaseException: + del self._runtime_overrides[len(self._runtime_overrides) - len(entries) :] + raise + return entries + + def freeze_flags(self) -> None: + """Lock every static flag group (the resolution end point: after the + load-time stages, before serving). ``flags.capture`` stays writable.""" + self.flags.freeze() + def _resolve_flags(self, server_args: ServerArgs) -> None: declarations = getattr(server_args, "_resolved_overrides", None) - if declarations is None: - return + if declarations is None and not self._runtime_overrides: + # Stash-less publish. For a config-shaped object (a dataclass: + # mock ServerArgs fixtures, dummy-path instances that skipped the + # monolith) still materialize the whitelist from its own fields, + # so flag reads match legacy server_args reads. Skip only for + # field-less sentinels (tests publishing object()). + if not dataclasses.is_dataclass(server_args): + return + from sglang.srt.arg_groups.arg_utils import resolvable_fields + + if any( + field not in vars(server_args) + for field in resolvable_fields(type(server_args)) + ): + # Bare object.__new__ fixtures: dataclass defaults live on + # the class, not the instance — nothing was populated, so + # treat it as a sentinel (hasattr would see the class + # defaults and materialize them, clobbering resolved flags). + return + declarations = () + declarations = list(declarations or ()) + self._runtime_overrides from sglang.srt.arg_groups.overrides import ( apply_model_overrides, assert_flag_parity, @@ -422,6 +520,10 @@ class RuntimeContext: server_args, {field for _source, decl in declarations for field in decl}, ) + # The capture tier is not part of the static resolution: carry it + # across re-resolves so runtime-stage recording cannot clobber a + # capture-time write (set_server_args re-seeds it per lifecycle). + flags.capture = self.flags.capture self.flags = flags @@ -453,3 +555,4 @@ def reset_context() -> None: """ _CONTEXT._server_args = None _CONTEXT.flags = Flags() + _CONTEXT._runtime_overrides = [] diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 041ae4c7f..71ee64ef1 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -600,6 +600,7 @@ class ServerArgs: "mxfp4) is supported for CUDA 12.8+ and PyTorch 2.8.0+" ), choices=["auto", "fp8_e5m2", "fp8_e4m3", "bf16", "bfloat16", "fp4_e2m1"], + resolvable=True, ), ] = "auto" enable_fp32_lm_head: A[ @@ -809,7 +810,10 @@ class ServerArgs: ] = False disable_overlap_schedule: A[ bool, - "Disable the overlap scheduler, which overlaps the CPU scheduler with GPU model worker.", + Arg( + help="Disable the overlap scheduler, which overlaps the CPU scheduler with GPU model worker.", + resolvable=True, + ), ] = False num_continuous_decode_steps: A[ int, @@ -1426,6 +1430,7 @@ class ServerArgs: Arg( help="Choose the kernels for decode attention layers (have priority over --attention-backend).", choices=ATTENTION_BACKEND_CHOICES, + resolvable=True, ), ] = None prefill_attention_backend: A[ @@ -1433,6 +1438,7 @@ class ServerArgs: Arg( help="Choose the kernels for prefill attention layers (have priority over --attention-backend).", choices=ATTENTION_BACKEND_CHOICES, + resolvable=True, ), ] = None sampling_backend: A[ @@ -1492,6 +1498,7 @@ class ServerArgs: Arg( help="DSA (DeepSeek Sparse Attention) prefill backend. If not specified, auto-detects based on hardware and kv_cache_dtype.", choices=DSA_CHOICES, + resolvable=True, ), ] = None dsa_decode_backend: A[ @@ -1499,6 +1506,7 @@ class ServerArgs: Arg( help="DSA (DeepSeek Sparse Attention) decode backend. If not specified, auto-detects based on hardware and kv_cache_dtype.", choices=DSA_CHOICES, + resolvable=True, ), ] = None dsa_topk_backend: A[ @@ -1594,6 +1602,7 @@ class ServerArgs: Arg( help="Choose the runner backend for MoE in speculative decoding.", choices=MOE_RUNNER_BACKEND_CHOICES, + resolvable=True, ), ] = None speculative_moe_a2a_backend: A[ @@ -1601,6 +1610,7 @@ class ServerArgs: Arg( help="Choose the backend for MoE A2A in speculative decoding", choices=MOE_A2A_BACKEND_CHOICES, + resolvable=True, ), ] = None speculative_draft_model_quantization: A[ @@ -1819,7 +1829,10 @@ class ServerArgs: ] = False disable_shared_experts_fusion: A[ bool, - "Disable the built-in shared experts fusion optimization for DeepSeek V3/R1. Note: DeepEP Waterfill (--enable-deepep-waterfill) still routes shared expert through DeepEP as an extra MoE slot, so shared expert is not separated from the MoE path when Waterfill is enabled.", + Arg( + help="Disable the built-in shared experts fusion optimization for DeepSeek V3/R1. Note: DeepEP Waterfill (--enable-deepep-waterfill) still routes shared expert through DeepEP as an extra MoE slot, so shared expert is not separated from the MoE path when Waterfill is enabled.", + resolvable=True, + ), ] = False enforce_shared_experts_fusion: A[ bool, @@ -1856,8 +1869,19 @@ class ServerArgs: Arg( help="The strategy to use for mamba radix cache.", choices=MAMBA_RADIX_CACHE_STRATEGY_CHOICES, + resolvable=True, ), ] = "auto" + uses_mamba_radix_cache: A[ + bool, + Arg( + help="(Derived) whether the model routes through the hybrid-mamba " + "radix cache handling; resolved from the model architecture, no " + "CLI surface.", + no_cli=True, + resolvable=True, + ), + ] = False mamba_track_interval: A[ int, "The interval to track the mamba state during decode.", @@ -2209,6 +2233,7 @@ class ServerArgs: "single-node or multi-node systems via MNNVL fabric. " "Fuses allreduce with Residual + RMSNorm for supported MoE models." ), + resolvable=True, ), ] = None enable_aiter_allreduce_fusion: A[bool, "Enable Aiter AllReduce Fusion."] = False @@ -3676,80 +3701,26 @@ class ServerArgs: return capture_sizes - def _set_default_dsa_kv_cache_dtype(self, major: int, quantization: str) -> str: - user_set_prefill = self.dsa_prefill_backend is not None - user_set_decode = self.dsa_decode_backend is not None - - # If user specified a backend but didn't explicitly set kv_cache_dtype, - # suggest them to be explicit about kv_cache_dtype to avoid surprises - if (user_set_prefill or user_set_decode) and self.kv_cache_dtype == "auto": - logger.warning( - "When specifying --dsa-prefill-backend or --dsa-decode-backend, " - "you should also explicitly set --kv-cache-dtype (e.g., 'fp8_e4m3' or 'bfloat16'). " - "DeepSeek V3.2 defaults to FP8 KV cache which may not be compatible with all backends." - ) - - if self.kv_cache_dtype == "auto": - if major >= 10: - self.kv_cache_dtype = "fp8_e4m3" - else: - self.kv_cache_dtype = "bfloat16" - logger.warning( - f"Setting KV cache dtype to {self.kv_cache_dtype} for DeepSeek DSA on SM{major} device." - ) - if self.kv_cache_dtype == "bf16": - self.kv_cache_dtype = "bfloat16" - assert self.kv_cache_dtype in [ - "bfloat16", - "fp8_e4m3", - ], "DeepSeek DSA only supports bf16/bfloat16 or fp8_e4m3 kv_cache_dtype" - - def _set_default_dsa_backends(self, kv_cache_dtype: str, major: int) -> str: - from sglang.srt.arg_groups.hisparse_hook import ( - apply_hisparse_dsa_backend_defaults, + def _set_default_dsa_kv_cache_dtype(self, major: int, quantization: str) -> None: + # Moved to the resolution pipeline (arg_groups/overrides.py: + # _dsa_kv_cache_dtype_default), invoked here at its legacy slot. + from sglang.srt.arg_groups.overrides import ( + _dsa_kv_cache_dtype_default, + run_post_process_pass, ) - user_set_prefill = self.dsa_prefill_backend is not None - user_set_decode = self.dsa_decode_backend is not None + run_post_process_pass(self, _dsa_kv_cache_dtype_default) - if apply_hisparse_dsa_backend_defaults( - self, user_set_prefill, user_set_decode, kv_cache_dtype - ): - return - - if not user_set_prefill and not user_set_decode and is_hip(): - self.dsa_prefill_backend = "tilelang" - self.dsa_decode_backend = "tilelang" - elif kv_cache_dtype == "fp8_e4m3": - if major >= 10: - if not user_set_prefill: - self.dsa_prefill_backend = "trtllm" - if not user_set_decode: - self.dsa_decode_backend = "trtllm" - else: - # Hopper FP8 defaults to flashmla_kv for both prefill and decode. - if not user_set_prefill: - self.dsa_prefill_backend = "flashmla_kv" - if not user_set_decode: - self.dsa_decode_backend = "flashmla_kv" - else: - # set prefill/decode backends based on hardware architecture. - if major >= 10: - if not user_set_prefill: - self.dsa_prefill_backend = "flashmla_sparse" - if not user_set_decode: - self.dsa_decode_backend = "trtllm" - else: - # Hopper defaults for bfloat16 - if not user_set_prefill: - self.dsa_prefill_backend = "flashmla_sparse" - if not user_set_decode: - self.dsa_decode_backend = "fa3" - - logger.warning( - f"Set DSA backends for {self.kv_cache_dtype} KV Cache: prefill={self.dsa_prefill_backend}, decode={self.dsa_decode_backend}." + def _set_default_dsa_backends(self, kv_cache_dtype: str, major: int) -> None: + # Moved to the resolution pipeline (arg_groups/overrides.py: + # _dsa_split_backend_resolution), invoked here at its legacy slot. + from sglang.srt.arg_groups.overrides import ( + _dsa_split_backend_resolution, + run_post_process_pass, ) + run_post_process_pass(self, _dsa_split_backend_resolution) + def _validate_hisparse_dsa_backend(self, attr: str, label: str): from sglang.srt.arg_groups.hisparse_hook import validate_hisparse_dsa_backend @@ -3906,35 +3877,15 @@ class ServerArgs: "Enable Aiter AllReduce Fusion for DeepseekV3ForCausalLM" ) - if ( - self.quantization == "modelopt_fp4" - and self.speculative_algorithm == "EAGLE" - and ( - self.speculative_moe_runner_backend is None - or self.speculative_moe_a2a_backend is None - ) - ): - if envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE.get(): - self.speculative_moe_runner_backend = "deep_gemm" - self.speculative_moe_a2a_backend = "deepep" - logger.info( - "Use deep_gemm moe runner and deepep a2a backend for bf16 nextn layer in deepseek fp4 checkpoint." - ) - # Validate usage of ep - if self.ep_size == 1: - raise ValueError( - "Invalid configuration: 'deep_gemm' speculative MoE runner backend with " - "'deepep' a2a backend requires expert parallelism (ep_size > 1). " - f"Current ep_size is {self.ep_size}. " - "Please set --ep-size > 1 (e.g., --ep-size 8) to use this configuration, " - "or change --speculative-moe-a2a-backend to 'none' if expert parallelism is not available." - ) - else: - self.speculative_moe_runner_backend = "triton" - self.speculative_moe_a2a_backend = "none" - logger.info( - "Use triton fused moe by default for bf16 nextn layer in deepseek fp4 checkpoint." - ) + # The fp4-checkpoint draft spec-MoE resolution moved to the + # resolution pipeline (arg_groups/overrides.py: + # _deepseek_spec_moe_resolution), invoked here at its legacy + # slot. + from sglang.srt.arg_groups.overrides import ( + _deepseek_spec_moe_resolution, + ) + + run_post_process_pass(self, _deepseek_spec_moe_resolution) elif model_arch in [ "DeepseekV4ForCausalLM", @@ -3943,12 +3894,16 @@ class ServerArgs: validate_deepseek_v4_cp(self) + # The SM120 marlin fallback moved to the resolution pipeline + # (arg_groups/overrides.py: _deepseek_v4_sm120_moe), invoked here + # at its legacy slot. + from sglang.srt.arg_groups.overrides import ( + _deepseek_v4_sm120_moe, + run_post_process_pass, + ) + + run_post_process_pass(self, _deepseek_v4_sm120_moe) if is_sm120_supported(): - if self.moe_runner_backend == "auto": - self.moe_runner_backend = "marlin" - logger.info( - "Use marlin as MoE runner backend on SM120 for DeepseekV4" - ) # SM120 lacks tcgen05/TMEM: disable features that depend on # DeepGEMM or require >99KB SMEM (topk_v2). envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False) @@ -4099,16 +4054,9 @@ class ServerArgs: # The quantization/moe_runner_backend resolution moved to the override # registry (arg_groups/overrides.py: _gemma4_overrides). elif model_arch == "MossVLForConditionalGeneration": - if self.is_attention_backend_not_set(): - self.prefill_attention_backend = "flashinfer" - logger.info( - "Use flashinfer as default prefill attention backend for Moss-VL" - ) - prefill_backend, _ = self.get_attention_backends() - assert prefill_backend == "flashinfer", ( - "MossVLForConditionalGeneration requires flashinfer prefill " - "attention backend for cross-attention custom mask support." - ) + # The prefill attention backend default + validation moved to the + # override registry (arg_groups/overrides.py: _moss_vl_overrides). + pass elif model_arch in ["Exaone4ForCausalLM", "ExaoneMoEForCausalLM"]: if hf_config.sliding_window_pattern is not None: # disable_hybrid_swa_memory moved to the override registry @@ -4132,16 +4080,14 @@ class ServerArgs: logger.info( f"Using {self.attention_backend} as attention backend for {model_arch}." ) - elif model_arch in ["KimiLinearForCausalLM"]: - self._handle_mamba_radix_cache(model_arch=model_arch) - elif model_arch in ["BailingMoeV2_5ForCausalLM"]: - self._handle_mamba_radix_cache(model_arch=model_arch) elif model_arch in ["NemotronHForCausalLM", "NemotronHPuzzleForCausalLM"]: - from sglang.srt.arg_groups.nemotron_h_hook import ( - apply_nemotron_h_defaults, + # Quantization / MoE runner / attention backend defaults moved to + # the override registry (arg_groups/overrides.py: + # _nemotron_h_overrides). + assert self.attention_backend != "triton", ( + "NemotronHForCausalLM does not support triton attention backend," + "as the first layer might not be an attention layer" ) - - apply_nemotron_h_defaults(self, model_arch) elif model_arch in [ "Qwen3MoeForCausalLM", "Qwen3VLMoeForConditionalGeneration", @@ -4150,25 +4096,11 @@ class ServerArgs: "InternS2PreviewForConditionalGeneration", "Qwen3_5ForConditionalGeneration", ]: - # The quantization/moe_runner_backend resolution moved to the override - # registry (arg_groups/overrides.py: _qwen3_moe_family_overrides). - - if model_arch in [ - "Qwen3NextForCausalLM", - "Qwen3_5MoeForConditionalGeneration", - "InternS2PreviewForConditionalGeneration", - "Qwen3_5ForConditionalGeneration", - ]: - # Attention backend + page size defaults moved to the override - # registry (arg_groups/overrides.py: _qwen3_5_hybrid_overrides). - self._handle_mamba_radix_cache(model_arch=model_arch) - - elif model_arch == "MiniCPMV4_6ForConditionalGeneration": - # 4.6 wraps a Qwen3.5 hybrid GDN backbone, so it needs the same - # mamba radix cache handling as Qwen3_5ForConditionalGeneration. - # (attention backend selection moved to the override registry: - # arg_groups/overrides.py _minicpm_v4_6_overrides) - self._handle_mamba_radix_cache(model_arch=model_arch) + # The quantization/moe_runner_backend resolution moved to the + # override registry (arg_groups/overrides.py: + # _qwen3_moe_family_overrides); the hybrid sub-family's attention + # backend + page size defaults to _qwen3_5_hybrid_overrides. + pass elif model_arch in ["Glm4MoeForCausalLM"]: # The quantization/moe_runner_backend/enable_tf32_matmul resolution @@ -4176,110 +4108,53 @@ class ServerArgs: # _glm4_moe_overrides). pass - elif model_arch in [ - "FalconH1ForCausalLM", - "JetNemotronForCausalLM", - "JetVLMForConditionalGeneration", - ]: - # Attention backend selection moved to the override registry - # (arg_groups/overrides.py: _falcon_h1_jet_overrides). - self._handle_mamba_radix_cache(model_arch=model_arch) - - elif model_arch == "GraniteMoeHybridForCausalLM": - hf_config = self.get_model_config().hf_config - has_mamba = any( - layer_type == "mamba" - for layer_type in getattr(hf_config, "layer_types", []) - ) - if has_mamba: - # Attention backend selection moved to the override registry - # (arg_groups/overrides.py: _granite_moe_hybrid_overrides). - self._handle_mamba_radix_cache(model_arch=model_arch) - elif model_arch in ["Lfm2ForCausalLM"]: # Attention backend selection moved to the override registry # (arg_groups/overrides.py: _lfm2_overrides). - self._handle_mamba_radix_cache(model_arch=model_arch) assert self.attention_backend != "triton", ( f"{model_arch} does not support triton attention backend, " "as the first layer might not be an attention layer" ) - elif model_arch in ["ZayaForCausalLM"]: - self._handle_mamba_radix_cache(model_arch=model_arch) - # MiniMaxM2ForCausalLM (enable_tf32_matmul) moved to the override registry # (arg_groups/overrides.py: _minimax_m2_overrides). # Qwen3VL aiter unified-attention page_size moved to the override registry # (arg_groups/overrides.py: _qwen3vl_overrides). - if envs.SGLANG_EMBEDDINGS_SPARSE_HEAD.is_set(): - self.disable_overlap_schedule = True - logger.warning( - "Overlap scheduler is disabled when using sparse head for embedding model." - ) + # Hybrid-mamba radix cache handling for the per-arch branch call sites + # dissolved above: the resolution pass self-guards on the arch union + # (and the Granite layer_types probe), so one call covers them all. + # Hybrid-spec archs already resolved at the pre-dispatch call above; + # for them this re-invocation is an idempotent no-op plus validation. + # Kept ahead of the sparse-head pass: the legacy per-branch calls + # resolved before that tail write of disable_overlap_schedule. + self._handle_mamba_radix_cache(model_arch=model_arch) - # Auto-enable FlashInfer AllReduce Fusion on SM90/SM100, for models with - # explicit support (DeepseekV3, GptOss, Glm4Moe, MistralLarge3, - # Qwen3/Qwen3-VL/Qwen3Next/Qwen3.5 MoE families). auto resolves to mnnvl on - # Blackwell (single- and multi-node) and trtllm on SM90 single-node systems. - if ( - self.flashinfer_allreduce_fusion_backend is None - and model_arch - in [ - "DeepseekV3ForCausalLM", - "DeepseekV32ForCausalLM", - "GptOssForCausalLM", - "GlmMoeDsaForCausalLM", - "Glm4MoeForCausalLM", - "Glm4MoeLiteForCausalLM", - "MistralLarge3ForCausalLM", - "Qwen3MoeForCausalLM", - "Qwen3VLMoeForConditionalGeneration", - "Qwen3NextForCausalLM", - "KimiK25ForConditionalGeneration", - "Qwen3_5MoeForConditionalGeneration", - "InternS2PreviewForConditionalGeneration", - "Qwen3_5ForConditionalGeneration", - "NemotronHForCausalLM", - "NemotronHPuzzleForCausalLM", - ] - and (is_sm90_supported() or is_sm100_supported()) - and self.tp_size > 1 - and not self.enable_dp_attention - and (self.nnodes == 1 or is_sm100_supported()) - and self.moe_a2a_backend == "none" - ): - self.flashinfer_allreduce_fusion_backend = "auto" - logger.info( - f"Auto-enabling FlashInfer AllReduce Fusion on SM90/SM10X for {model_arch}" - ) + from sglang.srt.arg_groups.overrides import ( + _sparse_head_overlap_disable, + run_post_process_pass, + ) - # Apply enforce_disable_flashinfer_allreduce_fusion after all model-specific adjustments - if self.enforce_disable_flashinfer_allreduce_fusion: - self.flashinfer_allreduce_fusion_backend = None - logger.info( - "FlashInfer allreduce fusion is forcibly disabled " - "via --enforce-disable-flashinfer-allreduce-fusion." - ) + run_post_process_pass(self, _sparse_head_overlap_disable) + + # The FlashInfer AllReduce Fusion auto-enable and the enforce-disable + # terminal moved to the resolution pipeline (arg_groups/overrides.py: + # _flashinfer_allreduce_fusion_auto_enable / + # _enforce_disable_allreduce_fusion), invoked here at their legacy + # slots. + from sglang.srt.arg_groups.overrides import ( + _enforce_disable_allreduce_fusion, + _flashinfer_allreduce_fusion_auto_enable, + ) + + run_post_process_pass(self, _flashinfer_allreduce_fusion_auto_enable) + run_post_process_pass(self, _enforce_disable_allreduce_fusion) def _support_mamba_cache_extra_buffer(self, model_arch: str): - if model_arch in [ - "Qwen3_5ForConditionalGeneration", - "Qwen3_5MoeForConditionalGeneration", - "Qwen3NextForCausalLM", - "InternS2PreviewForConditionalGeneration", - "MiniCPMV4_6ForConditionalGeneration", - "BailingMoeV2_5ForCausalLM", - "FalconH1ForCausalLM", - "GraniteMoeHybridForCausalLM", - "NemotronHForCausalLM", - "NemotronHPuzzleForCausalLM", - ]: - return self.linear_attn_backend == "triton" + from sglang.srt.arg_groups.overrides import supports_mamba_cache_extra_buffer - return False + return supports_mamba_cache_extra_buffer(self, model_arch) def _validate_mamba_no_buffer(self, model_arch: str): assert self.page_size in (1, None), "no_buffer only supports page_size=1." @@ -4307,20 +4182,17 @@ class ServerArgs: assert self.mamba_cache_chunk_size is not None def _handle_mamba_radix_cache(self, model_arch: str): - if self.disable_radix_cache: - return + # Resolution moved to the resolution pipeline (arg_groups/overrides.py: + # _mamba_radix_cache_resolution), invoked here at each legacy call + # slot; this handler keeps the validation. + from sglang.srt.arg_groups.overrides import ( + _mamba_radix_cache_resolution, + run_post_process_pass, + ) - self.uses_mamba_radix_cache = True - if self.mamba_radix_cache_strategy == "auto": - wants_overlap = not self.disable_overlap_schedule - wants_paging = self.page_size is not None and self.page_size > 1 - if ( - wants_overlap or wants_paging - ) and self._support_mamba_cache_extra_buffer(model_arch): - self.mamba_radix_cache_strategy = "extra_buffer" - else: - self.mamba_radix_cache_strategy = "no_buffer" - self.disable_overlap_schedule = True + run_post_process_pass(self, _mamba_radix_cache_resolution) + if not self.uses_mamba_radix_cache: + return if self.enable_mamba_extra_buffer(): self._validate_mamba_extra_buffer(model_arch) @@ -4487,29 +4359,12 @@ class ServerArgs: f"got {self.kv_cache_dtype}." ) - if ( - self.attention_backend == "cutedsl_mla" - or self.decode_attention_backend == "cutedsl_mla" - or self.prefill_attention_backend == "cutedsl_mla" - ): - assert ( - self.prefill_attention_backend != "cutedsl_mla" - ), "CuteDSL MLA only supports decoding for now" - if not is_sm100_supported(): - raise ValueError( - "CuteDSL MLA backend is only supported on Blackwell GPUs (SM100). Please use a different backend." - ) - if self.kv_cache_dtype not in [ - "fp8_e4m3", - "bf16", - "bfloat16", - "auto", - ]: - raise ValueError( - "CuteDSL MLA backend only supports kv-cache-dtype of fp8_e4m3, bf16, or auto." - ) - if self.prefill_attention_backend is None: - self.prefill_attention_backend = "trtllm_mla" + # The CuteDSL MLA validation + prefill fill moved to the resolution + # pipeline (arg_groups/overrides.py: _cutedsl_prefill_backend_fill), + # invoked here at its legacy slot. + from sglang.srt.arg_groups.overrides import _cutedsl_prefill_backend_fill + + run_post_process_pass(self, _cutedsl_prefill_backend_fill) if ( self.attention_backend == "trtllm_mha" @@ -5315,11 +5170,14 @@ class ServerArgs: self.expert_distribution_recorder_buffer_size = 1000 def _handle_pipeline_parallelism(self): - if self.pp_size > 1: - self.disable_overlap_schedule = True - logger.warning( - "Pipeline parallelism is incompatible with overlap schedule." - ) + # Moved to the resolution pipeline (arg_groups/overrides.py: + # _pipeline_parallel_overlap_disable), invoked here at its legacy slot. + from sglang.srt.arg_groups.overrides import ( + _pipeline_parallel_overlap_disable, + run_post_process_pass, + ) + + run_post_process_pass(self, _pipeline_parallel_overlap_disable) def _validate_prefill_only_disable_kv_cache_args(self): """Validate --prefill-only-disable-kv-cache flag/precondition constraints. @@ -5880,11 +5738,15 @@ class ServerArgs: ) self.enable_aiter_allreduce_fusion = False - if self.flashinfer_allreduce_fusion_backend is not None: - logger.warning( - "Disable --flashinfer-allreduce-fusion-backend because deterministic inference is enabled." - ) - self.flashinfer_allreduce_fusion_backend = None + # Moved to the resolution pipeline (arg_groups/overrides.py: + # _deterministic_allreduce_fusion_disable), invoked here at its + # legacy slot. + from sglang.srt.arg_groups.overrides import ( + _deterministic_allreduce_fusion_disable, + run_post_process_pass, + ) + + run_post_process_pass(self, _deterministic_allreduce_fusion_disable) # The forced-pytorch sampling write and the attention backend # fill/validation moved to the resolution pipeline @@ -6040,16 +5902,12 @@ class ServerArgs: from sglang.srt.arg_groups.overrides import ( _dllm_attention_backend, + _dllm_overlap_disable, run_post_process_pass, ) run_post_process_pass(self, _dllm_attention_backend) - - if not self.disable_overlap_schedule: - logger.warning( - "Overlap schedule is disabled because of using diffusion LLM inference" - ) - self.disable_overlap_schedule = True + run_post_process_pass(self, _dllm_overlap_disable) if not self.disable_radix_cache: # The page_size adjustment moved to the resolution pipeline diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 480930242..c801bbe8c 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -35,6 +35,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) +from sglang.srt.runtime_context import get_flags from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec @@ -107,7 +108,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): self.tp_size = model_runner.tp_size self.dp_size = model_runner.dp_size self.pp_size = model_runner.server_args.pp_size - self.enable_torch_compile = model_runner.server_args.enable_torch_compile + self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 03662c5da..af71fc36f 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -35,6 +35,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) +from sglang.srt.runtime_context import get_flags from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.spec_utils import fast_topk @@ -96,7 +97,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): self.tp_size = model_runner.tp_size self.dp_size = model_runner.dp_size self.pp_size = model_runner.server_args.pp_size - self.enable_torch_compile = model_runner.server_args.enable_torch_compile + self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 4ad723bcc..4c7db0ffc 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -32,6 +32,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) +from sglang.srt.runtime_context import get_flags from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput from sglang.srt.utils import ( require_attn_tp_gather, @@ -84,7 +85,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): self.device = model_runner.device self.device_module = torch.get_device_module(self.device) - self.enable_torch_compile = model_runner.server_args.enable_torch_compile + self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 1f325c779..7ad662cba 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -59,6 +59,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) +from sglang.srt.runtime_context import get_flags from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.spec_utils import fast_topk @@ -127,7 +128,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): self.tp_size = model_runner.tp_size self.dp_size = model_runner.server_args.dp_size self.pp_size = model_runner.server_args.pp_size - self.enable_torch_compile = model_runner.server_args.enable_torch_compile + self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) diff --git a/test/registered/ops/test_aiter_greedy_sample_amd.py b/test/registered/ops/test_aiter_greedy_sample_amd.py index 662296561..4dd688685 100644 --- a/test/registered/ops/test_aiter_greedy_sample_amd.py +++ b/test/registered/ops/test_aiter_greedy_sample_amd.py @@ -23,12 +23,15 @@ register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd") def _mock_global_server_args(backend="pytorch"): from sglang.srt.layers import sampler as sampler_mod + from sglang.srt.runtime_context import get_flags from sglang.srt.server_args import ServerArgs sampler_mod.get_global_server_args = lambda: ServerArgs( model_path="dummy", sampling_backend=backend, ) + # The sampler reads the resolved backend from the flags tier. + get_flags().sampling_backend = backend class _DummyTPGroup: device_group = None diff --git a/test/registered/rl/test_fp32_lm_head.py b/test/registered/rl/test_fp32_lm_head.py index dac52502c..cc7fba620 100644 --- a/test/registered/rl/test_fp32_lm_head.py +++ b/test/registered/rl/test_fp32_lm_head.py @@ -7,6 +7,7 @@ import torch.nn as nn import torch.nn.functional as F from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.runtime_context import get_flags from sglang.srt.server_args import ( ServerArgs, get_global_server_args, @@ -44,7 +45,7 @@ class TestLMHeadFP32(unittest.TestCase): def _make_logprocessor(self, vocab_size, enable_fp32): set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) - get_global_server_args().enable_dp_lm_head = False + get_flags().enable_dp_lm_head = False get_global_server_args().enable_fp32_lm_head = enable_fp32 cfg = SimpleNamespace(vocab_size=vocab_size, final_logit_softcapping=None) return LogitsProcessor(cfg, skip_all_gather=True, logit_scale=None) diff --git a/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py b/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py index a3cd88203..9d9e66474 100644 --- a/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py +++ b/test/registered/unit/batch_overlap/test_tbo_filter_batch_marker.py @@ -38,9 +38,11 @@ def _make_target_verify_batch(bs: int) -> ForwardBatch: def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch: fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3") + from sglang.srt.runtime_context import get_flags + with get_parallel().override(attn_tp_size=1), patch.object( tbo, "get_global_server_args", lambda: fake_args - ): + ), get_flags().attn.override(backend="fa3"): return TboForwardBatchPreparer.filter_batch( batch, start_token_index=lo, diff --git a/test/registered/unit/models/test_deepseek_v4_shared_expert_fusion.py b/test/registered/unit/models/test_deepseek_v4_shared_expert_fusion.py index 48aa53ff5..b08f53615 100644 --- a/test/registered/unit/models/test_deepseek_v4_shared_expert_fusion.py +++ b/test/registered/unit/models/test_deepseek_v4_shared_expert_fusion.py @@ -1,48 +1,58 @@ import unittest from types import SimpleNamespace -from unittest.mock import patch -from sglang.srt.models import deepseek_v4 as deepseek_v4_module from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM +from sglang.srt.runtime_context import get_context, get_flags, reset_context +from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=4, suite="base-a-test-cpu") class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase): + """The disable decision is a load-time resolution: it lands on the flags + tier through declare_load_time_override (dual-applied onto the published + config during the transition).""" + + def setUp(self): + self._saved_server_args = get_context()._server_args + + def tearDown(self): + if self._saved_server_args is None: + reset_context() + else: + get_context().set_server_args(self._saved_server_args) + def _make_model(self, n_shared_experts=1): return SimpleNamespace( config=SimpleNamespace(n_shared_experts=n_shared_experts) ) + def _publish(self, enforce): + server_args = ServerArgs(model_path="dummy") + server_args.enforce_shared_experts_fusion = enforce + get_context().set_server_args(server_args) + return server_args + def test_disables_shared_fusion_without_enforce(self): - server_args = SimpleNamespace( - disable_shared_experts_fusion=False, - enforce_shared_experts_fusion=False, - ) + server_args = self._publish(enforce=False) model = self._make_model() - with patch.object( - deepseek_v4_module, "get_global_server_args", return_value=server_args - ): - DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model) + DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model) self.assertEqual(model.num_fused_shared_experts, 0) + self.assertTrue(get_flags().disable_shared_experts_fusion) + # dual-apply transition: the published config carries the value too self.assertTrue(server_args.disable_shared_experts_fusion) def test_enables_shared_fusion_when_enforced(self): - server_args = SimpleNamespace( - disable_shared_experts_fusion=False, - enforce_shared_experts_fusion=True, - ) + server_args = self._publish(enforce=True) model = self._make_model() - with patch.object( - deepseek_v4_module, "get_global_server_args", return_value=server_args - ): - DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model) + DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model) self.assertEqual(model.num_fused_shared_experts, 1) + self.assertFalse(get_flags().disable_shared_experts_fusion) self.assertFalse(server_args.disable_shared_experts_fusion) diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index eac1cb4b6..39e46700f 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -184,47 +184,75 @@ class TestLoadBalanceMethod(unittest.TestCase): class TestHiSparseDsaBackendPolicy(unittest.TestCase): + # The backend selection moved to the resolution pipeline; these policy + # tests drive the pass through its read-only view. + @staticmethod + def _resolve(kv_cache_dtype, **kw): + from types import SimpleNamespace + + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _dsa_split_backend_resolution, + ) + + hf = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"]) + defaults = dict( + kv_cache_dtype=kv_cache_dtype, + dsa_prefill_backend=None, + dsa_decode_backend=None, + enable_hisparse=True, + ) + defaults.update(kw) + view = ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + with ( + patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True), + patch("sglang.srt.arg_groups.overrides.is_npu", return_value=False), + patch("sglang.srt.arg_groups.overrides.is_xpu", return_value=False), + patch("torch.cuda.get_device_capability", return_value=(9, 0)), + ): + declared = _dsa_split_backend_resolution(view) + return { + "dsa_prefill_backend": declared.get( + "dsa_prefill_backend", defaults["dsa_prefill_backend"] + ), + "dsa_decode_backend": declared.get( + "dsa_decode_backend", defaults["dsa_decode_backend"] + ), + } + @patch("sglang.srt.server_args.is_hip", return_value=False) def test_hisparse_defaults_to_flashmla_sparse_on_cuda_bfloat16(self, _mock_is_hip): - server_args = ServerArgs(model_path="dummy", enable_hisparse=True) + resolved = self._resolve("bfloat16") - server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9) - - self.assertEqual(server_args.dsa_prefill_backend, "flashmla_sparse") - self.assertEqual(server_args.dsa_decode_backend, "flashmla_sparse") + self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_sparse") + self.assertEqual(resolved["dsa_decode_backend"], "flashmla_sparse") @patch("sglang.srt.server_args.is_hip", return_value=False) def test_hisparse_defaults_to_flashmla_kv_on_cuda_fp8(self, _mock_is_hip): - server_args = ServerArgs(model_path="dummy", enable_hisparse=True) + resolved = self._resolve("fp8_e4m3") - server_args._set_default_dsa_backends(kv_cache_dtype="fp8_e4m3", major=9) - - self.assertEqual(server_args.dsa_prefill_backend, "flashmla_kv") - self.assertEqual(server_args.dsa_decode_backend, "flashmla_kv") + self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_kv") + self.assertEqual(resolved["dsa_decode_backend"], "flashmla_kv") @patch("sglang.srt.server_args.is_hip", return_value=True) def test_hisparse_defaults_to_tilelang_on_rocm(self, _mock_is_hip): - server_args = ServerArgs(model_path="dummy", enable_hisparse=True) + resolved = self._resolve("bfloat16") - server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9) - - self.assertEqual(server_args.dsa_prefill_backend, "tilelang") - self.assertEqual(server_args.dsa_decode_backend, "tilelang") + self.assertEqual(resolved["dsa_prefill_backend"], "tilelang") + self.assertEqual(resolved["dsa_decode_backend"], "tilelang") @patch("sglang.srt.server_args.is_hip", return_value=True) def test_hisparse_preserves_rocm_user_backend_and_defaults_missing_side( self, _mock_is_hip ): - server_args = ServerArgs( - model_path="dummy", - enable_hisparse=True, - dsa_prefill_backend="tilelang", - ) + resolved = self._resolve("bfloat16", dsa_prefill_backend="tilelang") - server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9) - - self.assertEqual(server_args.dsa_prefill_backend, "tilelang") - self.assertEqual(server_args.dsa_decode_backend, "tilelang") + self.assertEqual(resolved["dsa_prefill_backend"], "tilelang") + self.assertEqual(resolved["dsa_decode_backend"], "tilelang") @patch("sglang.srt.server_args.is_hip", return_value=True) def test_hisparse_accepts_aiter_backend_on_rocm(self, _mock_is_hip): diff --git a/test/registered/unit/test_legacy_global_ratchet.py b/test/registered/unit/test_legacy_global_ratchet.py index 286a53ebd..417a96062 100644 --- a/test/registered/unit/test_legacy_global_ratchet.py +++ b/test/registered/unit/test_legacy_global_ratchet.py @@ -25,7 +25,7 @@ _SRT_ROOT = Path(next(iter(sglang.srt.__path__))) # Baselines counted over python/sglang/srt/**/*.py, including each function's # own def line. Ratchet: decrease-only. _RATCHETS = [ - ("get_global_server_args", r"\bget_global_server_args\s*\(", 346), + ("get_global_server_args", r"\bget_global_server_args\s*\(", 278), ( "set_global_server_args_for_*", r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(", diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 14fc5cd90..83cf63fae 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -78,6 +78,18 @@ class TestModelOverridableWhitelist(CustomTestCase): "ep_size", "moe_dense_tp_size", "attn_cp_size", + "disable_overlap_schedule", + "uses_mamba_radix_cache", + "mamba_radix_cache_strategy", + "speculative_moe_runner_backend", + "speculative_moe_a2a_backend", + "disable_shared_experts_fusion", + "kv_cache_dtype", + "dsa_prefill_backend", + "dsa_decode_backend", + "prefill_attention_backend", + "decode_attention_backend", + "flashinfer_allreduce_fusion_backend", } ), ) @@ -790,6 +802,733 @@ class TestGoldenModelOverrides(_IsolatedPublish): self.assertEqual(_dllm_page_size(_view(dllm_algorithm=None)), {}) self.assertEqual(_dllm_page_size(_view(disable_radix_cache=True)), {}) + def test_overlap_disable_passes(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _dllm_overlap_disable, + _pipeline_parallel_overlap_disable, + _sparse_head_overlap_disable, + ) + + # pipeline parallelism: declares only when pp_size > 1 + self.assertEqual( + _pipeline_parallel_overlap_disable( + ResolvedView(SimpleNamespace(pp_size=1)) + ), + {}, + ) + self.assertEqual( + _pipeline_parallel_overlap_disable( + ResolvedView(SimpleNamespace(pp_size=2)) + ), + {"disable_overlap_schedule": True}, + ) + + # dllm: guarded on the algorithm and the current value + def _view(**kw): + defaults = dict( + dllm_algorithm="LowConfidence", disable_overlap_schedule=False + ) + defaults.update(kw) + return ResolvedView(SimpleNamespace(**defaults)) + + self.assertEqual(_dllm_overlap_disable(_view(dllm_algorithm=None)), {}) + self.assertEqual( + _dllm_overlap_disable(_view(disable_overlap_schedule=True)), {} + ) + self.assertEqual( + _dllm_overlap_disable(_view()), {"disable_overlap_schedule": True} + ) + + # embeddings sparse head: keyed on the env var being set + from sglang.srt.environ import envs + + view = ResolvedView(SimpleNamespace()) + with patch.object( + envs.SGLANG_EMBEDDINGS_SPARSE_HEAD, "is_set", return_value=False + ): + self.assertEqual(_sparse_head_overlap_disable(view), {}) + with patch.object( + envs.SGLANG_EMBEDDINGS_SPARSE_HEAD, "is_set", return_value=True + ): + self.assertEqual( + _sparse_head_overlap_disable(view), {"disable_overlap_schedule": True} + ) + + def test_deepseek_v4_overrides_at_callable_level(self): + from sglang.srt.arg_groups.overrides import _deepseek_v4_overrides + from sglang.srt.server_args import ServerArgs + + hf = SimpleNamespace(architectures=["DeepseekV4ForCausalLM"]) + + def _args(**kw): + defaults = dict( + device="cuda", + swa_full_tokens_ratio=ServerArgs.swa_full_tokens_ratio, + moe_runner_backend="auto", + get_model_config=lambda: SimpleNamespace(nvfp4_moe_meta=None), + ) + defaults.update(kw) + return SimpleNamespace(**defaults) + + self.assertEqual( + _deepseek_v4_overrides(_args(), hf), + { + "attention_backend": "dsv4", + "page_size": 256, + "swa_full_tokens_ratio": 0.1, + }, + ) + # NPU pool geometry + self.assertEqual( + _deepseek_v4_overrides(_args(device="npu"), hf)["page_size"], 128 + ) + # user-set window ratio survives + self.assertNotIn( + "swa_full_tokens_ratio", + _deepseek_v4_overrides(_args(swa_full_tokens_ratio=0.5), hf), + ) + # nvfp4 hybrid checkpoint routes the MoE runner + self.assertEqual( + _deepseek_v4_overrides( + _args( + get_model_config=lambda: SimpleNamespace(nvfp4_moe_meta=object()) + ), + hf, + )["moe_runner_backend"], + "flashinfer_trtllm_routed", + ) + + def test_deepseek_v4_sm120_moe_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _deepseek_v4_sm120_moe, + ) + + def _view(arch="DeepseekV4ForCausalLM", **kw): + hf = SimpleNamespace(architectures=[arch]) + defaults = dict(moe_runner_backend="auto") + defaults.update(kw) + return ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + + with patch.object(overrides_module, "is_sm120_supported", return_value=True): + self.assertEqual( + _deepseek_v4_sm120_moe(_view()), {"moe_runner_backend": "marlin"} + ) + self.assertEqual( + _deepseek_v4_sm120_moe(_view(moe_runner_backend="triton")), {} + ) + self.assertEqual(_deepseek_v4_sm120_moe(_view(arch="LlamaForCausalLM")), {}) + with patch.object(overrides_module, "is_sm120_supported", return_value=False): + self.assertEqual(_deepseek_v4_sm120_moe(_view()), {}) + + def test_nemotron_h_overrides_at_callable_level(self): + from sglang.srt.arg_groups.overrides import _nemotron_h_overrides + + def _hf(quant_algo="NVFP4"): + return SimpleNamespace( + architectures=["NemotronHForCausalLM"], + mlp_hidden_act="relu2", + quantization_config={"quant_algo": quant_algo}, + ) + + def _args(mc_quant, hf, **kw): + mc = SimpleNamespace(quantization=mc_quant, hf_config=hf) + defaults = dict( + quantization=None, + moe_runner_backend="auto", + moe_a2a_backend="none", + attention_backend=None, + get_model_config=lambda: mc, + ) + defaults.update(kw) + return SimpleNamespace(**defaults) + + hf = _hf() + with patch.object(overrides_module, "is_sm100_supported", return_value=True): + # modelopt checkpoint: quant algo resolution + sm100 defaults + self.assertEqual( + _nemotron_h_overrides(_args("modelopt", hf), hf), + { + "quantization": "modelopt_fp4", + "moe_runner_backend": "flashinfer_trtllm", + "attention_backend": "flashinfer", + }, + ) + hf_mixed = _hf("MIXED_PRECISION") + self.assertEqual( + _nemotron_h_overrides(_args("modelopt", hf_mixed), hf_mixed)[ + "quantization" + ], + "modelopt_mixed", + ) + with ( + patch.object(overrides_module, "is_sm100_supported", return_value=False), + patch.object(overrides_module, "is_cuda", return_value=True), + patch.object( + overrides_module, "get_device_capability", return_value=(9, 0) + ), + ): + # SM80-SM90 fp4: marlin + self.assertEqual( + _nemotron_h_overrides(_args("modelopt_fp4", hf), hf), + {"quantization": "modelopt_fp4", "moe_runner_backend": "marlin"}, + ) + # unquantized checkpoint: cutlass fallback, no quant declared + self.assertEqual( + _nemotron_h_overrides(_args(None, hf), hf), + {"moe_runner_backend": "flashinfer_cutlass"}, + ) + # non-modelopt quantized checkpoint: nothing declared + self.assertEqual(_nemotron_h_overrides(_args("fp8", hf), hf), {}) + # user-set moe backend survives + self.assertEqual( + _nemotron_h_overrides(_args(None, hf, moe_runner_backend="triton"), hf), + {}, + ) + + def test_speculative_moe_runner_default_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _speculative_moe_runner_default, + ) + + self.assertEqual( + _speculative_moe_runner_default( + ResolvedView( + SimpleNamespace( + speculative_moe_runner_backend=None, moe_runner_backend="triton" + ) + ) + ), + {"speculative_moe_runner_backend": "triton"}, + ) + # user-set draft backend survives + self.assertEqual( + _speculative_moe_runner_default( + ResolvedView( + SimpleNamespace( + speculative_moe_runner_backend="deep_gemm", + moe_runner_backend="auto", + ) + ) + ), + {}, + ) + + def test_dsa_split_backend_resolution_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _dsa_split_backend_resolution, + ) + + def _view(arch="DeepseekV32ForCausalLM", **kw): + hf = SimpleNamespace(architectures=[arch]) + defaults = dict( + kv_cache_dtype="fp8_e4m3", + dsa_prefill_backend=None, + dsa_decode_backend=None, + enable_hisparse=False, + ) + defaults.update(kw) + return ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + + with ( + patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True), + patch.object(overrides_module, "is_npu", return_value=False), + patch.object(overrides_module, "is_xpu", return_value=False), + patch.object(overrides_module, "is_hip", return_value=False), + patch("torch.cuda.get_device_capability", return_value=(9, 0)), + ): + # Hopper FP8 -> flashmla_kv both + self.assertEqual( + _dsa_split_backend_resolution(_view()), + { + "dsa_prefill_backend": "flashmla_kv", + "dsa_decode_backend": "flashmla_kv", + }, + ) + # Hopper bf16 -> flashmla_sparse / fa3 + self.assertEqual( + _dsa_split_backend_resolution(_view(kv_cache_dtype="bfloat16")), + { + "dsa_prefill_backend": "flashmla_sparse", + "dsa_decode_backend": "fa3", + }, + ) + # user-set prefill survives; only decode defaulted + self.assertEqual( + _dsa_split_backend_resolution(_view(dsa_prefill_backend="trtllm")), + {"dsa_decode_backend": "flashmla_kv"}, + ) + # hisparse arm takes precedence (CUDA fp8 -> flashmla_kv) + self.assertEqual( + _dsa_split_backend_resolution(_view(enable_hisparse=True)), + { + "dsa_prefill_backend": "flashmla_kv", + "dsa_decode_backend": "flashmla_kv", + }, + ) + # non-family arch declares nothing + self.assertEqual( + _dsa_split_backend_resolution(_view(arch="LlamaForCausalLM")), {} + ) + with ( + patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True), + patch.object(overrides_module, "is_npu", return_value=False), + patch.object(overrides_module, "is_xpu", return_value=False), + patch.object(overrides_module, "is_hip", return_value=True), + patch("torch.cuda.get_device_capability", return_value=(9, 4)), + ): + # ROCm with both unset -> tilelang + self.assertEqual( + _dsa_split_backend_resolution(_view(kv_cache_dtype="bfloat16")), + { + "dsa_prefill_backend": "tilelang", + "dsa_decode_backend": "tilelang", + }, + ) + + def test_flashinfer_allreduce_fusion_passes(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _deterministic_allreduce_fusion_disable, + _enforce_disable_allreduce_fusion, + _flashinfer_allreduce_fusion_auto_enable, + ) + + def _view(arch="Qwen3MoeForCausalLM", **kw): + hf = SimpleNamespace(architectures=[arch]) + defaults = dict( + flashinfer_allreduce_fusion_backend=None, + tp_size=2, + enable_dp_attention=False, + nnodes=1, + moe_a2a_backend="none", + enforce_disable_flashinfer_allreduce_fusion=False, + enable_deterministic_inference=False, + ) + defaults.update(kw) + return ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + + with ( + patch.object(overrides_module, "is_sm90_supported", return_value=True), + patch.object(overrides_module, "is_sm100_supported", return_value=False), + ): + self.assertEqual( + _flashinfer_allreduce_fusion_auto_enable(_view()), + {"flashinfer_allreduce_fusion_backend": "auto"}, + ) + # guards: unsupported arch / tp==1 / dp attention / a2a backend + self.assertEqual( + _flashinfer_allreduce_fusion_auto_enable( + _view(arch="LlamaForCausalLM") + ), + {}, + ) + self.assertEqual( + _flashinfer_allreduce_fusion_auto_enable(_view(tp_size=1)), {} + ) + self.assertEqual( + _flashinfer_allreduce_fusion_auto_enable( + _view(enable_dp_attention=True) + ), + {}, + ) + self.assertEqual( + _flashinfer_allreduce_fusion_auto_enable( + _view(moe_a2a_backend="deepep") + ), + {}, + ) + # SM90 multi-node: blocked (nnodes>1 needs SM100) + self.assertEqual( + _flashinfer_allreduce_fusion_auto_enable(_view(nnodes=2)), {} + ) + # user-set backend survives + self.assertEqual( + _flashinfer_allreduce_fusion_auto_enable( + _view(flashinfer_allreduce_fusion_backend="trtllm") + ), + {}, + ) + + # enforce-disable wins over everything + self.assertEqual( + _enforce_disable_allreduce_fusion( + _view( + flashinfer_allreduce_fusion_backend="auto", + enforce_disable_flashinfer_allreduce_fusion=True, + ) + ), + {"flashinfer_allreduce_fusion_backend": None}, + ) + self.assertEqual(_enforce_disable_allreduce_fusion(_view()), {}) + + # deterministic inference disables an enabled fusion + self.assertEqual( + _deterministic_allreduce_fusion_disable( + _view( + flashinfer_allreduce_fusion_backend="auto", + enable_deterministic_inference=True, + ) + ), + {"flashinfer_allreduce_fusion_backend": None}, + ) + self.assertEqual( + _deterministic_allreduce_fusion_disable( + _view(enable_deterministic_inference=True) + ), + {}, + ) + + def test_cutedsl_prefill_backend_fill_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _cutedsl_prefill_backend_fill, + ) + + def _view(**kw): + defaults = dict( + attention_backend=None, + decode_attention_backend="cutedsl_mla", + prefill_attention_backend=None, + kv_cache_dtype="auto", + ) + defaults.update(kw) + return ResolvedView(SimpleNamespace(**defaults)) + + with patch.object(overrides_module, "is_sm100_supported", return_value=True): + # decode-only cutedsl: prefill defaults to trtllm_mla + self.assertEqual( + _cutedsl_prefill_backend_fill(_view()), + {"prefill_attention_backend": "trtllm_mla"}, + ) + # user-set prefill survives + self.assertEqual( + _cutedsl_prefill_backend_fill(_view(prefill_attention_backend="fa3")), + {}, + ) + # cutedsl on the prefill side is rejected + with self.assertRaises(AssertionError): + _cutedsl_prefill_backend_fill( + _view(prefill_attention_backend="cutedsl_mla") + ) + # unsupported kv dtype rejected + with self.assertRaises(ValueError): + _cutedsl_prefill_backend_fill(_view(kv_cache_dtype="fp8_e5m2")) + # not a cutedsl config: nothing declared + self.assertEqual( + _cutedsl_prefill_backend_fill(_view(decode_attention_backend=None)), + {}, + ) + with patch.object(overrides_module, "is_sm100_supported", return_value=False): + with self.assertRaises(ValueError): + _cutedsl_prefill_backend_fill(_view()) + + def test_moss_vl_overrides_at_callable_level(self): + from sglang.srt.arg_groups.overrides import _moss_vl_overrides + + def _args(**kw): + defaults = dict( + attention_backend=None, + prefill_attention_backend=None, + decode_attention_backend=None, + ) + defaults.update(kw) + ns = SimpleNamespace(**defaults) + ns.is_attention_backend_not_set = lambda: ( + ns.attention_backend is None + and ns.prefill_attention_backend is None + and ns.decode_attention_backend is None + ) + ns.get_attention_backends = lambda: ( + ns.prefill_attention_backend or ns.attention_backend, + ns.decode_attention_backend or ns.attention_backend, + ) + return ns + + # nothing set: prefill defaults to flashinfer + self.assertEqual( + _moss_vl_overrides(_args(), None), + {"prefill_attention_backend": "flashinfer"}, + ) + # compatible user choice passes with no declaration + self.assertEqual( + _moss_vl_overrides(_args(attention_backend="flashinfer"), None), {} + ) + # incompatible user choice rejected + with self.assertRaises(AssertionError): + _moss_vl_overrides(_args(attention_backend="fa3"), None) + + def test_dsa_kv_cache_dtype_default_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _dsa_kv_cache_dtype_default, + ) + + def _view(**kw): + hf = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"]) + defaults = dict( + kv_cache_dtype="auto", + dsa_prefill_backend=None, + dsa_decode_backend=None, + ) + defaults.update(kw) + return ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + + with ( + patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True), + patch.object(overrides_module, "is_npu", return_value=False), + patch.object(overrides_module, "is_xpu", return_value=False), + ): + with patch("torch.cuda.get_device_capability", return_value=(9, 0)): + # Hopper: auto -> bfloat16 + self.assertEqual( + _dsa_kv_cache_dtype_default(_view()), + {"kv_cache_dtype": "bfloat16"}, + ) + # alias normalization + self.assertEqual( + _dsa_kv_cache_dtype_default(_view(kv_cache_dtype="bf16")), + {"kv_cache_dtype": "bfloat16"}, + ) + # explicit value survives (no declaration) + self.assertEqual( + _dsa_kv_cache_dtype_default(_view(kv_cache_dtype="fp8_e4m3")), {} + ) + # unsupported dtype rejected + with self.assertRaises(AssertionError): + _dsa_kv_cache_dtype_default(_view(kv_cache_dtype="fp8_e5m2")) + with patch("torch.cuda.get_device_capability", return_value=(10, 0)): + # Blackwell: auto -> fp8 + self.assertEqual( + _dsa_kv_cache_dtype_default(_view()), + {"kv_cache_dtype": "fp8_e4m3"}, + ) + + def test_deepseek_v4_kv_cache_dtype_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _deepseek_v4_kv_cache_dtype, + ) + + def _view(arch="DeepseekV4ForCausalLM", **kw): + hf = SimpleNamespace(architectures=[arch]) + defaults = dict(kv_cache_dtype="auto", device="cuda") + defaults.update(kw) + return ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + + self.assertEqual( + _deepseek_v4_kv_cache_dtype(_view()), {"kv_cache_dtype": "fp8_e4m3"} + ) + # NPU pins bfloat16 regardless of the auto default + self.assertEqual( + _deepseek_v4_kv_cache_dtype(_view(device="npu")), + {"kv_cache_dtype": "bfloat16"}, + ) + # explicit supported value survives + self.assertEqual( + _deepseek_v4_kv_cache_dtype(_view(kv_cache_dtype="bfloat16")), {} + ) + with self.assertRaises(AssertionError): + _deepseek_v4_kv_cache_dtype(_view(kv_cache_dtype="fp8_e5m2")) + self.assertEqual( + _deepseek_v4_kv_cache_dtype(_view(arch="LlamaForCausalLM")), {} + ) + + def test_deepseek_spec_moe_resolution_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _deepseek_spec_moe_resolution, + ) + from sglang.srt.environ import envs + + def _view(**kw): + hf = SimpleNamespace(architectures=["DeepseekV3ForCausalLM"]) + defaults = dict( + quantization="modelopt_fp4", + speculative_algorithm="EAGLE", + speculative_moe_runner_backend=None, + speculative_moe_a2a_backend=None, + ep_size=8, + ) + defaults.update(kw) + return ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + + with patch.object(overrides_module, "is_hip", return_value=True): + with patch.object( + envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE, "get", return_value=False + ): + self.assertEqual( + _deepseek_spec_moe_resolution(_view()), + { + "speculative_moe_runner_backend": "triton", + "speculative_moe_a2a_backend": "none", + }, + ) + # guards: quantization / algorithm / both fields user-set + self.assertEqual( + _deepseek_spec_moe_resolution(_view(quantization="fp8")), {} + ) + self.assertEqual( + _deepseek_spec_moe_resolution(_view(speculative_algorithm=None)), + {}, + ) + self.assertEqual( + _deepseek_spec_moe_resolution( + _view( + speculative_moe_runner_backend="triton", + speculative_moe_a2a_backend="none", + ) + ), + {}, + ) + with patch.object( + envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE, "get", return_value=True + ): + self.assertEqual( + _deepseek_spec_moe_resolution(_view()), + { + "speculative_moe_runner_backend": "deep_gemm", + "speculative_moe_a2a_backend": "deepep", + }, + ) + with self.assertRaises(ValueError): + _deepseek_spec_moe_resolution(_view(ep_size=1)) + # the arm is HIP-only + with patch.object(overrides_module, "is_hip", return_value=False): + self.assertEqual(_deepseek_spec_moe_resolution(_view()), {}) + + def test_mamba_radix_cache_resolution_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _mamba_radix_cache_resolution, + supports_mamba_cache_extra_buffer, + ) + + def _view(arch, layer_types=None, **kw): + hf = SimpleNamespace(architectures=[arch]) + if layer_types is not None: + hf.layer_types = layer_types + defaults = dict( + disable_radix_cache=False, + mamba_radix_cache_strategy="auto", + disable_overlap_schedule=False, + page_size=None, + linear_attn_backend="triton", + ) + defaults.update(kw) + return ResolvedView( + SimpleNamespace( + get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults + ) + ) + + # arch guard: non-mamba arch declares nothing + self.assertEqual(_mamba_radix_cache_resolution(_view("LlamaForCausalLM")), {}) + # radix cache disabled: nothing to resolve + self.assertEqual( + _mamba_radix_cache_resolution( + _view("Qwen3NextForCausalLM", disable_radix_cache=True) + ), + {}, + ) + # auto + overlap wanted + extra-buffer support -> extra_buffer + self.assertEqual( + _mamba_radix_cache_resolution(_view("Qwen3NextForCausalLM")), + { + "uses_mamba_radix_cache": True, + "mamba_radix_cache_strategy": "extra_buffer", + }, + ) + # auto + no extra-buffer support (Lfm2) -> no_buffer + overlap disable + self.assertEqual( + _mamba_radix_cache_resolution(_view("Lfm2ForCausalLM")), + { + "uses_mamba_radix_cache": True, + "mamba_radix_cache_strategy": "no_buffer", + "disable_overlap_schedule": True, + }, + ) + # neither overlap nor paging wanted -> no_buffer even when supported + declared = _mamba_radix_cache_resolution( + _view("Qwen3NextForCausalLM", disable_overlap_schedule=True, page_size=1) + ) + self.assertEqual(declared["mamba_radix_cache_strategy"], "no_buffer") + self.assertIs(declared["disable_overlap_schedule"], True) + # paging alone wants the extra buffer + self.assertEqual( + _mamba_radix_cache_resolution( + _view( + "Qwen3NextForCausalLM", disable_overlap_schedule=True, page_size=64 + ) + )["mamba_radix_cache_strategy"], + "extra_buffer", + ) + # user-set strategy: only the routing marker is declared + self.assertEqual( + _mamba_radix_cache_resolution( + _view( + "Qwen3NextForCausalLM", + mamba_radix_cache_strategy="extra_buffer_lazy", + ) + ), + {"uses_mamba_radix_cache": True}, + ) + # NemotronH routes through the pass (covered by the guard union, + # not the branch chain — its hook invokes the handler) + self.assertEqual( + _mamba_radix_cache_resolution(_view("NemotronHForCausalLM")), + { + "uses_mamba_radix_cache": True, + "mamba_radix_cache_strategy": "extra_buffer", + }, + ) + # GraniteMoeHybrid is guarded on mamba layer types + self.assertEqual( + _mamba_radix_cache_resolution( + _view("GraniteMoeHybridForCausalLM", layer_types=["attention"]) + ), + {}, + ) + self.assertEqual( + _mamba_radix_cache_resolution( + _view("GraniteMoeHybridForCausalLM", layer_types=["mamba", "attention"]) + )["mamba_radix_cache_strategy"], + "extra_buffer", + ) + # extra-buffer support requires the triton linear-attn backend + self.assertFalse( + supports_mamba_cache_extra_buffer( + SimpleNamespace(linear_attn_backend="fla"), "Qwen3NextForCausalLM" + ) + ) + def test_page_size_leaf_materializes_end_state(self): sa = self._construct("LlamaForCausalLM", "llama") declared = {f for _s, d in sa._resolved_overrides for f in d} diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 47b556bd5..358a19123 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -9,6 +9,7 @@ import unittest from unittest.mock import patch import sglang.srt.server_args as server_args_module +from sglang.srt.arg_groups.arg_utils import A, Arg from sglang.srt.runtime_context import ( Flags, ParallelContext, @@ -301,5 +302,146 @@ class TestFlagsTier(_IsolatedServerArgs): reset_context() # never leave the singleton frozen for other tests +@dataclasses.dataclass +class _FakeResolvedArgs: + """Publishable fixture with a resolvable whitelist (real flat leaves).""" + + page_size: A[int | None, Arg(help="p", resolvable=True)] = None + sampling_backend: A[str | None, Arg(help="s", resolvable=True)] = None + _resolved_overrides: list = dataclasses.field(default_factory=list) + + +class TestRuntimeResolutionStages(_IsolatedServerArgs): + """Runtime stages: post-publish declarations re-resolve the flags tier + atomically; freeze_flags() ends the resolution lifecycle.""" + + def _publish(self, **kw): + args = _FakeResolvedArgs(**kw) + get_context().set_server_args(args) + return args + + def test_record_before_publish_raises(self): + reset_context() + with self.assertRaises(ValueError): + get_context().record_runtime_overrides([("stage", {"page_size": 64})]) + + def test_record_updates_leaves_and_accumulates_stages(self): + args = self._publish(page_size=1, sampling_backend="flashinfer") + self.assertEqual(get_flags().page_size, 1) # publish-time materialize + # dual-apply transition: the call site keeps its imperative write + args.page_size = 64 + get_context().record_runtime_overrides([("stage.runner", {"page_size": 64})]) + self.assertEqual(get_flags().page_size, 64) + args.sampling_backend = "pytorch" + get_context().record_runtime_overrides( + [("stage.load", {"sampling_backend": "pytorch"})] + ) + self.assertEqual(get_flags().sampling_backend, "pytorch") + self.assertEqual(get_flags().page_size, 64) # earlier stage survives + + def test_record_parity_failure_rolls_back(self): + self._publish(page_size=1) + flags_before = get_flags() + with self.assertRaises(AssertionError): + # declared value diverges from the live server_args (no dual-apply) + get_context().record_runtime_overrides([("bad", {"page_size": 64})]) + self.assertIs(get_flags(), flags_before) # previous flags intact + self.assertEqual(get_context()._runtime_overrides, []) # rolled back + + def test_record_whitelist_violation_rolls_back(self): + self._publish() + with self.assertRaises(ValueError): + get_context().record_runtime_overrides([("bad", {"nope": 1})]) + self.assertEqual(get_context()._runtime_overrides, []) + + def test_freeze_ends_the_resolution_lifecycle(self): + args = self._publish(page_size=1) + try: + get_context().freeze_flags() + self.assertTrue(get_flags().frozen) + with self.assertRaises(RuntimeError): + get_context().record_runtime_overrides([("late", {"page_size": 64})]) + with self.assertRaises(RuntimeError): + get_context().set_server_args(args) + finally: + reset_context() + + def test_declare_load_time_override_dual_applies_and_records(self): + from sglang.srt.arg_groups.overrides import declare_load_time_override + + args = self._publish(page_size=1) + declare_load_time_override("model.load_time", {"page_size": 64}) + self.assertEqual(args.page_size, 64) # dual-applied onto server_args + self.assertEqual(get_flags().page_size, 64) # resolved into the leaf + self.assertEqual( + get_context()._runtime_overrides, + [("model.load_time", {"page_size": 64})], + ) + + def test_failed_republish_keeps_previous_lifecycle(self): + args = self._publish(page_size=1) + args.page_size = 64 + get_context().record_runtime_overrides([("stage", {"page_size": 64})]) + flags_before = get_flags() + bad = _FakeResolvedArgs(page_size=1) + bad._resolved_overrides = [("bad", {"nope": 1})] # gate rejects + with self.assertRaises(ValueError): + get_context().set_server_args(bad) + # previous publish fully intact: slot, flags, and the recorded stages + self.assertIs(get_context()._server_args, args) + self.assertIs(get_flags(), flags_before) + self.assertEqual( + get_context()._runtime_overrides, [("stage", {"page_size": 64})] + ) + + def test_capture_tier_seeded_at_publish_and_survives_stages(self): + # seeded from the published config + args = self._publish(page_size=1) + args.enable_torch_compile = True + get_context().set_server_args(args) # re-publish picks up the value + self.assertTrue(get_flags().capture.enable_torch_compile) + # capture-time write (B4) targets the capture leaf + get_flags().capture.enable_torch_compile = False + self.assertFalse(get_flags().capture.enable_torch_compile) + # a runtime-stage re-resolve must not clobber the capture write + args.page_size = 64 + get_context().record_runtime_overrides([("stage", {"page_size": 64})]) + self.assertFalse(get_flags().capture.enable_torch_compile) + # capture stays writable after freeze + try: + get_context().freeze_flags() + get_flags().capture.enable_torch_compile = True + self.assertTrue(get_flags().capture.enable_torch_compile) + finally: + reset_context() + + def test_bare_dataclass_publish_skips_materialization(self): + # object.__new__(ServerArgs) fixtures (no __init__, no field values) + # must publish without touching the flags tier — dataclass defaults + # live on the class, so materializing from them would clobber + # previously resolved flags with defaults. + from sglang.srt.server_args import ServerArgs + + self._publish(page_size=64) + self.assertEqual(get_flags().page_size, 64) + bare = object.__new__(ServerArgs) + get_context().set_server_args(bare) + self.assertIs(get_server_args(), bare) + self.assertEqual(get_flags().page_size, 64) # not clobbered + + def test_capture_tier_defaults_for_sentinel_publish(self): + get_context().set_server_args(object()) + self.assertFalse(get_flags().capture.enable_torch_compile) + + def test_republish_clears_runtime_overrides(self): + args = self._publish(page_size=1) + args.page_size = 64 + get_context().record_runtime_overrides([("stage", {"page_size": 64})]) + self.assertEqual(get_flags().page_size, 64) + self._publish(page_size=1) # fresh lifecycle + self.assertEqual(get_flags().page_size, 1) + self.assertEqual(get_context()._runtime_overrides, []) + + if __name__ == "__main__": unittest.main()