[refactor] Migrate the attention_backend resolution chain (stack 11/15) (#30073)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
5c95bf15c8
commit
abbb41a214
@@ -33,8 +33,23 @@ import logging
|
||||
from typing import Any, Callable, Dict, Iterable, List, Optional, Sequence, Tuple
|
||||
|
||||
from sglang.srt.arg_groups.arg_utils import model_overridable_fields
|
||||
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 is_flashinfer_available, is_xpu
|
||||
from sglang.srt.utils.common import (
|
||||
cpu_has_amx_support,
|
||||
get_device_sm,
|
||||
is_blackwell_supported,
|
||||
is_cpu,
|
||||
is_cuda,
|
||||
is_flashinfer_available,
|
||||
is_hip,
|
||||
is_npu,
|
||||
is_sm90_supported,
|
||||
is_sm100_supported,
|
||||
is_sm120_supported,
|
||||
is_xpu,
|
||||
xpu_has_xmx_support,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -250,6 +265,21 @@ def _exaone_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
@_register_for("GptOssForCausalLM")
|
||||
def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
overrides: Dict[str, Any] = {}
|
||||
# Set attention backend for GPT-OSS
|
||||
if server_args.is_attention_backend_not_set():
|
||||
if is_sm100_supported():
|
||||
overrides["attention_backend"] = "trtllm_mha"
|
||||
elif is_sm90_supported():
|
||||
overrides["attention_backend"] = "fa3"
|
||||
elif is_cpu() and cpu_has_amx_support():
|
||||
overrides["attention_backend"] = "intel_amx"
|
||||
elif is_xpu():
|
||||
overrides["attention_backend"] = "intel_xpu"
|
||||
elif is_hip():
|
||||
overrides["attention_backend"] = "aiter"
|
||||
else:
|
||||
overrides["attention_backend"] = "triton"
|
||||
if 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
|
||||
@@ -269,17 +299,111 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
and quantization_config.get("quant_method") == "mxfp4"
|
||||
):
|
||||
# use bf16 for mxfp4 triton kernels
|
||||
return {"dtype": "bfloat16"}
|
||||
overrides["dtype"] = "bfloat16"
|
||||
return overrides
|
||||
|
||||
|
||||
# Keep in sync with LLAMA4_MODEL_ARCHS (server_args.py).
|
||||
@_register_for("Llama4ForConditionalGeneration", "Llama4ForCausalLM")
|
||||
def _llama4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if server_args.device == "cpu":
|
||||
return {}
|
||||
# Auto-select attention backend for Llama4 if not specified
|
||||
if server_args.attention_backend is None:
|
||||
if is_sm100_supported():
|
||||
backend, platform = "trtllm_mha", "sm100"
|
||||
elif is_sm90_supported():
|
||||
backend, platform = "fa3", "sm90"
|
||||
elif is_hip():
|
||||
backend, platform = "aiter", "hip"
|
||||
elif server_args.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"
|
||||
)
|
||||
return {"attention_backend": backend}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for(
|
||||
"Gemma4ForConditionalGeneration",
|
||||
"Gemma4ForCausalLM",
|
||||
"Gemma4UnifiedForConditionalGeneration",
|
||||
)
|
||||
def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
default_attention_backend = "trtllm_mha" if is_sm100_supported() else "triton"
|
||||
if server_args.is_attention_backend_not_set():
|
||||
logger.info(
|
||||
f"Use {default_attention_backend} as default attention backend for Gemma4"
|
||||
)
|
||||
return {"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.
|
||||
if server_args.attention_backend is None:
|
||||
return {"attention_backend": default_attention_backend}
|
||||
return {}
|
||||
|
||||
|
||||
@_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:
|
||||
return {"attention_backend": "triton"}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for(
|
||||
"FalconH1ForCausalLM", "JetNemotronForCausalLM", "JetVLMForConditionalGeneration"
|
||||
)
|
||||
def _falcon_h1_jet_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm100_supported() and server_args.attention_backend is None:
|
||||
return {"attention_backend": "triton"}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for("GraniteMoeHybridForCausalLM")
|
||||
def _granite_moe_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
has_mamba = any(
|
||||
layer_type == "mamba" for layer_type in getattr(hf_config, "layer_types", [])
|
||||
)
|
||||
if has_mamba and is_sm100_supported() and server_args.attention_backend is None:
|
||||
return {"attention_backend": "flashinfer"}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for("Lfm2ForCausalLM")
|
||||
def _lfm2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm100_supported() and server_args.attention_backend is None:
|
||||
return {"attention_backend": "flashinfer"}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for("Glm4MoeForCausalLM")
|
||||
def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
logger.info(
|
||||
"Enable TF32 matmul for Glm4MoeForCausalLM model to improve gate gemm performance."
|
||||
)
|
||||
return {"enable_tf32_matmul": True}
|
||||
|
||||
|
||||
@_register_for("Olmo2ForCausalLM")
|
||||
def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
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."
|
||||
)
|
||||
return {"disable_hybrid_swa_memory": True}
|
||||
overrides["disable_hybrid_swa_memory"] = True
|
||||
if server_args.attention_backend is None:
|
||||
if is_cuda() and is_sm100_supported():
|
||||
overrides["attention_backend"] = "trtllm_mha"
|
||||
elif is_cuda() and get_device_sm() >= 80:
|
||||
overrides["attention_backend"] = "fa3"
|
||||
else:
|
||||
overrides["attention_backend"] = "triton"
|
||||
return overrides
|
||||
|
||||
|
||||
@register_model_override_predicate(
|
||||
@@ -288,6 +412,13 @@ def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
)
|
||||
def _step3p_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
overrides: Dict[str, Any] = {}
|
||||
if server_args.is_attention_backend_not_set():
|
||||
if is_blackwell_supported():
|
||||
logger.info("Auto-select fa4 attention backend for Step3p7 on Blackwell.")
|
||||
overrides["attention_backend"] = "fa4"
|
||||
elif is_sm90_supported():
|
||||
logger.info("Auto-select fa3 attention backend for Step3p7 on Hopper.")
|
||||
overrides["attention_backend"] = "fa3"
|
||||
if server_args.speculative_algorithm == "EAGLE":
|
||||
logger.info(
|
||||
"Enable multi-layer EAGLE speculative decoding for Step3p5ForCausalLM model."
|
||||
@@ -333,6 +464,155 @@ def _deterministic_sampling_backend(view: Any) -> dict:
|
||||
return {}
|
||||
|
||||
|
||||
def _deterministic_is_deepseek_model(view: Any) -> bool:
|
||||
"""Faithful copy of the deterministic handler's arch probe (pure read;
|
||||
the handler keeps its own copy for the later deepseek validation)."""
|
||||
from sglang.srt.connector import ConnectorType
|
||||
from sglang.srt.utils.common import parse_connector_type
|
||||
|
||||
if parse_connector_type(view.model_path) == ConnectorType.INSTANCE:
|
||||
return False
|
||||
try:
|
||||
hf_config = view.get_model_config().hf_config
|
||||
return hf_config.architectures[0] in [
|
||||
"DeepseekV2ForCausalLM",
|
||||
"DeepseekV3ForCausalLM",
|
||||
"DeepseekV32ForCausalLM",
|
||||
"MistralLarge3ForCausalLM",
|
||||
"PixtralForConditionalGeneration",
|
||||
"GlmMoeDsaForCausalLM",
|
||||
]
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
@register_post_process
|
||||
def _deterministic_attention_backend(view: Any) -> dict:
|
||||
if not view.enable_deterministic_inference:
|
||||
return {}
|
||||
from sglang.srt.server_args import DETERMINISTIC_ATTENTION_BACKEND_CHOICES
|
||||
|
||||
if view.attention_backend is None:
|
||||
# User didn't specify attention backend, fallback based on GPU architecture
|
||||
if is_sm100_supported() or is_sm120_supported():
|
||||
# Blackwell and newer architectures
|
||||
if _deterministic_is_deepseek_model(view):
|
||||
# fallback to triton for DeepSeek models because flashinfer
|
||||
# doesn't support deterministic inference for DeepSeek models yet
|
||||
backend = "triton"
|
||||
else:
|
||||
# fallback to flashinfer on Blackwell for non-DeepSeek models
|
||||
backend = "flashinfer"
|
||||
else:
|
||||
# Hopper (SM90) and older architectures
|
||||
backend = "fa3"
|
||||
logger.warning(
|
||||
f"Attention backend not specified. Falling back to '{backend}' for deterministic inference. "
|
||||
f"You can explicitly set --attention-backend to one of {DETERMINISTIC_ATTENTION_BACKEND_CHOICES}."
|
||||
)
|
||||
return {"attention_backend": backend}
|
||||
elif view.attention_backend not in DETERMINISTIC_ATTENTION_BACKEND_CHOICES:
|
||||
# User explicitly specified an incompatible attention backend
|
||||
raise ValueError(
|
||||
f"Currently only {DETERMINISTIC_ATTENTION_BACKEND_CHOICES} attention backends are supported for deterministic inference, "
|
||||
f"but you explicitly specified '{view.attention_backend}'."
|
||||
)
|
||||
return {}
|
||||
|
||||
|
||||
@register_post_process
|
||||
def _attention_backend_default(view: Any) -> dict:
|
||||
if view.prefill_attention_backend is not None and (
|
||||
view.prefill_attention_backend == view.decode_attention_backend
|
||||
): # override the default attention backend
|
||||
return {"attention_backend": view.prefill_attention_backend}
|
||||
if view.attention_backend is None:
|
||||
backend = view._get_default_attn_backend(
|
||||
view.use_mla_backend(), view.get_model_config()
|
||||
)
|
||||
logger.info(
|
||||
f"Attention backend not specified. Use {backend} backend by default."
|
||||
)
|
||||
return {"attention_backend": backend}
|
||||
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":
|
||||
logger.warning(
|
||||
"FlashAttention3 only supports fp8_e4m3 if using FP8; "
|
||||
"Setting attention backend to triton."
|
||||
)
|
||||
return {"attention_backend": "triton"}
|
||||
return {}
|
||||
|
||||
|
||||
@register_post_process
|
||||
def _attention_backend_platform_fallbacks(view: Any) -> dict:
|
||||
if (
|
||||
view.attention_backend == "intel_amx"
|
||||
and view.device == "cpu"
|
||||
and not cpu_has_amx_support()
|
||||
):
|
||||
logger.warning(
|
||||
"The current platform does not support Intel AMX, will fallback to torch_native backend."
|
||||
)
|
||||
return {"attention_backend": "torch_native"}
|
||||
if (
|
||||
view.attention_backend == "intel_xpu"
|
||||
and view.device == "xpu"
|
||||
and not xpu_has_xmx_support()
|
||||
):
|
||||
logger.warning(
|
||||
"The current platform does not support Intel XMX, will fallback to triton backend."
|
||||
)
|
||||
return {"attention_backend": "triton"}
|
||||
return {}
|
||||
|
||||
|
||||
@register_post_process
|
||||
def _attention_backend_dual_chunk(view: Any) -> dict:
|
||||
if (
|
||||
getattr(view.get_model_config().hf_config, "dual_chunk_attention_config", None)
|
||||
is not None
|
||||
):
|
||||
if view.attention_backend is None:
|
||||
logger.info("Dual chunk attention is turned on by default.")
|
||||
return {"attention_backend": "dual_chunk_flash_attn"}
|
||||
elif view.attention_backend != "dual_chunk_flash_attn":
|
||||
raise ValueError(
|
||||
"Dual chunk attention is enabled, but attention backend is set to "
|
||||
f"{view.attention_backend}. Please set it to 'dual_chunk_flash_attn'."
|
||||
)
|
||||
return {}
|
||||
|
||||
|
||||
@register_post_process
|
||||
def _dllm_attention_backend(view: Any) -> dict:
|
||||
if view.dllm_algorithm is None:
|
||||
return {}
|
||||
if is_hip():
|
||||
if view.attention_backend not in ["triton", "aiter"]:
|
||||
logger.warning(
|
||||
"Attention backend is set to triton for diffusion LLM inference on AMD GPUs"
|
||||
)
|
||||
return {"attention_backend": "triton"}
|
||||
elif is_npu():
|
||||
if view.attention_backend != "ascend":
|
||||
logger.warning(
|
||||
"Attention backend is overridden to 'ascend' when running on NPU for diffusion LLM inference."
|
||||
)
|
||||
return {"attention_backend": "ascend"}
|
||||
elif view.cuda_graph_config.decode.backend != Backend.DISABLED:
|
||||
if view.attention_backend != "flashinfer":
|
||||
logger.warning(
|
||||
"Attention backend is set to flashinfer because of enabling cuda graph in diffusion LLM inference"
|
||||
)
|
||||
return {"attention_backend": "flashinfer"}
|
||||
return {}
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class OverrideRecord:
|
||||
"""Provenance of one resolved write: ``base`` is the value before this
|
||||
@@ -444,6 +724,27 @@ def apply_declarations_to_server_args(
|
||||
setattr(server_args, field, value)
|
||||
|
||||
|
||||
def refresh_declared_fields(server_args: Any, fields: Iterable[str]) -> None:
|
||||
"""Transition helper for legacy code that overwrites a resolved field
|
||||
AFTER the override collection in ``__post_init__`` (e.g.
|
||||
``ModelRunner.model_specific_adjustment`` forcing ``attention_backend``
|
||||
for HRM-Text). Redeclares the live value so publish parity holds and the
|
||||
flags tier materializes the adjusted end state.
|
||||
"""
|
||||
_missing = object()
|
||||
declarations = server_args._resolved_overrides
|
||||
for field in fields:
|
||||
effective = _missing
|
||||
for _source, decl in declarations:
|
||||
if field in decl:
|
||||
effective = decl[field]
|
||||
if effective is _missing:
|
||||
continue
|
||||
live = getattr(server_args, field)
|
||||
if effective != live:
|
||||
declarations.append((f"runtime_adjustment[{field}]", {field: live}))
|
||||
|
||||
|
||||
def assert_flag_parity(
|
||||
flags: Any,
|
||||
server_args: Any,
|
||||
|
||||
@@ -1184,6 +1184,13 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if not server_args.disable_chunked_prefix_cache:
|
||||
log_info_on_rank0(logger, "Chunked prefix cache is turned on.")
|
||||
|
||||
# The imperative adjustments above may overwrite fields the resolution passes
|
||||
# already declared (HRM-Text forces attention_backend); redeclare the
|
||||
# adjusted values so publish parity holds.
|
||||
from sglang.srt.arg_groups.overrides import refresh_declared_fields
|
||||
|
||||
refresh_declared_fields(server_args, ("attention_backend",))
|
||||
|
||||
def check_quantized_moe_compatibility(self):
|
||||
if (
|
||||
quantization_config := getattr(
|
||||
|
||||
@@ -281,6 +281,10 @@ class _StaticFlags(_FlagGroupBase):
|
||||
class AttnFlags(_StaticFlags):
|
||||
"""Attention-family resolved flags (leaves arrive with the V3 sweeps)."""
|
||||
|
||||
# Resolved attention backend; the pristine user request stays on
|
||||
# server_args.attention_backend.
|
||||
backend: str | None = None
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class MoeFlags(_StaticFlags):
|
||||
@@ -327,8 +331,10 @@ class Flags(_StaticFlags):
|
||||
# Resolved-config field name → dotted flag-leaf path (e.g. a V3 sweep adds
|
||||
# "use_mla_backend": "attn.use_mla_backend"). Fields not listed default to a
|
||||
# flat leaf of the same name on the Flags container. Populated per field
|
||||
# family as readers migrate; empty in the skeleton.
|
||||
FLAG_LEAF_MAP: dict[str, str] = {}
|
||||
# family as readers migrate.
|
||||
FLAG_LEAF_MAP: dict[str, str] = {
|
||||
"attention_backend": "attn.backend",
|
||||
}
|
||||
|
||||
|
||||
def resolve_flag_leaf(
|
||||
|
||||
@@ -93,7 +93,6 @@ from sglang.srt.utils.common import (
|
||||
nullable_str,
|
||||
parse_connector_type,
|
||||
torch_release,
|
||||
xpu_has_xmx_support,
|
||||
)
|
||||
from sglang.srt.utils.hf_transformers_utils import check_gguf_file
|
||||
from sglang.srt.utils.network import NetworkAddress, get_free_port, wait_port_available
|
||||
@@ -1409,6 +1408,7 @@ class ServerArgs:
|
||||
Arg(
|
||||
help="Choose the kernels for attention layers.",
|
||||
choices=ATTENTION_BACKEND_CHOICES,
|
||||
model_overridable=True,
|
||||
),
|
||||
] = None
|
||||
decode_attention_backend: A[
|
||||
@@ -4059,23 +4059,8 @@ class ServerArgs:
|
||||
envs.SGLANG_EAGER_INPUT_NO_COPY.set(True)
|
||||
|
||||
elif model_arch in ["GptOssForCausalLM"]:
|
||||
# Set attention backend for GPT-OSS
|
||||
if self.is_attention_backend_not_set():
|
||||
if is_sm100_supported():
|
||||
self.attention_backend = "trtllm_mha"
|
||||
elif is_sm90_supported():
|
||||
self.attention_backend = "fa3"
|
||||
elif is_cpu() and cpu_has_amx_support():
|
||||
self.attention_backend = "intel_amx"
|
||||
elif is_xpu():
|
||||
self.attention_backend = "intel_xpu"
|
||||
elif is_hip():
|
||||
self.attention_backend = "aiter"
|
||||
else:
|
||||
self.attention_backend = "triton"
|
||||
|
||||
# XPU dtype validation moved to the override registry
|
||||
# (arg_groups/overrides.py: _gpt_oss_overrides).
|
||||
# Attention backend selection + XPU dtype validation moved to the
|
||||
# override registry (arg_groups/overrides.py: _gpt_oss_overrides).
|
||||
|
||||
supported_backends = [
|
||||
"triton",
|
||||
@@ -4221,35 +4206,13 @@ class ServerArgs:
|
||||
"Step3p5ForCausalLM" in model_arch
|
||||
or "Step3p7ForConditionalGeneration" in model_arch
|
||||
):
|
||||
if self.is_attention_backend_not_set():
|
||||
if is_blackwell_supported():
|
||||
self.attention_backend = "fa4"
|
||||
logger.info(
|
||||
"Auto-select fa4 attention backend for Step3p7 on Blackwell."
|
||||
)
|
||||
elif is_sm90_supported():
|
||||
self.attention_backend = "fa3"
|
||||
logger.info(
|
||||
"Auto-select fa3 attention backend for Step3p7 on Hopper."
|
||||
)
|
||||
# EAGLE multi-layer + hierarchical-cache SWA writes moved to the
|
||||
# override registry (arg_groups/overrides.py: _step3p_overrides).
|
||||
# Attention backend selection + EAGLE multi-layer +
|
||||
# hierarchical-cache SWA writes moved to the override registry
|
||||
# (arg_groups/overrides.py: _step3p_overrides).
|
||||
pass
|
||||
elif model_arch in LLAMA4_MODEL_ARCHS and self.device != "cpu":
|
||||
# Auto-select attention backend for Llama4 if not specified
|
||||
if self.attention_backend is None:
|
||||
if is_sm100_supported():
|
||||
self.attention_backend, platform = "trtllm_mha", "sm100"
|
||||
elif is_sm90_supported():
|
||||
self.attention_backend, platform = "fa3", "sm90"
|
||||
elif is_hip():
|
||||
self.attention_backend, platform = "aiter", "hip"
|
||||
elif self.device == "xpu":
|
||||
self.attention_backend, platform = "intel_xpu", "xpu"
|
||||
else:
|
||||
self.attention_backend, platform = "triton", "other platforms"
|
||||
logger.warning(
|
||||
f"Use {self.attention_backend} as attention backend on {platform} for Llama4 model"
|
||||
)
|
||||
# Attention backend auto-select moved to the override registry
|
||||
# (arg_groups/overrides.py: _llama4_overrides).
|
||||
assert self.attention_backend in {
|
||||
"fa3",
|
||||
"aiter",
|
||||
@@ -4271,21 +4234,8 @@ class ServerArgs:
|
||||
"Gemma4ForCausalLM",
|
||||
"Gemma4UnifiedForConditionalGeneration",
|
||||
):
|
||||
default_attention_backend = (
|
||||
"trtllm_mha" if is_sm100_supported() else "triton"
|
||||
)
|
||||
if self.is_attention_backend_not_set():
|
||||
self.attention_backend = default_attention_backend
|
||||
logger.info(
|
||||
f"Use {self.attention_backend} as default attention backend for Gemma4"
|
||||
)
|
||||
else:
|
||||
# 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.
|
||||
if self.attention_backend is None:
|
||||
self.attention_backend = default_attention_backend
|
||||
|
||||
# Default attention backend selection moved to the override registry
|
||||
# (arg_groups/overrides.py: _gemma4_overrides).
|
||||
prefill_backend, decode_backend = self.get_attention_backends()
|
||||
accepted_backends = ("trtllm_mha", "triton", "ascend", "intel_xpu")
|
||||
assert (
|
||||
@@ -4325,15 +4275,8 @@ class ServerArgs:
|
||||
self.attention_backend in accepted_backends
|
||||
), f"One of the attention backends in {accepted_backends} is required for {model_arch}, but got {self.attention_backend}"
|
||||
elif model_arch in ["Olmo2ForCausalLM"]:
|
||||
# disable_hybrid_swa_memory moved to the override registry
|
||||
# (arg_groups/overrides.py: _olmo2_overrides).
|
||||
if self.attention_backend is None:
|
||||
if is_cuda() and is_sm100_supported():
|
||||
self.attention_backend = "trtllm_mha"
|
||||
elif is_cuda() and get_device_sm() >= 80:
|
||||
self.attention_backend = "fa3"
|
||||
else:
|
||||
self.attention_backend = "triton"
|
||||
# disable_hybrid_swa_memory + attention backend selection moved to
|
||||
# the override registry (arg_groups/overrides.py: _olmo2_overrides).
|
||||
|
||||
# Flashinfer appears to degrade performance when sliding window attention
|
||||
# is used for the Olmo2 architecture. Olmo2 does not use sliding window attention
|
||||
@@ -4420,8 +4363,8 @@ class ServerArgs:
|
||||
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.
|
||||
if is_sm100_supported() and self.attention_backend is None:
|
||||
self.attention_backend = "triton"
|
||||
# (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)
|
||||
|
||||
elif model_arch in ["Glm4MoeForCausalLM"]:
|
||||
@@ -4447,18 +4390,16 @@ class ServerArgs:
|
||||
logger.info(
|
||||
"Use flashinfer_trtllm as MoE runner backend on sm100 for Glm4MoeForCausalLM"
|
||||
)
|
||||
self.enable_tf32_matmul = True
|
||||
logger.info(
|
||||
"Enable TF32 matmul for Glm4MoeForCausalLM model to improve gate gemm performance."
|
||||
)
|
||||
# enable_tf32_matmul moved to the override registry
|
||||
# (arg_groups/overrides.py: _glm4_moe_overrides).
|
||||
|
||||
elif model_arch in [
|
||||
"FalconH1ForCausalLM",
|
||||
"JetNemotronForCausalLM",
|
||||
"JetVLMForConditionalGeneration",
|
||||
]:
|
||||
if is_sm100_supported() and self.attention_backend is None:
|
||||
self.attention_backend = "triton"
|
||||
# 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":
|
||||
@@ -4468,13 +4409,13 @@ class ServerArgs:
|
||||
for layer_type in getattr(hf_config, "layer_types", [])
|
||||
)
|
||||
if has_mamba:
|
||||
if is_sm100_supported() and self.attention_backend is None:
|
||||
self.attention_backend = "flashinfer"
|
||||
# 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"]:
|
||||
if is_sm100_supported() and self.attention_backend is None:
|
||||
self.attention_backend = "flashinfer"
|
||||
# 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, "
|
||||
@@ -4690,22 +4631,20 @@ class ServerArgs:
|
||||
|
||||
def _handle_attention_backend_compatibility(self):
|
||||
model_config = self.get_model_config()
|
||||
use_mla_backend = self.use_mla_backend()
|
||||
|
||||
if self.prefill_attention_backend is not None and (
|
||||
self.prefill_attention_backend == self.decode_attention_backend
|
||||
): # override the default attention backend
|
||||
self.attention_backend = self.prefill_attention_backend
|
||||
# The attention_backend write clusters of this handler moved to the
|
||||
# resolution pipeline (arg_groups/overrides.py), each invoked below at
|
||||
# its legacy slot; the interleaved non-attention adjustments stay.
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
_attention_backend_default,
|
||||
_attention_backend_dual_chunk,
|
||||
_attention_backend_fa3_fp8_fallback,
|
||||
_attention_backend_platform_fallbacks,
|
||||
run_post_process_pass,
|
||||
)
|
||||
|
||||
# Pick the default attention backend if not specified
|
||||
if self.attention_backend is None:
|
||||
self.attention_backend = self._get_default_attn_backend(
|
||||
use_mla_backend, model_config
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"Attention backend not specified. Use {self.attention_backend} backend by default."
|
||||
)
|
||||
# Split-backend override + default fill.
|
||||
run_post_process_pass(self, _attention_backend_default)
|
||||
|
||||
# Torch native and flex attention backends
|
||||
if self.attention_backend == "torch_native":
|
||||
@@ -4858,12 +4797,7 @@ class ServerArgs:
|
||||
)
|
||||
self.page_size = 64
|
||||
|
||||
if self.attention_backend == "fa3" and self.kv_cache_dtype == "fp8_e5m2":
|
||||
logger.warning(
|
||||
"FlashAttention3 only supports fp8_e4m3 if using FP8; "
|
||||
"Setting attention backend to triton."
|
||||
)
|
||||
self.attention_backend = "triton"
|
||||
run_post_process_pass(self, _attention_backend_fa3_fp8_fallback)
|
||||
|
||||
if (
|
||||
(
|
||||
@@ -4889,25 +4823,7 @@ class ServerArgs:
|
||||
self.mem_fraction_static *= 0.85
|
||||
|
||||
# Other platforms backends
|
||||
if (
|
||||
self.attention_backend == "intel_amx"
|
||||
and self.device == "cpu"
|
||||
and not cpu_has_amx_support()
|
||||
):
|
||||
logger.warning(
|
||||
"The current platform does not support Intel AMX, will fallback to torch_native backend."
|
||||
)
|
||||
self.attention_backend = "torch_native"
|
||||
|
||||
if (
|
||||
self.attention_backend == "intel_xpu"
|
||||
and self.device == "xpu"
|
||||
and not xpu_has_xmx_support()
|
||||
):
|
||||
logger.warning(
|
||||
"The current platform does not support Intel XMX, will fallback to triton backend."
|
||||
)
|
||||
self.attention_backend = "triton"
|
||||
run_post_process_pass(self, _attention_backend_platform_fallbacks)
|
||||
|
||||
prefill_backend, decode_backend = self.get_attention_backends()
|
||||
if self.use_mla_backend() and prefill_backend == "intel_xpu":
|
||||
@@ -4930,18 +4846,7 @@ class ServerArgs:
|
||||
self.page_size = 128
|
||||
|
||||
# Dual chunk flash attention backend
|
||||
if (
|
||||
getattr(model_config.hf_config, "dual_chunk_attention_config", None)
|
||||
is not None
|
||||
):
|
||||
if self.attention_backend is None:
|
||||
self.attention_backend = "dual_chunk_flash_attn"
|
||||
logger.info("Dual chunk attention is turned on by default.")
|
||||
elif self.attention_backend != "dual_chunk_flash_attn":
|
||||
raise ValueError(
|
||||
"Dual chunk attention is enabled, but attention backend is set to "
|
||||
f"{self.attention_backend}. Please set it to 'dual_chunk_flash_attn'."
|
||||
)
|
||||
run_post_process_pass(self, _attention_backend_dual_chunk)
|
||||
if self.attention_backend == "dual_chunk_flash_attn":
|
||||
logger.warning(
|
||||
"Mixed chunk and radix cache are disabled when using dual-chunk flash attention backend"
|
||||
@@ -6312,10 +6217,11 @@ class ServerArgs:
|
||||
)
|
||||
self.flashinfer_allreduce_fusion_backend = None
|
||||
|
||||
# The forced-pytorch sampling write moved to the resolution
|
||||
# pipeline (arg_groups/overrides.py:
|
||||
# _deterministic_sampling_backend), invoked at its legacy slot.
|
||||
# The forced-pytorch sampling write and the attention backend
|
||||
# fill/validation moved to the resolution pipeline
|
||||
# (arg_groups/overrides.py), invoked at their legacy slots.
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
_deterministic_attention_backend,
|
||||
_deterministic_sampling_backend,
|
||||
run_post_process_pass,
|
||||
)
|
||||
@@ -6338,29 +6244,7 @@ class ServerArgs:
|
||||
pass
|
||||
|
||||
# Check attention backend
|
||||
if self.attention_backend is None:
|
||||
# User didn't specify attention backend, fallback based on GPU architecture
|
||||
if is_sm100_supported() or is_sm120_supported():
|
||||
# Blackwell and newer architectures
|
||||
if is_deepseek_model:
|
||||
# fallback to triton for DeepSeek models because flashinfer doesn't support deterministic inference for DeepSeek models yet
|
||||
self.attention_backend = "triton"
|
||||
else:
|
||||
# fallback to flashinfer on Blackwell for non-DeepSeek models
|
||||
self.attention_backend = "flashinfer"
|
||||
else:
|
||||
# Hopper (SM90) and older architectures
|
||||
self.attention_backend = "fa3"
|
||||
logger.warning(
|
||||
f"Attention backend not specified. Falling back to '{self.attention_backend}' for deterministic inference. "
|
||||
f"You can explicitly set --attention-backend to one of {DETERMINISTIC_ATTENTION_BACKEND_CHOICES}."
|
||||
)
|
||||
elif self.attention_backend not in DETERMINISTIC_ATTENTION_BACKEND_CHOICES:
|
||||
# User explicitly specified an incompatible attention backend
|
||||
raise ValueError(
|
||||
f"Currently only {DETERMINISTIC_ATTENTION_BACKEND_CHOICES} attention backends are supported for deterministic inference, "
|
||||
f"but you explicitly specified '{self.attention_backend}'."
|
||||
)
|
||||
run_post_process_pass(self, _deterministic_attention_backend)
|
||||
|
||||
if is_deepseek_model:
|
||||
if self.attention_backend not in ["fa3", "triton"]:
|
||||
@@ -6471,7 +6355,9 @@ class ServerArgs:
|
||||
def _handle_dllm_inference(self):
|
||||
if self.dllm_algorithm is None:
|
||||
return
|
||||
# On AMD/HIP, disable cuda graph for DLLM and use triton backend
|
||||
# On AMD/HIP, disable cuda graph for DLLM (the attention_backend
|
||||
# resolution moved to the pipeline: arg_groups/overrides.py
|
||||
# _dllm_attention_backend, invoked below at its legacy slot).
|
||||
if is_hip():
|
||||
if (
|
||||
self.cuda_graph_config.decode.backend != Backend.DISABLED
|
||||
@@ -6482,23 +6368,14 @@ class ServerArgs:
|
||||
)
|
||||
self.cuda_graph_config.decode.backend = Backend.DISABLED
|
||||
self.cuda_graph_config.prefill.backend = Backend.DISABLED
|
||||
if self.attention_backend not in ["triton", "aiter"]:
|
||||
logger.warning(
|
||||
"Attention backend is set to triton for diffusion LLM inference on AMD GPUs"
|
||||
)
|
||||
self.attention_backend = "triton"
|
||||
elif is_npu():
|
||||
if self.attention_backend != "ascend":
|
||||
logger.warning(
|
||||
"Attention backend is overridden to 'ascend' when running on NPU for diffusion LLM inference."
|
||||
)
|
||||
self.attention_backend = "ascend"
|
||||
elif self.cuda_graph_config.decode.backend != Backend.DISABLED:
|
||||
if self.attention_backend != "flashinfer":
|
||||
logger.warning(
|
||||
"Attention backend is set to flashinfer because of enabling cuda graph in diffusion LLM inference"
|
||||
)
|
||||
self.attention_backend = "flashinfer"
|
||||
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
_dllm_attention_backend,
|
||||
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"
|
||||
|
||||
Reference in New Issue
Block a user