From 704e51283605d76d5961d6e014330964fdb1523b Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Fri, 14 Aug 2026 08:59:52 +0800 Subject: [PATCH] [GDN] Honor configured linear-attn verify backend in the kernel dispatcher (#34592) Co-authored-by: Claude Fable 5 --- .../layers/attention/linear/gdn_backend.py | 20 ++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index 7539cc1cd..6b49fae3b 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -15,6 +15,7 @@ from sglang.srt.layers.attention.linear.utils import ( build_verify_intermediate_state_indices, get_linear_attn_decode_backend, get_linear_attn_prefill_backend, + get_linear_attn_verify_backend, ) from sglang.srt.layers.radix_linear_attention import RadixLinearAttention from sglang.srt.mem_cache.memory_pool import MambaPool @@ -110,6 +111,7 @@ class GDNKernelDispatcher: self, decode_backend: LinearAttnKernelBackend, prefill_backend: LinearAttnKernelBackend, + verify_backend: Optional[LinearAttnKernelBackend] = None, ): triton_kernel = TritonGDNKernel() self.tree_verify_kernel = triton_kernel @@ -177,10 +179,15 @@ class GDNKernelDispatcher: else: raise ValueError(f"Unsupported GDN prefill backend: {prefill_backend}") - # Verify kernel: use FlashInfer when the selected FlashInfer kernel - # supports MTP verify. SM90 uses the fp32-state path; SM100 uses the - # bf16-state adapter in FlashInferGDNKernel. - if ( + # Verify kernel. An explicitly configured verify backend wins; the + # historical auto rule (FlashInfer when the selected FlashInfer kernel + # supports MTP verify) only applies when no explicit choice was made. + # SM90 FlashInfer verify requires a fp32 SSM state, so e.g. + # --mamba-ssm-dtype bfloat16 setups must be able to force Triton here. + if verify_backend is not None and verify_backend.is_triton(): + self.verify_kernel = triton_kernel + self.verify_kernel_is_flashinfer = False + elif ( decode_backend.is_flashinfer() or prefill_backend.is_flashinfer() ) and flashinfer_kernel.supports_target_verify: self.verify_kernel = flashinfer_kernel @@ -346,7 +353,10 @@ class GDNAttnBackend(MambaAttnBackendBase): decode_backend = get_linear_attn_decode_backend() prefill_backend = get_linear_attn_prefill_backend() - self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, prefill_backend) + verify_backend = get_linear_attn_verify_backend() + self.kernel_dispatcher = GDNKernelDispatcher( + decode_backend, prefill_backend, verify_backend + ) # Sized past the pool for attn_tp-padded warmup/MLP-sync batches (see helper). self.verify_intermediate_state_indices = ( build_verify_intermediate_state_indices(