[Config] Round 5.2: the per-model declarations get their own modules (#37087)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-30 02:24:35 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 7e751153eb
commit e51a3ae65e
44 changed files with 2898 additions and 1730 deletions
@@ -0,0 +1,376 @@
"""What the per-model override declarations are written against.
The declarations themselves live one directory down, in
``arg_groups/model_overrides/``: one module per model family, mirroring the
``models/`` naming. This module is what they all import -- the registry they
register into, the read-only views they are handed, and the few accessors that
answer questions about the model. It deliberately depends on nothing in
``overrides.py``, so a family module never has to import its way back up.
"""
import logging
from typing import Any, Callable, Dict, List, Optional, Tuple
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import is_mps, is_no_spec_infer_or_topk_one
logger = logging.getLogger(__name__)
# Constant per-architecture overrides (populated by the migration sweeps).
MODEL_OVERRIDES: Dict[str, Dict[str, Any]] = {
# These models run in bfloat16 regardless of the requested dtype
# (faithful port of the legacy unconditional arch branch).
"MistralLarge3ForCausalLM": {"dtype": "bfloat16"},
"PixtralForConditionalGeneration": {"dtype": "bfloat16"},
}
# Derived per-architecture override providers, in registration order.
_MODEL_OVERRIDE_FNS: Dict[str, List[Callable[..., dict]]] = {}
# Predicate-keyed providers, in registration order — for legacy branches
# matched by substring/predicate on the architecture string rather than an
# exact name (e.g. '"Step3p5ForCausalLM" in model_arch').
_PREDICATE_OVERRIDE_FNS: List[Tuple[Callable[[str], bool], Callable[..., dict]]] = []
def register_model_override(architecture: str):
"""Register a derived-override provider for ``architecture``.
The decorated callable receives ``(server_args, hf_config)``, must not
mutate either, and returns a ``{field: resolved_value}`` dict (possibly
empty when nothing applies). Providers needing derived model data beyond
the HF config go through ``model_config_of(server_args)`` (cached,
read-only) — never anything mutating.
"""
def decorator(fn: Callable[..., dict]) -> Callable[..., dict]:
_MODEL_OVERRIDE_FNS.setdefault(architecture, []).append(fn)
return fn
return decorator
def register_model_override_predicate(predicate: Callable[[str], bool]):
"""Register a derived-override provider keyed by an architecture
predicate. Same callable contract as ``register_model_override``."""
def decorator(fn: Callable[..., dict]) -> Callable[..., dict]:
_PREDICATE_OVERRIDE_FNS.append((predicate, fn))
return fn
return decorator
def _invoke_provider(
fn: Callable[..., dict], server_args: Any, hf_config: Any
) -> Dict[str, Any]:
declared = fn(server_args, hf_config)
if not isinstance(declared, dict):
raise TypeError(
f"model override provider {fn.__qualname__} must return a dict, "
f"got {type(declared).__name__}"
)
return declared
class ResolvedView:
"""Read-only view of the resolving configuration handed to post-process
passes: the accumulated declarations overlaid on the pristine
``server_args`` (residual imperative writes of non-resolved fields show
through the fallthrough) — exactly the state the legacy handler at the
same slot observed. Writes are rejected: passes return declarations.
"""
__slots__ = ("_server_args", "_overlay")
def __init__(self, server_args: Any, overlay: Optional[Dict[str, Any]] = None):
object.__setattr__(self, "_server_args", server_args)
object.__setattr__(self, "_overlay", overlay or {})
def __getattr__(self, name: str) -> Any:
overlay = object.__getattribute__(self, "_overlay")
if name in overlay:
return overlay[name]
return getattr(object.__getattribute__(self, "_server_args"), name)
def __setattr__(self, name: str, value: Any) -> None:
raise AttributeError(
"ResolvedView is read-only; post-process passes return declarations"
)
class ResolvingConfig:
"""Live read view of the resolution result: the declaration stash over the
record's fields, looked up per read.
``ResolvedView`` snapshots the overlay when it is built, which is what a
post-process pass wants -- it reads the state at its slot. A resolver that
reads *after* declaring, or after calling something that declares, needs the
current answer instead, so this one walks the stash on every read. It falls
through to the field, which is where the raw input lives.
"""
__slots__ = ("_server_args",)
def __init__(self, server_args: Any):
object.__setattr__(self, "_server_args", server_args)
def __getattr__(self, name: str) -> Any:
server_args = object.__getattribute__(self, "_server_args")
for _source, declared in reversed(
getattr(server_args, "_resolved_overrides", None) or ()
):
if name in declared:
return declared[name]
return getattr(server_args, name)
def __setattr__(self, name: str, value: Any) -> None:
raise AttributeError(
"ResolvingConfig is read-only; resolution writes through declarations"
)
def resolving_view(server_args: Any) -> ResolvingConfig:
"""A live read view of what resolution has decided so far."""
return ResolvingConfig(server_args)
def _declaration_overlay(server_args: Any) -> Dict[str, Any]:
"""What the declarations say so far, last writer wins.
Nothing writes the fields, so a mid-resolution reader needs this to see a
decision at all; the fields keep what the caller supplied."""
overlay: Dict[str, Any] = {}
for _source, declared in getattr(server_args, "_resolved_overrides", None) or ():
overlay.update(declared)
return overlay
def resolved_view(server_args: Any) -> ResolvedView:
"""Read-only view of the resolving configuration: the declarations
overlaid on the fields, snapshotted per call.
For mid-resolution code that is not a pass (``__post_init__`` handlers and
hooks) that must answer with what resolution decided -- a declaration-only resolver (a model-specific
override, a registry entry) never writes the field, so a field read there
answers with the raw input."""
return ResolvedView(server_args, overlay=_declaration_overlay(server_args))
def attention_backends_of(cfg: Any) -> tuple:
"""(prefill, decode) attention backends of a config-shaped object (a
ResolvedView mid-resolution, or pristine server_args at dispatch time):
split fields fall back to the base backend."""
prefill = (
cfg.prefill_attention_backend
if cfg.prefill_attention_backend
else cfg.attention_backend
)
decode = (
cfg.decode_attention_backend
if cfg.decode_attention_backend
else cfg.attention_backend
)
return prefill, decode
def _register_for(*architectures: str):
"""Register one provider for several architectures (family lists)."""
def decorator(fn: Callable[..., dict]) -> Callable[..., dict]:
for architecture in architectures:
register_model_override(architecture)(fn)
return fn
return decorator
def record_of(view: Any) -> Any:
"""The record a view reads through.
For the few helpers a view cannot serve: `get_default_attn_backend` reads
through *both* overlays, so it needs the record the two views are built
from rather than either one of them.
"""
return object.__getattribute__(view, "_server_args")
def is_attention_backend_not_set(cfg: Any):
"""None of the three attention backends has been decided yet.
Takes the view rather than the record: every read is a view read, and the
callers that hold a view (the override providers) would otherwise have to
reach back through it for a record.
"""
return (
cfg.attention_backend is None
and cfg.prefill_attention_backend is None
and cfg.decode_attention_backend is None
)
def use_mla_backend(server_args: Any):
from sglang.srt.configs.model_config import AttentionArch
model_config = model_config_of(server_args)
return model_config.attention_arch == AttentionArch.MLA
def model_config_of(server_args: Any):
"""The model configuration this record describes, built once and memoised.
Takes a view as readily as the record: a view is a read overlay of one
record, the memo has to live on that record either way, and the callers
that hold a view would otherwise all have to unwrap it themselves.
"""
if isinstance(server_args, (ResolvedView, ResolvingConfig)):
server_args = record_of(server_args)
# Lazy init to avoid circular import
cfg = resolving_view(server_args)
from sglang.srt.configs.model_config import ModelConfig
memo = getattr(server_args, "_model_config", None)
if memo is not None:
# The key is the path this record carried when the cache was
# filled. The GGUF and ModelScope handlers declare a different
# `model_path`, and a configuration built before them describes
# another checkpoint. `ModelConfig` re-points its own `model_path`
# at the local pull directory when the weights sit behind an
# object-store URI, so its field is not the key. A configuration a
# fixture supplied carries no key and is handed back as it is.
built_from = getattr(server_args, "_model_config_built_from", None)
if built_from is None or built_from == cfg.model_path:
return memo
model_config = ModelConfig.from_server_args(server_args)
server_args._model_config = model_config
server_args._model_config_built_from = cfg.model_path
if model_config.is_hybrid_swa:
logger.info(
"Hybrid SWA model detected. architectures=%s",
model_config.hf_config.architectures,
)
return model_config
def mamba_extra_buffer_of(cfg: Any) -> bool:
"""Mid-resolution equivalent of runtime_context.mamba_extra_buffer_enabled:
reads the (possibly overlaid) strategy from a config-shaped object.
This is the one definition of the predicate: ``ServerArgs`` delegates its
member to it, and the runtime_context accessor is its post-publish sibling
(which cannot reuse it, because the two leaves land in different bags)."""
return cfg.disable_radix_cache is False and cfg.mamba_radix_cache_strategy in (
"extra_buffer",
"extra_buffer_lazy",
)
def get_default_attn_backend(server_args: Any, use_mla_backend: bool, model_config):
"""
Auto select the fastest attention backend.
1. Models with MHA Architecture (e.g: Llama, QWen)
1.1 We will turn on FA3 on hopper unless user use spec decode with topk > 1 or page_size > 1.
1.2 Use trtllm_mha for SM100/SM103 (Blackwell B200/GB200/B300) excluding spec with topk > 1.
Note: trtllm_mha does not support SM120, which will fall back to flashinfer.
1.3 In other cases, we will use flashinfer if available, otherwise use triton.
2. Models with MLA Architecture and using FA3
2.1 We will use FA3 backend on hopper.
2.2 We will use Flashinfer backend on blackwell.
2.3 Otherwise, we will use triton backend.
"""
cfg = resolving_view(server_args)
# OOT platforms provide their own default attention backend.
if current_platform.is_out_of_tree():
return current_platform.get_default_attention_backend()
# Whisper requires flashinfer for cross-attention CUDA graph support.
if "WhisperForConditionalGeneration" in (
model_config.hf_config.architectures or []
):
return "flashinfer"
if not use_mla_backend:
# MHA architecture
if get_platform().is_hopper_with_cuda_12_3 and is_no_spec_infer_or_topk_one(
resolved_view(server_args)
):
# Note: flashinfer 0.6.1 caused performance regression on Hopper attention kernel
# Before the kernel is fixed, we choose fa3 as the default backend on Hopper MHA
# ref: https://github.com/sgl-project/sglang/issues/17411
return "fa3"
elif (
get_platform().is_sm100
and is_no_spec_infer_or_topk_one(resolved_view(server_args))
and (
cfg.speculative_algorithm is None
or cfg.speculative_eagle_topk is not None
)
):
# trtllm_mha requires equal K/V row widths; fa4 carries
# v_head_dim through.
if model_config.has_asymmetric_kv:
return "fa4"
return "trtllm_mha"
elif get_platform().is_hip:
return "aiter"
elif is_mps():
return "torch_native"
else:
# FlashInfer does not support attention sinks.
if get_platform().has_flashinfer and not model_config.has_attention_sinks:
return "flashinfer"
return "triton"
else:
# MLA architecture
if get_platform().is_hopper_with_cuda_12_3:
return "fa3"
elif get_platform().is_sm100:
return "flashinfer"
elif get_platform().is_hip:
head_num = model_config.get_num_kv_heads(cfg.tp_size)
# TODO current aiter only support head number 16 or 128 head number
if head_num == 128 or head_num == 16:
return "aiter"
else:
return "triton"
elif is_mps():
return "torch_native"
else:
return "triton"
def _dspark_verify_on_decode_backend(
backend: Optional[str], q_len: int, kv_cache_dtype: Optional[str]
) -> bool:
"""Whether the MLA decode backend can serve a q_len-wide target verify."""
if backend == "trtllm_mla":
return True
if backend == "tokenspeed_mla":
return kv_cache_dtype == "fp8_e4m3" and q_len <= 8
if backend == "cutedsl_mla":
# cute-dsl monolithic MLA decode folds the verify tokens into the head
# dim (fold_sq), so it serves any DSPARK verify width. Needs flashinfer
# >= 0.6.15 (older builds reject q_len >= 5).
return True
return False
def _is_mxfp4_pack_quantized(hf_config: Any) -> bool:
qc = getattr(
getattr(hf_config, "text_config", hf_config), "quantization_config", None
)
if not isinstance(qc, dict):
return False
groups = qc.get("config_groups") or {}
return any(
"mxfp4" in str(g.get("format", ""))
for g in groups.values()
if isinstance(g, dict)
)
@@ -0,0 +1,37 @@
"""Per-model config-time override declarations, one module per family,
mirroring the ``models/`` naming.
Importing this package is what registers them. An architecture may be claimed
by more than one module here -- one supplies its attention shape, another its
MoE runner -- but two of them must never declare the *same* field for it:
nobody would own that value, and which module supplied it would come down to
the order of the imports below. ``test_model_override_split.py`` forbids the
overlap, which is why this list needs no particular order.
"""
from sglang.srt.arg_groups.model_overrides import deepseek_v2 # noqa: F401
from sglang.srt.arg_groups.model_overrides import deepseek_v4 # noqa: F401
from sglang.srt.arg_groups.model_overrides import exaone # noqa: F401
from sglang.srt.arg_groups.model_overrides import falcon_h1 # noqa: F401
from sglang.srt.arg_groups.model_overrides import gemma2_gemma3 # noqa: F401
from sglang.srt.arg_groups.model_overrides import gemma4 # noqa: F401
from sglang.srt.arg_groups.model_overrides import glm4_moe # noqa: F401
from sglang.srt.arg_groups.model_overrides import gpt_oss # noqa: F401
from sglang.srt.arg_groups.model_overrides import granitemoehybrid # noqa: F401
from sglang.srt.arg_groups.model_overrides import inkling # noqa: F401
from sglang.srt.arg_groups.model_overrides import interns2_mobius # noqa: F401
from sglang.srt.arg_groups.model_overrides import kimi_k3 # noqa: F401
from sglang.srt.arg_groups.model_overrides import lfm2 # noqa: F401
from sglang.srt.arg_groups.model_overrides import llama4 # noqa: F401
from sglang.srt.arg_groups.model_overrides import mimo_v2 # noqa: F401
from sglang.srt.arg_groups.model_overrides import minicpm # noqa: F401
from sglang.srt.arg_groups.model_overrides import minicpmv # noqa: F401
from sglang.srt.arg_groups.model_overrides import minimax_m2 # noqa: F401
from sglang.srt.arg_groups.model_overrides import minimax_m3 # noqa: F401
from sglang.srt.arg_groups.model_overrides import moss_vl # noqa: F401
from sglang.srt.arg_groups.model_overrides import muse_glimmer # noqa: F401
from sglang.srt.arg_groups.model_overrides import nemotron_h # noqa: F401
from sglang.srt.arg_groups.model_overrides import olmo2 # noqa: F401
from sglang.srt.arg_groups.model_overrides import qwen3_5 # noqa: F401
from sglang.srt.arg_groups.model_overrides import qwen3_moe # noqa: F401
from sglang.srt.arg_groups.model_overrides import qwen3_vl # noqa: F401
@@ -0,0 +1,145 @@
"""Config-time override declarations for deepseek_v2.
Architectures: DeepseekV32ForCausalLM, DeepseekV3ForCausalLM, Dots3NoteForCausalLM, GlmMoeDsaForCausalLM, KimiK25ForConditionalGeneration, LongcatFlashForCausalLM, LongcatFlashForCausalLMNextN, MistralLarge3ForCausalLM, PixtralForConditionalGeneration.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
is_attention_backend_not_set,
resolving_view,
use_mla_backend,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for(
"DeepseekV3ForCausalLM",
"DeepseekV32ForCausalLM",
"KimiK25ForConditionalGeneration",
"MistralLarge3ForCausalLM",
"PixtralForConditionalGeneration",
"GlmMoeDsaForCausalLM",
"LongcatFlashForCausalLM",
"LongcatFlashForCausalLMNextN",
"Dots3NoteForCausalLM",
)
def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
"""Order-safe declarations of the DeepSeek/DSA branch. The CP parallel
writes (enable_dp_attention/ep_size/moe_a2a_backend have post-monolith
writers), the kv-cache/split-backend defaults, the quant/moe block (read
before it by _set_default_dsa_kv_cache_dtype) and the env writes stay in
the branch."""
cfg = resolving_view(server_args)
from sglang.srt.configs.model_config import is_deepseek_dsa
overrides: Dict[str, Any] = {}
if is_deepseek_dsa(hf_config): # DeepSeek 3.2/GLM 5
# Set attention backend for DeepSeek
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = "dsa"
logger.info("Use dsa attention backend for DeepSeek with DSA.")
if not get_platform().is_npu and not get_platform().is_xpu: # CUDA or ROCm GPU
if cfg.enable_prefill_cp:
logger.warning(
"Context parallel feature is still under experiment. It has only been verified on Hopper platform."
)
overrides["enable_dp_attention"] = True
overrides["moe_dense_tp_size"] = 1
if cfg.cp_strategy == "zigzag":
overrides["moe_a2a_backend"] = "deepep"
overrides["ep_size"] = cfg.tp_size
logger.warning(
"zigzag DSA CP requires moe_dense_tp_size=1, "
"moe_a2a_backend=deepep, ep_size=tp_size, batch_size=1."
)
else:
assert (
cfg.dp_size == 1
), "interleave DSA CP does not support DP attention."
assert (
cfg.tp_size <= 8
), "Context parallel only supports single machine (tp_size <= 8). Cross-machine CP has precision issues."
# Note(kpham-sgl): Keep attn_tp_size == 1 under DSA CP.
# DSACPLayerCommunicator does not all-reduce attention-TP
# partial o_proj outputs before replicated dense FFNs.
attn_cp_size = cfg.tp_size // cfg.dp_size
overrides["attn_cp_size"] = attn_cp_size
logger.warning(
"Enabled DSA context parallel: "
f"strategy={cfg.cp_strategy}, dp_size={cfg.dp_size}, "
f"moe_dense_tp_size={overrides['moe_dense_tp_size']}, "
f"ep_size={overrides.get('ep_size', cfg.ep_size)}, tp_size={cfg.tp_size}, "
f"attn_cp_size={attn_cp_size}, "
f"kv_cache_dtype={cfg.kv_cache_dtype}, "
f"moe_a2a_backend={overrides.get('moe_a2a_backend', cfg.moe_a2a_backend)}, "
f"cuda_graph_config[prefill].backend=disabled"
)
# Deferred import to avoid a circular import at module-load
# time (dsa.utils imports the runtime-context accessors).
from sglang.srt.layers.attention.dsa.utils import (
aiter_can_use_preshuffle_paged_mqa,
)
if get_platform().is_hip and not aiter_can_use_preshuffle_paged_mqa():
# Legacy ROCm DSA path: aiter's gluon paged-MQA kernel is
# unavailable (Triton<3.5 and AITER_ENABLE_AOT_GLUON_PA_MQA_LOGITS
# not set, or SGLANG_DSA_HIP_DISABLE_PRESHUFFLE=1 / SGLANG_USE_AITER=0).
overrides["page_size"] = 1
logger.warning(
"Setting page size to 1 for DeepSeek DSA on ROCm "
"(aiter preshuffle paged-MQA path unavailable: "
"needs Triton>=3.5.0 or AITER_ENABLE_AOT_GLUON_PA_MQA_LOGITS=1)."
)
else:
overrides["page_size"] = 64
logger.warning("Setting page size to 64 for DeepSeek DSA.")
else:
# DeepSeek V3/R1/V3.1
if get_platform().is_sm100:
if (
cfg.attention_backend is None
and cfg.prefill_attention_backend is None
and cfg.decode_attention_backend is None
):
overrides["attention_backend"] = "trtllm_mla"
logger.info(
"Use trtllm_mla as attention backend on sm100 for DeepseekV3ForCausalLM"
)
# MLA prefill CP auto-config. Mirrors the NSA CP block above
# (minus the in-seq/round-robin mode split, which MLA CP does not support)
if cfg.enable_prefill_cp and use_mla_backend(server_args):
logger.warning(
"MLA prefill context parallel is still experimental. "
"Verified on Hopper with the fa3 backend."
)
overrides["enable_dp_attention"] = True
# TODO(kpham-sgl) Supports moe_dense_tp_size != 1.
overrides["moe_dense_tp_size"] = 1
overrides["moe_a2a_backend"] = "deepep"
overrides["ep_size"] = cfg.tp_size
logger.warning(
"For MLA CP, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, batch_size == 1"
)
# FIXME(kpham-sgl): Keep attn_tp_size == 1 under MLA CP.
# DSACPLayerCommunicator does not all-reduce attention-TP
# partial o_proj outputs before replicated dense FFNs.
attn_cp_size = cfg.tp_size // cfg.dp_size
overrides["attn_cp_size"] = attn_cp_size
logger.warning(
f"Enable Context Parallel opt for MLA, "
f"Setting dp_size == {cfg.dp_size} and "
f"attn_cp_size == {attn_cp_size}, "
f"moe_dense_tp_size == {overrides['moe_dense_tp_size']}, "
f"ep_size == {overrides['ep_size']}, "
f"tp_size == {cfg.tp_size}, "
f"moe_a2a_backend {overrides['moe_a2a_backend']}, "
f"cuda_graph_config[prefill].backend=disabled"
)
return overrides
@@ -0,0 +1,76 @@
"""Config-time override declarations for deepseek_v4.
Architectures: DeepseekV4ForCausalLM.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
model_config_of,
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_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."""
cfg = resolving_view(server_args)
from sglang.srt.server_args import ServerArgs
model_arch = hf_config.architectures[0]
overrides: Dict[str, Any] = {"attention_backend": "dsv4"}
page_size = 256
if cfg.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 cfg.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}.")
if cfg.moe_runner_backend == "auto":
model_config = model_config_of(server_args)
# nvidia/DeepSeek-V4-Pro-NVFP4 uses the routed TRT-LLM runner.
if 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."
)
elif (
cfg.device == "cuda"
and not get_platform().is_hip
and cfg.moe_a2a_backend == "none"
and not envs.SGLANG_DSV4_FP4_DEQUANT.get()
and model_config.is_fp4_experts
and (
get_platform().is_sm90
or get_platform().is_sm100
or get_platform().is_sm120
)
):
overrides["moe_runner_backend"] = "flashinfer_mxfp4"
logger.info(
"Use flashinfer_mxfp4 as MoE runner backend for " f"{model_arch}."
)
return overrides
@@ -0,0 +1,23 @@
"""Config-time override declarations for exaone.
Architectures: Exaone4ForCausalLM, ExaoneMoEForCausalLM.
"""
import logging
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
)
logger = logging.getLogger(__name__)
@_register_for("Exaone4ForCausalLM", "ExaoneMoEForCausalLM")
def _exaone_overrides(server_args: Any, hf_config: Any) -> dict:
if hf_config.sliding_window_pattern is not None:
logger.warning(
f"Disabling hybrid SWA memory for {hf_config.architectures[0]} as it is not yet supported."
)
return {"disable_hybrid_swa_memory": True}
return {}
@@ -0,0 +1,22 @@
"""Config-time override declarations for falcon_h1.
Architectures: FalconH1ForCausalLM, JetNemotronForCausalLM, JetVLMForConditionalGeneration.
"""
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
@_register_for(
"FalconH1ForCausalLM", "JetNemotronForCausalLM", "JetVLMForConditionalGeneration"
)
def _falcon_h1_jet_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if get_platform().is_sm100 and cfg.attention_backend is None:
return {"attention_backend": "triton"}
return {}
@@ -0,0 +1,29 @@
"""Config-time override declarations for gemma2_gemma3.
Architectures: Gemma2ForCausalLM, Gemma3ForCausalLM, Gemma3ForConditionalGeneration, Gemma3nForCausalLM, Gemma3nForConditionalGeneration.
"""
import logging
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
)
logger = logging.getLogger(__name__)
@_register_for(
"Gemma2ForCausalLM",
"Gemma3ForCausalLM",
"Gemma3ForConditionalGeneration",
"Gemma3nForCausalLM",
"Gemma3nForConditionalGeneration",
)
def _gemma2_gemma3_overrides(server_args: Any, hf_config: Any) -> dict:
# FIXME: https://github.com/sgl-project/sglang/pull/7367 is not compatible with gemma2 model.
# It failed at this test: https://github.com/sgl-project/sglang/actions/runs/16255155597/job/45890331952#step:4:736
logger.warning(
f"Disable hybrid SWA memory for {hf_config.architectures[0]} as it is not yet supported."
)
return {"disable_hybrid_swa_memory": True}
@@ -0,0 +1,47 @@
"""Config-time override declarations for gemma4.
Architectures: Gemma4ForCausalLM, Gemma4ForConditionalGeneration, Gemma4UnifiedForConditionalGeneration.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
is_attention_backend_not_set,
model_config_of,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for(
"Gemma4ForConditionalGeneration",
"Gemma4ForCausalLM",
"Gemma4UnifiedForConditionalGeneration",
)
def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
default_attention_backend = "trtllm_mha" if get_platform().is_sm100 else "triton"
if is_attention_backend_not_set(cfg):
logger.info(
f"Use {default_attention_backend} as default attention backend for Gemma4"
)
overrides["attention_backend"] = default_attention_backend
# If only one split backend is set, keep the other side on a
# Gemma4-compatible fallback instead of letting generic backend selection
# choose an unsupported backend later.
elif cfg.attention_backend is None:
overrides["attention_backend"] = default_attention_backend
if get_platform().is_sm100 and cfg.moe_runner_backend == "auto":
if model_config_of(server_args).quantization == "modelopt_fp4":
overrides["quantization"] = "modelopt_fp4"
overrides["moe_runner_backend"] = "flashinfer_trtllm"
logger.info(
"Use flashinfer_trtllm as MoE runner backend on "
"SM100 for Gemma-4 (modelopt_fp4)"
)
return overrides
@@ -0,0 +1,50 @@
"""Config-time override declarations for glm4_moe.
Architectures: Glm4MoeForCausalLM.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for("Glm4MoeForCausalLM")
def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
if get_platform().is_sm100:
quantization_config = getattr(hf_config, "quantization_config", None)
quant_method = (
quantization_config.get("quant_method")
if quantization_config is not None
else None
)
quantization = cfg.quantization
if (
quantization is None
and not server_args._quantization_explicitly_unset
and quant_method is not None
):
overrides["quantization"] = quant_method
quantization = quant_method
if (
quantization in {"modelopt_fp4", None}
and cfg.moe_a2a_backend == "none"
and cfg.moe_runner_backend == "auto"
):
overrides["moe_runner_backend"] = "flashinfer_trtllm"
logger.info(
"Use flashinfer_trtllm as MoE runner backend on sm100 for Glm4MoeForCausalLM"
)
logger.info(
"Enable TF32 matmul for Glm4MoeForCausalLM model to improve gate gemm performance."
)
overrides["enable_tf32_matmul"] = True
return overrides
@@ -0,0 +1,127 @@
"""Config-time override declarations for gpt_oss.
Architectures: GptOssForCausalLM.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
is_attention_backend_not_set,
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import (
get_nvidia_driver_version,
is_cpu,
is_mps,
is_triton_kernels_available,
)
logger = logging.getLogger(__name__)
@_register_for("GptOssForCausalLM")
def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
# Set attention backend for GPT-OSS
if is_attention_backend_not_set(cfg):
if get_platform().is_sm100:
overrides["attention_backend"] = "trtllm_mha"
elif get_platform().is_sm90:
overrides["attention_backend"] = "fa3"
elif is_cpu() and get_platform().has_amx:
overrides["attention_backend"] = "intel_amx"
elif get_platform().is_xpu:
overrides["attention_backend"] = "intel_xpu"
elif get_platform().is_hip:
overrides["attention_backend"] = "aiter"
elif not (is_mps() and use_mlx()):
# Exempt MLX only -- it owns attention in its own runner. macOS
# without MLX still falls through to triton and fails fast below,
# rather than landing on torch_native (no sliding window, no sinks).
overrides["attention_backend"] = "triton"
if get_platform().is_xpu:
# Check for bf16 dtype on Intel XPU. Reads the pristine dtype request,
# which equals the legacy mid-branch read: dtype had no earlier writer
# for this arch.
if cfg.dtype == "auto":
logger.warning(
"GptOssForCausalLM on Intel XPU currently supports bfloat16 dtype only"
)
elif cfg.dtype not in ["bfloat16"]:
raise NotImplementedError(
f"GptOssForCausalLM on Intel XPU only supports bfloat16 dtype, "
f"but got '{cfg.dtype}'. Please use --dtype bfloat16 or remove --dtype to use auto."
)
quantization_config = getattr(hf_config, "quantization_config", None)
is_mxfp4_quant_format = (
quantization_config is not None
and quantization_config.get("quant_method") == "mxfp4"
)
if is_mxfp4_quant_format:
# use bf16 for mxfp4 triton kernels
overrides["dtype"] = "bfloat16"
if cfg.moe_runner_backend == "auto":
if get_platform().is_sm100 and is_mxfp4_quant_format:
overrides["moe_runner_backend"] = "flashinfer_mxfp4"
logger.warning(
"Detected SM100 and MXFP4 quantization format for GPT-OSS model, enabling FlashInfer MXFP4 MOE kernel."
)
elif get_platform().is_sm120 and is_mxfp4_quant_format:
overrides["moe_runner_backend"] = "flashinfer_mxfp4"
logger.warning(
"Detected SM120 and MXFP4 quantization format for GPT-OSS model, "
"enabling FlashInfer CUTLASS MXFP4 MOE kernel."
)
elif (
get_platform().is_hip and envs.SGLANG_USE_AITER.get()
) and is_mxfp4_quant_format:
overrides["moe_runner_backend"] = "auto"
logger.warning(
"Detected ROCm and MXFP4 quantization format for GPT-OSS model, enabling aiter MXFP4 MOE kernel."
)
## The AITER MXFP4 fused-MoE path for GPT-OSS expects the
## SEPARATED gate/up tile layout (matches the
## `gptoss_fp4_tuned_fmoe.csv` flydsl entries and the
## Mxfp4MoEMethod weight shuffle). Other AITER MXFP4
## callers default to INTERLEAVE; opt this path out
## unless the user explicitly overrode it.
# envs.SGLANG_USE_AITER_MOE_GU_ITLV.set(False)
elif get_platform().is_hip and envs.SGLANG_USE_AITER.get():
# For GPT-OSS bf16 on ROCm with aiter, use triton backend
# because aiter CK kernel doesn't support all GEMM dimensions
overrides["moe_runner_backend"] = "triton"
logger.warning(
"Detected ROCm with SGLANG_USE_AITER for GPT-OSS bf16 model, using triton MOE kernel."
)
elif get_platform().is_musa and envs.SGLANG_DEEPEP_BF16_DISPATCH.get():
overrides["moe_runner_backend"] = "deep_gemm"
logger.warning(
"Detected MUSA with SGLANG_DEEPEP_BF16_DISPATCH for bf16 model, using deep_gemm kernel."
)
elif (
cfg.ep_size == 1
and is_triton_kernels_available()
and cfg.quantization is None
and not (is_cpu() and get_platform().has_amx)
):
# The triton_kernels package segfaults on Blackwell (B200)
# with NVIDIA driver >= 595. Fall back to triton backend.
if get_platform().is_blackwell and get_nvidia_driver_version() >= (595,):
overrides["moe_runner_backend"] = "triton"
logger.warning(
"Detected GPT-OSS model on Blackwell with driver >= 595, "
"using triton MOE kernel to avoid triton_kernels SIGSEGV."
)
else:
overrides["moe_runner_backend"] = "triton_kernel"
logger.warning(
"Detected GPT-OSS model, enabling triton_kernels MOE kernel."
)
return overrides
@@ -0,0 +1,23 @@
"""Config-time override declarations for granitemoehybrid.
Architectures: GraniteMoeHybridForCausalLM.
"""
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
@_register_for("GraniteMoeHybridForCausalLM")
def _granite_moe_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
has_mamba = any(
layer_type == "mamba" for layer_type in getattr(hf_config, "layer_types", [])
)
if has_mamba and get_platform().is_sm100 and cfg.attention_backend is None:
return {"attention_backend": "flashinfer"}
return {}
@@ -0,0 +1,76 @@
"""Config-time override declarations for inkling.
Architectures: InklingForConditionalGeneration, InklingForConditionalGenerationMTP.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
is_attention_backend_not_set,
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for(
"InklingForConditionalGeneration",
"InklingForConditionalGenerationMTP",
)
def _inkling_overrides(server_args: Any, hf_config: Any) -> dict:
"""Inkling architecture defaults: SWA / mamba KV-pool ratios tuned for the
hybrid-SWA layout, the extra-buffer mamba strategy, and the unified radix
tree (which Inkling requires — models/inkling.py asserts it). The full-graph
prefill default is set separately (inline, before cuda-graph resolution) —
see ServerArgs.__post_init__ / _apply_inkling_prefill_cuda_graph_default. The
server-arg defaults each yield to an explicit user value (compared against
the ServerArgs class default); the prefill declaration is materialized
before _parse_cuda_graph_config folds cuda_graph_backend_prefill into
prefill.backend, and an explicit --cuda-graph-backend-prefill /
--disable-prefill-cuda-graph still wins. The unified-radix env write follows
the MiniMax-M3 handler precedent (env is not a resolvable server-arg)."""
cfg = resolving_view(server_args)
from sglang.srt.server_args import ServerArgs
overrides: Dict[str, Any] = {}
# NOTE: the full-graph prefill default is NOT set here. cuda-graph config is
# resolved in __post_init__ before declarations are materialized, so a
# cuda_graph_backend_prefill declared here lands too late (the breakable
# default would already have been auto-disabled for this multimodal arch).
# It is set inline before _handle_cuda_graph_config instead.
if cfg.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio:
overrides["swa_full_tokens_ratio"] = 0.1
if cfg.mamba_full_memory_ratio == ServerArgs.mamba_full_memory_ratio:
overrides["mamba_full_memory_ratio"] = 0.1
# Inkling requires the extra-buffer mamba strategy (inkling.py asserts
# enable_mamba_extra_buffer()); the generic "auto" resolution does not cover
# Inkling, so pin it here. Yields to an explicit --mamba-scheduler-strategy.
#
# The default comparison answers "unset" only while nothing has declared the
# field first. `_mamba_radix_cache_resolution` would, from the slot just
# above `collect_model_override_declarations`, for an architecture whose
# linear-attention spec sets `uses_mamba_radix_cache`. Inkling has no such
# spec; giving it one silently stops this pin from firing, so compare
# against the unresolved token ("auto") if that day comes.
if cfg.mamba_radix_cache_strategy == ServerArgs.mamba_radix_cache_strategy:
overrides["mamba_radix_cache_strategy"] = "extra_buffer"
# Inkling attention runs only on the fa4 (Blackwell) or triton backends --
# models/inkling_common/attn.py asserts attention_backend in {fa4, triton}.
# The generic resolver would otherwise pick trtllm_mha (SM100) / fa3
# (Hopper), so a bare launch fails on the first attention forward. Pin a
# supported default when the user left every attention-backend flag unset
# (mirrors the MiniMax-M3 SM100 fa4-default above); an explicit
# --attention-backend / --prefill/decode-attention-backend still wins.
if is_attention_backend_not_set(cfg):
inkling_attn_backend = "fa4" if get_platform().is_sm100 else "triton"
overrides["attention_backend"] = inkling_attn_backend
logger.info(
f"Use {inkling_attn_backend} as the attention backend for Inkling "
"(requires fa4 or triton)."
)
envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.set(True)
return overrides
@@ -0,0 +1,20 @@
"""Config-time override declarations for interns2_mobius.
Architectures: InternS2MobiusForConditionalGeneration.
"""
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
@_register_for("InternS2MobiusForConditionalGeneration")
def _interns2_mobius_baseline_overrides(server_args: Any, hf_config: Any) -> dict:
"""Select the only MoE runner validated for the 2,560-expert baseline."""
cfg = resolving_view(server_args)
if cfg.moe_runner_backend == "auto":
return {"moe_runner_backend": "triton_kernel"}
return {}
@@ -0,0 +1,191 @@
"""Config-time override declarations for kimi_k3.
Architectures: KimiK3ForConditionalGeneration.
"""
import inspect
import logging
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_dspark_verify_on_decode_backend,
_is_mxfp4_pack_quantized,
_register_for,
attention_backends_of,
is_attention_backend_not_set,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import get_device_name, is_mnnvl_fabric_device
logger = logging.getLogger(__name__)
def _require_kimi_k3_cutedsl_dcp_support() -> None:
try:
from flashinfer.decode import trtllm_batch_decode_with_kv_cache_mla
parameters = inspect.signature(trtllm_batch_decode_with_kv_cache_mla).parameters
except (ImportError, TypeError, ValueError) as exc:
raise RuntimeError(
"Kimi-K3 DCP with decode_attention_backend='cutedsl_mla' requires "
"FlashInfer 0.6.17 or newer with "
"trtllm_batch_decode_with_kv_cache_mla exposing enable_dcp."
) from exc
if "enable_dcp" not in parameters:
raise RuntimeError(
"Kimi-K3 DCP with decode_attention_backend='cutedsl_mla' requires "
"enable_dcp in the signature of "
"flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla; upgrade "
"to FlashInfer 0.6.17 or newer."
)
@_register_for("KimiK3ForConditionalGeneration")
def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if cfg.dcp_size > 1:
overrides = {}
if cfg.enable_symm_mem:
logger.warning(
"Kimi-K3 DCP disables --enable-symm-mem due to decode CUDA "
"graph correctness issues."
)
overrides["enable_symm_mem"] = False
if cfg.speculative_algorithm == "DSPARK":
from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode,
read_ragged_verify_mode,
)
ragged_mode = read_ragged_verify_mode()
if ragged_mode is not RaggedVerifyMode.STATIC:
raise ValueError(
"Kimi-K3 DCP + DSPARK currently requires "
"SGLANG_RAGGED_VERIFY_MODE=static; compact/cap-accept are "
f"not validated under DCP (got {ragged_mode.value!r})."
)
# DSPARK target-verify + draft-extend must run on the decode
# (cutedsl_mla) backend, whose _run_decode_kernel implements the DCP
# signature (causal_seqs / cp_world / cp_rank). The default
# "prefill" routes verify to trtllm_mla, whose base _run_decode_kernel
# lacks that DCP path (TypeError: unexpected kwarg 'causal_seqs').
overrides["speculative_attention_mode"] = "decode"
prefill_backend, decode_backend = attention_backends_of(cfg)
if decode_backend == "cutedsl_mla" or decode_backend is None:
_require_kimi_k3_cutedsl_dcp_support()
logger.info(
"Kimi-K3 DCP keeps decode attention backend 'cutedsl_mla' "
f"(prefill={prefill_backend!r} -> 'trtllm_mla')."
)
overrides.update(
prefill_attention_backend="trtllm_mla",
decode_attention_backend="cutedsl_mla",
)
elif decode_backend == "tokenspeed_mla":
logger.info(
"Kimi-K3 DCP overrides attention backends: "
f"prefill={prefill_backend!r}, decode={decode_backend!r} -> "
"'tokenspeed_mla'."
)
logger.info(
"Kimi-K3 DCP with tokenspeed mla backend overrides KV cache dtype: "
f"{cfg.kv_cache_dtype!r} -> 'fp8_e4m3'."
)
overrides.update(
prefill_attention_backend="tokenspeed_mla",
decode_attention_backend="tokenspeed_mla",
kv_cache_dtype="fp8_e4m3",
)
else:
raise AssertionError(
f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla' or 'tokenspeed_mla', got {decode_backend!r}."
)
if cfg.dcp_replicate_q_proj is None:
logger.info("Kimi-K3 DCP enables replicated Q projection by default.")
overrides["dcp_replicate_q_proj"] = True
device_name = get_device_name()
dcp_comm_backend = "fi_a2a" if is_mnnvl_fabric_device() else "a2a"
logger.info(
"Kimi-K3 DCP selects communication backend on "
f"{device_name!r}: {cfg.dcp_comm_backend!r} -> "
f"{dcp_comm_backend!r}."
)
overrides["dcp_comm_backend"] = dcp_comm_backend
return overrides
if not (get_platform().is_sm100 and get_platform().device_sm in (100, 103)):
return {}
backends_unset = is_attention_backend_not_set(cfg)
if cfg.speculative_algorithm != "DSPARK":
if not backends_unset:
return {}
logger.info(
"Use trtllm_mla as the default prefill and decode attention "
"backend for Kimi-K3 on SM100/SM103."
)
return {
"decode_attention_backend": "trtllm_mla",
"prefill_attention_backend": "trtllm_mla",
}
# DSPARK: verify runs on the decode backend (mode=decode below), so this
# picks the verify kernel -- mode=prefill routes it to flashinfer, which is
# slow and syncs, while plain decode is cold under dspark.
q_len = cfg.speculative_num_draft_tokens or (
cfg.speculative_dspark_block_size + 1
if cfg.speculative_dspark_block_size is not None
# Checkpoint auto-infer happens after overrides; K3 draft uses block 7.
else 8
)
overrides = {}
if backends_unset:
backend = "trtllm_mla"
overrides["decode_attention_backend"] = backend
overrides["prefill_attention_backend"] = "trtllm_mla"
else:
# Explicit backend knobs keep priority, but the mode is a separate knob
# that still needs declaring -- else verify stays on the prefill backend,
# whose host-side plan (flashinfer by default) forces a per-step D2H.
_, backend = attention_backends_of(cfg)
if _dspark_verify_on_decode_backend(backend, q_len, cfg.kv_cache_dtype):
overrides["speculative_attention_mode"] = "decode"
logger.info(
"Kimi-K3 DSPARK on SM100/SM103: decode/verify attention backend "
f"{backend} (speculative_attention_mode=decode)."
)
else:
logger.warning(
f"Kimi-K3 DSPARK: decode attention backend {backend!r} cannot serve "
f"target verify at q_len={q_len}, so verify runs on the prefill "
"backend (speculative_attention_mode=prefill). A host-plan prefill "
"backend costs a per-step seq_lens D2H sync; leave the attention "
"backend knobs unset for the sync-free default."
)
return overrides
@_register_for("KimiK3ForConditionalGeneration")
def _kimi_k3_moe_runner_overrides(server_args: Any, hf_config: Any) -> dict:
# MoE runner default, independent of the attention-backend gate above.
# trtllm-gen fused MoE (flashinfer_mxfp4) beats marlin on both the decode
# (M=bs) and the target-verify (M=bs*(gamma+1)) regimes on SM100/SM103.
# SM107 uses the same packed-MXFP4 runner; leaving auto unresolved falls
# back to BF16 weight materialization during model loading.
cfg = resolving_view(server_args)
if cfg.moe_runner_backend != "auto":
return {}
if not (get_platform().is_sm100 and get_platform().device_sm in (100, 103, 107)):
return {}
if not _is_mxfp4_pack_quantized(hf_config):
return {}
logger.info(
"Kimi-K3 on SM100/SM103/SM107: moe_runner_backend=flashinfer_mxfp4 "
"(FlashInfer SiTU kernels)."
)
return {"moe_runner_backend": "flashinfer_mxfp4"}
@@ -0,0 +1,20 @@
"""Config-time override declarations for lfm2.
Architectures: Lfm2ForCausalLM, Lfm2MoeForCausalLM.
"""
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
@_register_for("Lfm2ForCausalLM", "Lfm2MoeForCausalLM")
def _lfm2_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if get_platform().is_sm100 and cfg.attention_backend is None:
return {"attention_backend": "flashinfer"}
return {}
@@ -0,0 +1,47 @@
"""Config-time override declarations for llama4.
Architectures: Llama4ForCausalLM, Llama4ForConditionalGeneration.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
# Keep in sync with LLAMA4_MODEL_ARCHS (server_args.py).
@_register_for("Llama4ForConditionalGeneration", "Llama4ForCausalLM")
def _llama4_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if cfg.device == "cpu":
return {}
overrides: Dict[str, Any] = {}
# Auto-select attention backend for Llama4 if not specified
if cfg.attention_backend is None:
if get_platform().is_sm100:
backend, platform = "trtllm_mha", "sm100"
elif get_platform().is_sm90:
backend, platform = "fa3", "sm90"
elif get_platform().is_hip:
backend, platform = "aiter", "hip"
elif cfg.device == "xpu":
backend, platform = "intel_xpu", "xpu"
else:
backend, platform = "triton", "other platforms"
logger.warning(
f"Use {backend} as attention backend on {platform} for Llama4 model"
)
overrides["attention_backend"] = backend
if get_platform().is_sm100 and cfg.moe_runner_backend == "auto":
if cfg.quantization in {"fp8", "modelopt_fp8"}:
overrides["moe_runner_backend"] = "flashinfer_trtllm"
logger.info(
"Use flashinfer_trtllm as MoE runner backend on SM100 for Llama4"
)
return overrides
@@ -0,0 +1,37 @@
"""Config-time override declarations for mimo_v2.
Architectures: MiMoV2FlashForCausalLM, MiMoV2ForCausalLM.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import get_quantization_config
logger = logging.getLogger(__name__)
# Keep in sync with MIMO_V2_MODEL_ARCHS (server_args.py / configs/hf_config.py).
@_register_for("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM")
def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
if cfg.speculative_algorithm == "EAGLE":
logger.info("Enable multi-layer EAGLE speculative decoding for MiMoV2 model.")
overrides["enable_multi_layer_eagle"] = True
# On Blackwell "auto" falls through to the triton fused-MoE runner, ~12%
# slower at bs=1 decode. FP4 checkpoints use flashinfer_mxfp4 instead.
if (
get_platform().is_sm100
and cfg.moe_runner_backend == "auto"
and get_quantization_config(hf_config) == "fp8"
):
overrides["moe_runner_backend"] = "flashinfer_trtllm"
logger.info("MiMoV2 FP8 on SM100: moe_runner_backend=flashinfer_trtllm.")
return overrides
@@ -0,0 +1,66 @@
"""Config-time override declarations for minicpm.
Architectures: MiniCPMForCausalLM, MiniCPMSALAForCausalLM.
"""
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
is_attention_backend_not_set,
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_platform
@_register_for("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM")
def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if cfg.enable_dp_attention:
raise ValueError("MiniCPM does not support DP attention")
has_sparse_attention = getattr(hf_config, "has_minicpm_sparse_attention", False)
has_hybrid_attention = has_sparse_attention or getattr(
hf_config, "has_lightning_layers", False
)
overrides: Dict[str, Any] = {}
if has_hybrid_attention:
if cfg.enable_hierarchical_cache:
raise ValueError("MiniCPM SALA does not support hierarchical cache")
overrides["disable_radix_cache"] = True
if envs.SGLANG_MINICPM_FORCE_DENSE.get():
dense_backends = {
"minicpm_flashattn": ("fa4" if get_platform().is_blackwell else "fa3"),
"minicpm_flashinfer": "flashinfer",
}
# Literal keys keep the written-field set statically derivable; a loop
# variable hides it from the census in test_chain_read_ratchet.py.
dense_attention = dense_backends.get(cfg.attention_backend)
if dense_attention is not None:
overrides["attention_backend"] = dense_attention
dense_prefill = dense_backends.get(cfg.prefill_attention_backend)
if dense_prefill is not None:
overrides["prefill_attention_backend"] = dense_prefill
dense_decode = dense_backends.get(cfg.decode_attention_backend)
if dense_decode is not None:
overrides["decode_attention_backend"] = dense_decode
elif has_sparse_attention:
uses_sparse_backend = is_attention_backend_not_set(cfg) or any(
backend in ("minicpm_flashattn", "minicpm_flashinfer")
for backend in (
cfg.attention_backend,
cfg.prefill_attention_backend,
cfg.decode_attention_backend,
)
)
if uses_sparse_backend and cfg.disaggregation_mode != "null":
raise ValueError(
"MiniCPM sparse attention does not support PD disaggregation"
)
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = (
"minicpm_flashinfer"
if get_platform().is_blackwell
else "minicpm_flashattn"
)
return overrides
@@ -0,0 +1,20 @@
"""Config-time override declarations for minicpmv.
Architectures: MiniCPMV4_6ForConditionalGeneration.
"""
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
@_register_for("MiniCPMV4_6ForConditionalGeneration")
def _minicpm_v4_6_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if get_platform().is_sm100 and cfg.attention_backend is None:
return {"attention_backend": "triton"}
return {}
@@ -0,0 +1,36 @@
"""Config-time override declarations for minimax_m2.
Architectures: MiniMaxM2ForCausalLM.
"""
import logging
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
model_config_of,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for("MiniMaxM2ForCausalLM")
def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides = {"enable_tf32_matmul": True}
logger.info(
"Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance."
)
if (
get_platform().is_sm100
and cfg.moe_runner_backend == "auto"
and model_config_of(server_args).quantization == "modelopt_fp4"
):
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
logger.info(
"Use flashinfer_trtllm_routed as MoE runner backend on SM10X "
"for MiniMaxM2ForCausalLM with modelopt_fp4."
)
return overrides
@@ -0,0 +1,150 @@
"""Config-time override declarations for minimax_m3.
Architectures: MiniMaxM3SparseForCausalLM, MiniMaxM3SparseForConditionalGeneration.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
is_attention_backend_not_set,
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import get_quantization_config
logger = logging.getLogger(__name__)
@_register_for("MiniMaxM3SparseForCausalLM", "MiniMaxM3SparseForConditionalGeneration")
def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
quant_method = get_quantization_config(hf_config)
quant_resolved = cfg.quantization
if (
quant_resolved is None
and not server_args._quantization_explicitly_unset
and quant_method is not None
):
overrides["quantization"] = quant_method
quant_resolved = quant_method
if get_platform().is_hip:
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = "triton"
if cfg.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
overrides["moe_runner_backend"] = "triton"
if not envs.USE_ROCM_AITER_ROPE_BACKEND.is_set():
envs.USE_ROCM_AITER_ROPE_BACKEND.set("0")
aiter_fusion_resolved = cfg.enable_aiter_allreduce_fusion
if cfg.ep_size > 1 and cfg.moe_a2a_backend == "none" and aiter_fusion_resolved:
logger.warning(
"Disable --enable-aiter-allreduce-fusion for MiniMax-M3 "
"standard EP on ROCm because the deferred fused all-reduce "
"corrupts sparse MoE partial outputs."
)
overrides["enable_aiter_allreduce_fusion"] = False
aiter_fusion_resolved = False
# By default MiniMax-M3 on ROCm keeps NCCL all-reduce (custom AR off)
# whenever aiter all-reduce fusion is not used. Opting in via
# SGLANG_M3_ALLOW_CUSTOM_AR keeps custom all-reduce enabled so the
# quick-reduce path (ROCM_QUICK_REDUCE_QUANTIZATION=INT4/INT6/INT8) can
# accelerate the large prefill all-reduce.
if not aiter_fusion_resolved and not envs.SGLANG_M3_ALLOW_CUSTOM_AR.get():
overrides["disable_custom_all_reduce"] = True
elif get_platform().is_sm100:
if is_attention_backend_not_set(cfg):
if (
cfg.kv_cache_dtype == "fp8_e4m3"
and not envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get()
):
# fp8 attention GEMMs activate whenever possible
# (m3_fp8_attn_gemm_enabled); only trtllm_mha serves the dense
# fp8-q path, so prefer it over fa4 for fp8 KV. The
# SGLANG_DISABLE_M3_FP8_ATTN_GEMM kill switch keeps the fa4
# default (pre-fp8 behavior).
overrides["attention_backend"] = "trtllm_mha"
else:
overrides["attention_backend"] = "fa4"
backend_resolved = overrides.get("attention_backend", cfg.attention_backend)
page_resolved = cfg.page_size
# fa4 (fmha_sm100) and trtllm_mha both allow the page_size == 128
# sparse block MSA needs (trtllm_mha via trtllm-gen's dynamic
# tokens-per-page kernels).
if page_resolved is None and backend_resolved in ("fa4", "trtllm_mha"):
overrides["page_size"] = 128
page_resolved = 128
if cfg.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
overrides["moe_runner_backend"] = "deep_gemm"
elif cfg.moe_runner_backend == "auto" and quant_resolved == "modelopt_mixed":
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
logger.info(
"MiniMax-M3 on SM100: attention_backend="
f"{overrides.get('attention_backend', cfg.attention_backend)}, page_size={page_resolved}, "
f"moe_runner_backend={overrides.get('moe_runner_backend', cfg.moe_runner_backend)}."
)
elif get_platform().is_sm90:
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = "fa3"
page_resolved = cfg.page_size
if (
page_resolved is None
and overrides.get("attention_backend", cfg.attention_backend) == "fa3"
):
overrides["page_size"] = 128
page_resolved = 128
logger.info(
"MiniMax-M3 on Hopper: attention_backend="
f"{overrides.get('attention_backend', cfg.attention_backend)}, page_size={page_resolved} "
"(MSA is SM100-only; sparse attention runs on the Triton path)."
)
# fp8 attention GEMMs have no opt-in flag: m3_fp8_attn_gemm_enabled
# (server_args.py) derives the mode from kv_cache_dtype (fp8_e4m3) +
# attention_backend (trtllm_mha) + SM100 at runtime. Surface the
# resolution here: warn on fp8_e5m2 (fmha_sm100's variant lookup would
# silently dispatch the e4m3 kernel, so e5m2 stays on the widening Triton
# path), log when the fp8 GEMM mode is active, and log when the
# SGLANG_DISABLE_M3_FP8_ATTN_GEMM kill switch suppresses it.
if cfg.kv_cache_dtype == "fp8_e5m2":
logger.warning(
"MiniMax-M3 with kv_cache_dtype fp8_e5m2: fp8 attention GEMMs stay "
"DISABLED (fmha_sm100's variant lookup would silently dispatch the "
"e4m3 kernel for e5m2); sparse attention runs on the widening "
"Triton path. Use --kv-cache-dtype fp8_e4m3 for fp8 attention GEMMs."
)
elif (
cfg.kv_cache_dtype == "fp8_e4m3"
and overrides.get("attention_backend", cfg.attention_backend) == "trtllm_mha"
and get_platform().is_sm100
):
if envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get():
logger.info(
"MiniMax-M3 fp8 attention GEMMs DISABLED by "
"SGLANG_DISABLE_M3_FP8_ATTN_GEMM: bf16 indexer + widening "
"Triton sparse path, bf16 q; dense layers keep trtllm_mha's "
"fp8 KV cache."
)
else:
logger.info(
"MiniMax-M3 fp8 attention GEMMs active (kv_cache_dtype fp8_e4m3 + "
"trtllm_mha on SM100): fp8 main/index KV, fp8-cast q, fp8 "
"sparse/MSA kernels. Set SGLANG_DISABLE_M3_FP8_ATTN_GEMM=1 to "
"force the pre-fp8 numerics."
)
moe_runner_resolved = overrides.get("moe_runner_backend", cfg.moe_runner_backend)
if quant_resolved is None and moe_runner_resolved in ("auto", "deep_gemm"):
if moe_runner_resolved == "deep_gemm":
logger.warning(
"MiniMax-M3: the deep_gemm MoE runner produces corrupted output "
"on bf16 full weights; overriding --moe-runner-backend to 'triton'."
)
overrides["moe_runner_backend"] = "triton"
return overrides
@@ -0,0 +1,34 @@
"""Config-time override declarations for moss_vl.
Architectures: MossVLForConditionalGeneration.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
attention_backends_of,
is_attention_backend_not_set,
resolved_view,
resolving_view,
)
logger = logging.getLogger(__name__)
@_register_for("MossVLForConditionalGeneration")
def _moss_vl_overrides(server_args: Any, hf_config: Any) -> dict:
overrides: Dict[str, Any] = {}
if is_attention_backend_not_set(resolving_view(server_args)):
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 attention_backends_of(resolved_view(server_args))[0]
)
assert prefill_backend == "flashinfer", (
"MossVLForConditionalGeneration requires flashinfer prefill "
"attention backend for cross-attention custom mask support."
)
return overrides
@@ -0,0 +1,24 @@
"""Config-time override declarations for muse_glimmer.
Architectures: MuseGlimmerForCausalLM, MuseGlimmerForConditionalGeneration.
"""
import logging
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for("MuseGlimmerForConditionalGeneration", "MuseGlimmerForCausalLM")
def _muse_glimmer_fp4_gemm_runner_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if get_platform().is_sm120 and cfg.fp4_gemm_runner_backend == "auto":
logger.info("Use marlin as FP4 GEMM runner backend on SM120 for Muse Glimmer")
return {"fp4_gemm_runner_backend": "marlin"}
return {}
@@ -0,0 +1,123 @@
"""Config-time override declarations for nemotron_h.
Architectures: NemotronHForCausalLM, NemotronHPuzzleForCausalLM.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
is_attention_backend_not_set,
model_config_of,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_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)."""
cfg = resolving_view(server_args)
model_arch = hf_config.architectures[0]
model_config = model_config_of(server_args)
overrides: Dict[str, Any] = {}
is_modelopt = model_config.quantization in [
"modelopt",
"modelopt_fp8",
"modelopt_fp4",
"modelopt_mixed",
]
quantization = cfg.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
has_w4a16_moe_layers = False
if is_modelopt and quantization == "modelopt_mixed":
has_w4a16_moe_layers = any(
info.get("quant_algo") == "W4A16_NVFP4" and ".experts." in name
for name, info in hf_config.quantization_config.get(
"quantized_layers", {}
).items()
)
if has_w4a16_moe_layers:
if cfg.moe_a2a_backend != "none":
raise ValueError("W4A16_NVFP4 MoE layers require --moe-a2a-backend=none.")
if cfg.moe_runner_backend not in ("auto", "marlin"):
raise ValueError(
"W4A16_NVFP4 MoE layers require --moe-runner-backend=marlin."
)
if cfg.moe_runner_backend == "auto":
overrides["moe_runner_backend"] = "marlin"
logger.info(
"Use marlin as MoE runner backend for "
f"{model_arch} with W4A16_NVFP4 MoE layers"
)
elif (is_modelopt or model_config.quantization is None) and (
cfg.moe_runner_backend == "auto"
):
if get_platform().is_sm100 and cfg.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 get_platform().is_cuda
and (8, 0) <= get_platform().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 get_platform().is_blackwell and is_attention_backend_not_set(cfg):
if cfg.speculative_algorithm is not None:
speculative_algorithm = cfg.speculative_algorithm.upper()
if get_platform().is_sm100 and cfg.speculative_eagle_topk in (
None,
1,
):
overrides["attention_backend"] = "trtllm_mha"
if cfg.page_size is None:
overrides["page_size"] = 64
if cfg.mamba_radix_cache_strategy == "auto":
overrides["mamba_radix_cache_strategy"] = "extra_buffer"
if (
cfg.speculative_draft_attention_backend is None
and speculative_algorithm in ("EAGLE", "NEXTN", "DSPARK")
):
overrides["speculative_draft_attention_backend"] = "trtllm_mha"
else:
overrides["attention_backend"] = "triton"
if (
cfg.speculative_draft_attention_backend is None
and speculative_algorithm in ("EAGLE", "NEXTN", "DFLASH", "DSPARK")
):
overrides["speculative_draft_attention_backend"] = "flashinfer"
elif get_platform().is_sm100:
overrides["attention_backend"] = "trtllm_mha"
return overrides
@@ -0,0 +1,34 @@
"""Config-time override declarations for olmo2.
Architectures: Olmo2ForCausalLM.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for("Olmo2ForCausalLM")
def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
# FIXME: https://github.com/sgl-project/sglang/pull/7367 is not compatible with Olmo3 model.
logger.warning(
f"Disabling hybrid SWA memory for {hf_config.architectures[0]} as it is not yet supported."
)
overrides["disable_hybrid_swa_memory"] = True
if cfg.attention_backend is None:
if get_platform().is_cuda and get_platform().is_sm100:
overrides["attention_backend"] = "trtllm_mha"
elif get_platform().is_cuda and get_platform().device_sm >= 80:
overrides["attention_backend"] = "fa3"
else:
overrides["attention_backend"] = "triton"
return overrides
@@ -0,0 +1,54 @@
"""Config-time override declarations for qwen3_5.
Architectures: InternS2MobiusForConditionalGeneration, InternS2PreviewForConditionalGeneration, Qwen3NextForCausalLM, Qwen3_5ForConditionalGeneration, Qwen3_5MoeForConditionalGeneration.
"""
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
get_default_attn_backend,
mamba_extra_buffer_of,
model_config_of,
resolved_view,
resolving_view,
use_mla_backend,
)
from sglang.srt.runtime_context import get_platform
@_register_for(
"Qwen3NextForCausalLM",
"Qwen3_5MoeForConditionalGeneration",
"InternS2PreviewForConditionalGeneration",
"InternS2MobiusForConditionalGeneration",
"Qwen3_5ForConditionalGeneration",
)
def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if not get_platform().is_sm100 or cfg.attention_backend is not None:
return {}
sm100_default_attn_backend = "triton"
# trtllm_mha requires speculative_eagle_topk == 1 and page_size > 1.
# get_default_attn_backend handles the eagle_topk check.
# There is only one case where page_size=1 is required,
# which is when radix cache is enabled and both extra_buffer
# and spec decoding are disabled.
default_attn_backend = get_default_attn_backend(
server_args,
use_mla_backend=use_mla_backend(server_args),
model_config=model_config_of(server_args),
)
# The mamba radix-cache pass runs before this dispatch: read the
# declared strategy through the view (the legacy branch observed the
# already-written field here).
if default_attn_backend == "trtllm_mha" and not (
not mamba_extra_buffer_of(resolved_view(server_args))
and not cfg.disable_radix_cache
and cfg.speculative_algorithm is None
):
sm100_default_attn_backend = "trtllm_mha"
return {
"attention_backend": sm100_default_attn_backend,
"page_size": 64 if sm100_default_attn_backend == "trtllm_mha" else 1,
}
@@ -0,0 +1,50 @@
"""Config-time override declarations for qwen3_moe.
Architectures: InternS2PreviewForConditionalGeneration, Qwen3MoeForCausalLM, Qwen3NextForCausalLM, Qwen3VLMoeForConditionalGeneration, Qwen3_5ForConditionalGeneration, Qwen3_5MoeForConditionalGeneration.
"""
import logging
from typing import Any, Dict
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import get_quantization_config
logger = logging.getLogger(__name__)
@_register_for(
"Qwen3MoeForCausalLM",
"Qwen3VLMoeForConditionalGeneration",
"Qwen3NextForCausalLM",
"Qwen3_5MoeForConditionalGeneration",
"InternS2PreviewForConditionalGeneration",
"Qwen3_5ForConditionalGeneration",
)
def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
if get_platform().is_sm100:
quant_method = get_quantization_config(hf_config)
quantization = cfg.quantization
if (
quantization is None
and not server_args._quantization_explicitly_unset
and quant_method is not None
):
overrides["quantization"] = quant_method
quantization = quant_method
if (
(quantization in ("fp8", "modelopt_fp4") or quantization is None)
and cfg.moe_a2a_backend == "none"
and cfg.moe_runner_backend == "auto"
):
overrides["moe_runner_backend"] = "flashinfer_trtllm"
logger.info(
"Use flashinfer_trtllm as MoE runner backend on sm100 for "
f"{hf_config.architectures[0]}"
)
return overrides
@@ -0,0 +1,32 @@
"""Config-time override declarations for qwen3_vl.
Architectures: Qwen3VLForConditionalGeneration.
"""
import logging
from typing import Any
from sglang.srt.arg_groups.model_override_base import (
_register_for,
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@_register_for("Qwen3VLForConditionalGeneration")
def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
if (
get_platform().is_hip
and envs.SGLANG_USE_AITER_UNIFIED_ATTN.get()
and cfg.page_size is None
):
logger.info(
"Setting page_size=16 for aiter unified attention on Qwen3VLForConditionalGeneration."
)
return {"page_size": 16}
return {}
File diff suppressed because it is too large Load Diff
@@ -910,3 +910,32 @@ def handle_multimodal_feature_transport(server_args: Any):
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.set(
"1" if requested_transport == "cuda_ipc" else "0"
)
_ssl_verify_warned = False
def ssl_verify_of(cfg: Any):
"""What to pass as the requests library's ``verify=``.
A CA file means validate against it. SSL configured without one means
verification off -- self-signed certificates in development -- and that is
worth saying out loud, once. No SSL means the system CA bundle.
The warning is once per process: the message is about how this process was
configured, and a second engine repeating it says nothing new.
"""
global _ssl_verify_warned
if cfg.ssl_ca_certs:
return cfg.ssl_ca_certs
if cfg.ssl_certfile:
if not _ssl_verify_warned:
logger.warning(
"SSL is enabled but --ssl-ca-certs was not provided. Certificate "
"verification is DISABLED for internal health checks. For "
"production deployments, provide --ssl-ca-certs or use CA-signed "
"certificates."
)
_ssl_verify_warned = True
return False
return True
+3 -2
View File
@@ -483,6 +483,7 @@ from sglang.srt.entrypoints.v1_loads import router as v1_loads_router
v1_loads_router.route_class = ORJSONRoute
app.include_router(v1_loads_router)
from sglang.srt.arg_groups.serving_hook import ssl_verify_of
from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router
from sglang.srt.runtime_context import (
describe_kv_events_publisher,
@@ -2199,7 +2200,7 @@ def _execute_server_warmup(server_args: ServerArgs):
if server_args.api_key:
headers["Authorization"] = f"Bearer {server_args.api_key}"
ssl_verify = server_args.ssl_verify()
ssl_verify = ssl_verify_of(server_args)
# Wait until the server is launched
success = False
@@ -2375,7 +2376,7 @@ def _freeze_gc_after_server_warmup(server_args: ServerArgs):
server_args.url() + "/freeze_gc",
headers=freeze_headers,
timeout=10,
verify=server_args.ssl_verify(),
verify=ssl_verify_of(server_args),
)
res.raise_for_status()
except requests.exceptions.RequestException:
@@ -5,6 +5,7 @@ from typing import List, Optional, Tuple
import requests
import torch
from sglang.srt.arg_groups.serving_hook import ssl_verify_of
from sglang.srt.entrypoints.EngineBase import EngineBase
from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.server_args import ServerArgs
@@ -24,7 +25,7 @@ def launch_server_process(server_args: ServerArgs) -> multiprocessing.Process:
timeout = 300.0 # Increased timeout to 5 minutes for downloading large models
start_time = time.perf_counter()
ssl_verify = server_args.ssl_verify()
ssl_verify = ssl_verify_of(server_args)
with requests.Session() as session:
while time.perf_counter() - start_time < timeout:
@@ -74,7 +75,7 @@ class HttpServerEngineAdapter(EngineBase):
"""
url = f"{self.server_args.url()}/{endpoint}"
response = requests.post(
url, json=payload or {}, verify=self.server_args.ssl_verify()
url, json=payload or {}, verify=ssl_verify_of(self.server_args)
)
response.raise_for_status()
return response.json()
@@ -11,6 +11,10 @@ import msgspec
import torch
import torch.distributed as dist
from sglang.srt.arg_groups.overrides import (
modelexpress_transport_of,
modelexpress_url_of,
)
from sglang.srt.configs.device_config import DeviceConfig
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS
@@ -219,8 +223,8 @@ def build_load_config(
remote_instance_weight_loader_backend=get_model().remote_instance_weight_loader_backend,
remote_instance_weight_loader_transfer_engine=remote_instance_weight_transporter_engine,
remote_instance_weight_loader_transfer_engine_session_id=remote_instance_weight_transporter_session_id,
modelexpress_url=server_args.modelexpress_url,
modelexpress_transport=server_args.modelexpress_transport,
modelexpress_url=modelexpress_url_of(server_args),
modelexpress_transport=modelexpress_transport_of(server_args),
modelopt_config=modelopt_config,
rl_quant_profile=get_model().rl_quant_profile,
draft_model_idx=draft_model_idx,
+4 -4
View File
@@ -1717,8 +1717,8 @@ def pre_capture_activation_reserve_mb(gpu_mem: float | None) -> float:
Derived from published leaves across four bags (``disagg`` / ``schedule`` /
``exec.graph`` / ``spec``) plus the configured parallel sizes, so it follows
a post-publish override; ``ServerArgs.pre_capture_activation_reserve_mb`` is
the pre-publish equivalent and
a post-publish override; ``pre_capture_activation_reserve_mb_of`` in
``arg_groups.overrides`` is the config-shaped equivalent and
``TestDerivedPredicatesAgreeAcrossTiers`` pins the two equal.
"""
schedule = get_schedule()
@@ -2100,7 +2100,7 @@ def describe_kv_events_publisher(server_args: Any) -> Optional[dict]:
helpers the scheduler binds through — so the advertisement cannot
drift from the sockets.
"""
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.arg_groups.overrides import kv_event_block_size_of, resolving_view
# Lazy import so loading server_args doesn't pull in
# disaggregation / msgspec / zmq at module top level.
@@ -2135,7 +2135,7 @@ def describe_kv_events_publisher(server_args: Any) -> Optional[dict]:
"endpoint_host": host,
"endpoint_port_base": port,
"topic": cfg.topic,
"block_size": resolved.kv_event_block_size,
"block_size": kv_event_block_size_of(resolved),
"dp_size": resolved.dp_size,
}
# Load range, from the same resolver SchedulerLoadPublisher binds
+1 -85
View File
@@ -58,7 +58,6 @@ from sglang.srt.arg_groups.overrides import (
remote_instance_transfer_engine_of,
resolution_projection,
resolving_view,
supports_mamba_cache_extra_buffer,
)
from sglang.srt.environ import envs
from sglang.srt.function_call.function_call_parser import FunctionCallParser
@@ -3759,33 +3758,7 @@ class ServerArgs:
# CUDA graph configuration resolution
# ------------------------------------------------------------------
def pre_capture_activation_reserve_mb(self, gpu_mem: Optional[float]) -> float:
# Runtime activation working-set reserve for eager decode above the captured
# max_bs and transient prefill/logits; also covers fixed state caches.
cfg = resolving_view(self)
if cfg.disaggregation_mode == "decode":
running_requests = (
cfg.max_running_requests or cfg.cuda_graph_config.decode.max_bs or 1
)
activation_tokens = max(
running_requests * (cfg.speculative_num_draft_tokens or 1), 2048
)
elif cfg.chunked_prefill_size > 0:
activation_tokens = max(cfg.chunked_prefill_size, 2048)
else:
activation_tokens = max(cfg.max_prefill_tokens, 2048)
reserved_mem = (
512 + activation_tokens * 1.5 + cfg.tp_size * cfg.pp_size / 8 * 1024
)
if gpu_mem is not None and gpu_mem > 60 * 1024:
reserved_mem = max(reserved_mem, 10 * 1024)
return reserved_mem
def _support_mamba_cache_extra_buffer(self, model_arch: str):
return supports_mamba_cache_extra_buffer(self, model_arch)
# ===== END TO BE REFACTORED ====
# ===== END TO BE REFACTORED ====
LANGUAGE_MODEL_ONLY_ARCHITECTURES = ("MuseGlimmerForConditionalGeneration",)
@@ -4137,32 +4110,6 @@ class ServerArgs:
return cfg.startup_weight_load_mode == "overlap"
def ssl_verify(self):
"""Return the value for the requests library's verify= parameter.
When SSL is configured:
- If a CA certificate file is provided, return its path so requests
validates the server certificate against that CA.
- Otherwise, return False to disable certificate verification
(suitable for self-signed certificates in development/testing).
A warning is logged once when this happens.
When SSL is not configured, return True to use the system's default
CA bundle.
"""
if self.ssl_ca_certs:
return self.ssl_ca_certs
if self.ssl_certfile:
if not getattr(self, "_ssl_verify_warned", False):
logger.warning(
"SSL is enabled but --ssl-ca-certs was not provided. "
"Certificate verification is DISABLED for internal "
"health checks. For production deployments, provide "
"--ssl-ca-certs or use CA-signed certificates."
)
self._ssl_verify_warned = True
return False
return True
def __setattr__(self, name, value):
# Once resolution has finished the record is the READ-ONLY raw input
# the config bags were projected from. Resolved config changes go to the bags via
@@ -4190,42 +4137,11 @@ class ServerArgs:
check_server_args(self)
@property
def _parsed_modelexpress_config(self) -> dict:
cache = getattr(self, "_mx_config_cache", None)
if cache is not None:
return cache
if self.modelexpress_config is None:
result = {}
elif isinstance(self.modelexpress_config, str):
result = json.loads(self.modelexpress_config)
else:
result = self.modelexpress_config
self._mx_config_cache = result
return result
@property
def modelexpress_url(self) -> Optional[str]:
return self._parsed_modelexpress_config.get("url")
@property
def modelexpress_transport(self) -> str:
"""Transport backend for modelexpress."""
return self._parsed_modelexpress_config.get("transport", "nixl")
def remote_instance_weight_loader_use_transfer_engine(self, load_format=None):
"""``load_format`` overrides the seed's: a draft runner loading under
``--speculative-draft-load-format`` needs its own transfer engine."""
return remote_instance_transfer_engine_of(resolving_view(self), load_format)
@property
def kv_event_block_size(self) -> int:
"""Width KV events are emitted at: under DCP the radix tree pages at
``page_size * dcp_size`` (``mem_cache/kv_cache_builder.py``).
"""
cfg = resolving_view(self)
return cfg.page_size * self.dcp_size
# --------------------------------------------------------------------------
# Module-level ServerArgs helpers and runtime shims.