[Qwen3.5] Set full attn_backend to trtllm_mha on SM100 by default when possible (#19030)

This commit is contained in:
hlu1
2026-03-02 23:14:53 +08:00
committed by GitHub
parent 2d183c4e6d
commit 468e3dc56b
+52 -44
View File
@@ -1597,25 +1597,6 @@ class ServerArgs:
elif model_arch in [ elif model_arch in [
"Qwen3MoeForCausalLM", "Qwen3MoeForCausalLM",
"Qwen3VLMoeForConditionalGeneration", "Qwen3VLMoeForConditionalGeneration",
]:
if is_sm100_supported():
quant_method = get_quantization_config(hf_config)
if self.quantization is None and quant_method is not None:
self.quantization = quant_method
if (
(
self.quantization in ("fp8", "modelopt_fp4")
or self.quantization is None
)
and self.moe_a2a_backend == "none"
and self.moe_runner_backend == "auto"
):
self.moe_runner_backend = "flashinfer_trtllm"
logger.info(
"Use flashinfer_trtllm as MoE runner backend on sm100 for "
f"{model_arch}"
)
elif model_arch in [
"Qwen3NextForCausalLM", "Qwen3NextForCausalLM",
"Qwen3_5MoeForConditionalGeneration", "Qwen3_5MoeForConditionalGeneration",
"Qwen3_5ForConditionalGeneration", "Qwen3_5ForConditionalGeneration",
@@ -1634,13 +1615,38 @@ class ServerArgs:
): ):
self.moe_runner_backend = "flashinfer_trtllm" self.moe_runner_backend = "flashinfer_trtllm"
logger.info( logger.info(
f"Use flashinfer_trtllm as MoE runner backend on sm100 for {model_arch}" "Use flashinfer_trtllm as MoE runner backend on sm100 for "
f"{model_arch}"
) )
if model_arch in [
"Qwen3NextForCausalLM",
"Qwen3_5MoeForConditionalGeneration",
"Qwen3_5ForConditionalGeneration",
]:
sm100_default_attn_backend = "triton"
if is_sm100_supported():
# 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 = self._get_default_attn_backend(
use_mla_backend=self.use_mla_backend(),
model_config=self.get_model_config(),
)
if default_attn_backend == "trtllm_mha" and not (
not self.enable_mamba_extra_buffer()
and not self.disable_radix_cache
and self.speculative_algorithm is None
):
sm100_default_attn_backend = "trtllm_mha"
self._handle_mamba_radix_cache( self._handle_mamba_radix_cache(
model_arch=model_arch, model_arch=model_arch,
support_mamba_cache=True, support_mamba_cache=True,
support_mamba_cache_extra_buffer=True, support_mamba_cache_extra_buffer=True,
sm100_default_attention_backend="triton", sm100_default_attention_backend=sm100_default_attn_backend,
) )
elif model_arch in ["Glm4MoeForCausalLM"]: elif model_arch in ["Glm4MoeForCausalLM"]:
@@ -1815,17 +1821,7 @@ class ServerArgs:
"flashinfer" if is_flashinfer_available() else "pytorch" "flashinfer" if is_flashinfer_available() else "pytorch"
) )
def _handle_attention_backend_compatibility(self): def _get_default_attn_backend(self, use_mla_backend: bool, model_config):
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
# Pick the default attention backend if not specified
if self.attention_backend is None:
""" """
Auto select the fastest attention backend. Auto select the fastest attention backend.
@@ -1839,14 +1835,13 @@ class ServerArgs:
2.2 We will use Flashinfer backend on blackwell. 2.2 We will use Flashinfer backend on blackwell.
2.3 Otherwise, we will use triton backend. 2.3 Otherwise, we will use triton backend.
""" """
if not use_mla_backend: if not use_mla_backend:
# MHA architecture # MHA architecture
if is_hopper_with_cuda_12_3() and is_no_spec_infer_or_topk_one(self): if is_hopper_with_cuda_12_3() and is_no_spec_infer_or_topk_one(self):
# Note: flashinfer 0.6.1 caused performance regression on Hopper attention kernel # 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 # Before the kernel is fixed, we choose fa3 as the default backend on Hopper MHA
# ref: https://github.com/sgl-project/sglang/issues/17411 # ref: https://github.com/sgl-project/sglang/issues/17411
self.attention_backend = "fa3" return "fa3"
elif ( elif (
is_sm100_supported() is_sm100_supported()
and is_no_spec_infer_or_topk_one(self) and is_no_spec_infer_or_topk_one(self)
@@ -1855,28 +1850,41 @@ class ServerArgs:
or self.speculative_eagle_topk is not None or self.speculative_eagle_topk is not None
) )
): ):
self.attention_backend = "trtllm_mha" return "trtllm_mha"
elif is_hip(): elif is_hip():
self.attention_backend = "aiter" return "aiter"
else: else:
self.attention_backend = ( return "flashinfer" if is_flashinfer_available() else "triton"
"flashinfer" if is_flashinfer_available() else "triton"
)
else: else:
# MLA architecture # MLA architecture
if is_hopper_with_cuda_12_3(): if is_hopper_with_cuda_12_3():
self.attention_backend = "fa3" return "fa3"
elif is_sm100_supported(): elif is_sm100_supported():
self.attention_backend = "flashinfer" return "flashinfer"
elif is_hip(): elif is_hip():
head_num = model_config.get_num_kv_heads(self.tp_size) head_num = model_config.get_num_kv_heads(self.tp_size)
# TODO current aiter only support head number 16 or 128 head number # TODO current aiter only support head number 16 or 128 head number
if head_num == 128 or head_num == 16: if head_num == 128 or head_num == 16:
self.attention_backend = "aiter" return "aiter"
else: else:
self.attention_backend = "triton" return "triton"
else: else:
self.attention_backend = "triton" return "triton"
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
# 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( logger.info(
f"Attention backend not specified. Use {self.attention_backend} backend by default." f"Attention backend not specified. Use {self.attention_backend} backend by default."