[refactor] Config resolution pipeline: full-stack review (10-PR series, review only) (#30137)

This commit is contained in:
Cheng Wan
2026-07-05 00:00:07 -07:00
committed by GitHub
parent ce733f106b
commit 8fb99bbaf8
74 changed files with 2035 additions and 628 deletions
@@ -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."""
+2 -24
View File
@@ -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(
@@ -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"
)
+588 -1
View File
@@ -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:
@@ -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()
@@ -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,
@@ -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
)
+2 -2
View File
@@ -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
@@ -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:
+3 -2
View File
@@ -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
+1 -1
View File
@@ -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):
+8
View File
@@ -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):
+3 -2
View File
@@ -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:
@@ -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)
@@ -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()
@@ -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
@@ -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().
@@ -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
+2 -3
View File
@@ -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)
+2 -3
View File
@@ -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)
+2 -2
View File
@@ -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)
@@ -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)
@@ -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":
+3 -3
View File
@@ -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)
+11 -6
View File
@@ -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,
+9 -4
View File
@@ -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.",
@@ -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)
+2 -3
View File
@@ -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)
+2 -2
View File
@@ -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
+2 -3
View File
@@ -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)
+2 -3
View File
@@ -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
+10 -7
View File
@@ -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,
+10 -7
View File
@@ -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,
@@ -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()
+3 -3
View File
@@ -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()
+9 -4
View File
@@ -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.",
+3 -4
View File
@@ -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()
+3 -3
View File
@@ -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
+2 -2
View File
@@ -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()
+2 -2
View File
@@ -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)
+2 -3
View File
@@ -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)
+2 -3
View File
@@ -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
+2 -2
View File
@@ -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()
+2 -3
View File
@@ -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)
+2 -2
View File
@@ -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:
+2 -3
View File
@@ -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)
+3 -3
View File
@@ -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
+2 -2
View File
@@ -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:
+9 -6
View File
@@ -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)."
+2 -2
View File
@@ -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))
+2 -2
View File
@@ -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
+2 -3
View File
@@ -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)
+2 -3
View File
@@ -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
+3 -3
View File
@@ -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))
+2 -2
View File
@@ -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:
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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:
+2 -2
View File
@@ -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:
+2 -2
View File
@@ -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:
+107 -4
View File
@@ -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 = []
+137 -279
View File
@@ -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
@@ -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)
@@ -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)
@@ -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)
@@ -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)