[refactor] Config resolution pipeline: full-stack review (10-PR series, review only) (#30137)
This commit is contained in:
@@ -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."""
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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.",
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)."
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user