[XPU] Use a fused GDN kernel from sgl-kernel for Qwen3.5 (#33354)

This commit is contained in:
Xia Weiwen
2026-08-24 12:07:50 +08:00
committed by GitHub
parent 8dcfb3b5e7
commit 5b5b29d4e2
9 changed files with 307 additions and 3 deletions
@@ -0,0 +1,138 @@
import torch
from sgl_kernel import gdn_attention as sgl_kernel_gdn_attention
from sglang.srt.layers.attention.linear.gdn_backend import GDNAttnBackend
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
class XpuGDNAttnBackend(GDNAttnBackend):
"""XPU specialization of ``GDNAttnBackend``.
Adds an optional fused path that dispatches the whole conv1d + gating +
delta-rule pipeline to the vendored vLLM SYCL kernel exposed as
``torch.ops.sgl_kernel.gdn_attention``. This is opt-in via
``--linear-attn-backend intel_xpu`` (the default remains ``triton``, same
as other platforms).
"""
def supports_fused_gdn(self, layer, forward_batch: ForwardBatch) -> bool:
"""Conservative guard: only the plain decode / non-prefix-cached,
non-speculative extend cases are handled by the fused kernel."""
mode = forward_batch.forward_mode
backends = self.linear_attn_backends
selected = (
backends.verify
if mode.is_target_verify()
else (backends.decode if mode.is_decode_or_idle() else backends.prefill)
)
if not selected.is_intel_xpu():
return False
if not hasattr(torch.ops.sgl_kernel, "gdn_attention"):
# User explicitly asked for intel_xpu but the op isn't built.
raise RuntimeError(
"--linear-attn-backend intel_xpu requires the "
"torch.ops.sgl_kernel.gdn_attention op, but it is not "
"available. Rebuild sgl-kernel-xpu or use "
"--linear-attn-backend triton."
)
if mode.is_target_verify() or mode.is_draft_extend_v2():
return False
fm = self.forward_metadata
if getattr(fm, "has_mamba_track_mask", False):
# chunked prefix-cache intermediate-state tracking unsupported
return False
if getattr(fm, "query_start_loc", None) is None:
return False
# GDN (not KDA) shared weights must be plain tensors
if not isinstance(layer.conv_weights, torch.Tensor):
return False
if layer.bias is not None and not isinstance(layer.bias, torch.Tensor):
return False
return True
def forward_fused_gdn(
self,
layer: RadixLinearAttention,
forward_batch: ForwardBatch,
projected_states_qkvz: torch.Tensor,
projected_states_ba: torch.Tensor,
):
"""Run the fused SYCL GDN op and return ``(core_attn_out, z)``.
Caches stay in the SGLang pool layout and are updated in place. The conv
pool is ``[cache, dim, width-1]``; we pass a transposed view so the op sees
its logical ``[cache, width-1, dim]`` layout while the kernels index via
explicit width/dim strides (no gather/transpose/scatter copies). The ssm
pool already matches the op layout. ``mamba_cache_indices`` indexes the
full pool directly for both conv and ssm.
"""
fm = self.forward_metadata
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id)
conv_states = layer_cache.conv[0] # [cache, dim, width-1]
ssm_states = layer_cache.temporal # [cache, nv, hv, hk]
cache_indices = fm.mamba_cache_indices
query_start_loc = fm.query_start_loc
device = projected_states_qkvz.device
dtype = projected_states_qkvz.dtype
bs = forward_batch.batch_size
num_actual_tokens = projected_states_qkvz.shape[0]
if forward_batch.forward_mode.is_decode_or_idle():
num_decodes, num_prefills = bs, 0
has_initial_state = torch.ones(bs, dtype=torch.bool, device=device)
else:
num_decodes, num_prefills = 0, bs
has_initial_state = forward_batch.extend_prefix_lens > 0
# Full-pool, zero-copy: transposed view for conv + native ssm pool, indexed
# directly by the full-pool cache indices.
conv_view = conv_states.transpose(1, 2) # [cache, width-1, dim] view
state_idx = cache_indices.to(torch.int32).contiguous()
core_attn_out = torch.empty(
num_actual_tokens,
layer.num_v_heads,
layer.head_v_dim,
dtype=dtype,
device=device,
)
z = torch.empty_like(core_attn_out)
sgl_kernel_gdn_attention(
core_attn_out=core_attn_out,
z=z,
projected_states_qkvz=projected_states_qkvz,
projected_states_ba=projected_states_ba,
num_k_heads=layer.num_k_heads,
num_v_heads=layer.num_v_heads,
head_k_dim=layer.head_k_dim,
head_v_dim=layer.head_v_dim,
conv_state=conv_view,
ssm_state=ssm_states,
conv_weights=layer.conv_weights,
conv_bias=layer.bias,
activation=layer.activation,
A_log=layer.A_log,
dt_bias=layer.dt_bias,
num_prefills=num_prefills,
num_decodes=num_decodes,
num_spec_decodes=0,
has_initial_state=has_initial_state,
non_spec_query_start_loc=query_start_loc,
non_spec_token_indx=None,
non_spec_state_indices_tensor=state_idx,
spec_query_start_loc=None,
spec_token_indx=None,
spec_state_indices_tensor=None,
num_accepted_tokens=None,
num_actual_tokens=num_actual_tokens,
# Heads/tensors are already per-rank sharded; kernel needs no
# further in-kernel sharding, so this is always 1, not --tp-size.
tp_size=1,
reorder_input=True,
)
# conv/ssm states were updated in place via the pool views.
return core_attn_out, z
@@ -378,6 +378,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
is_blackwell, is_blackwell,
is_npu, is_npu,
is_sm120_supported, is_sm120_supported,
is_xpu,
) )
if not is_npu(): if not is_npu():
@@ -389,6 +390,11 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
GDNAttnBackend, GDNAttnBackend,
flashinfer_gdn_prefill_default, flashinfer_gdn_prefill_default,
) )
if is_xpu():
from sglang.srt.hardware_backend.xpu.attention.xpu_gdn_backend import (
XpuGDNAttnBackend as GDNAttnBackend,
)
else: else:
from sglang.srt.hardware_backend.npu.attention.ascend_gdn_backend import ( from sglang.srt.hardware_backend.npu.attention.ascend_gdn_backend import (
AscendGDNAttnBackend as GDNAttnBackend, AscendGDNAttnBackend as GDNAttnBackend,
@@ -142,6 +142,13 @@ class GDNKernelDispatcher:
cutedsl_kernel = None cutedsl_kernel = None
if decode_backend.is_triton(): if decode_backend.is_triton():
self.decode_kernel = triton_kernel self.decode_kernel = triton_kernel
elif decode_backend.is_intel_xpu():
if not is_xpu():
raise ValueError("--linear-attn-backend intel_xpu requires Intel XPU")
# The fused SYCL kernel is dispatched via XpuGDNAttnBackend.forward_fused_gdn,
# outside this dispatcher; Triton is the dispatcher-level kernel for requests
# that hook doesn't handle (e.g. verify).
self.decode_kernel = triton_kernel
elif decode_backend.is_cutedsl(): elif decode_backend.is_cutedsl():
if not is_cuda(): if not is_cuda():
raise ValueError("GDN CuTe DSL backend requires CUDA") raise ValueError("GDN CuTe DSL backend requires CUDA")
@@ -169,6 +176,12 @@ class GDNKernelDispatcher:
if prefill_backend.is_triton(): if prefill_backend.is_triton():
self.extend_kernel = triton_kernel self.extend_kernel = triton_kernel
elif prefill_backend.is_intel_xpu():
if not is_xpu():
raise ValueError("--linear-attn-backend intel_xpu requires Intel XPU")
# See the decode branch above: intel_xpu uses Triton as its
# dispatcher-level fallback kernel.
self.extend_kernel = triton_kernel
elif prefill_backend.is_cutedsl(): elif prefill_backend.is_cutedsl():
if not is_cuda(): if not is_cuda():
raise ValueError("GDN CuTe DSL backend requires CUDA") raise ValueError("GDN CuTe DSL backend requires CUDA")
@@ -384,6 +397,7 @@ class GDNAttnBackend(MambaAttnBackendBase):
), f"{self.conv_states_shape[-1]=} should be less than {FLA_CHUNK_SIZE}" ), f"{self.conv_states_shape[-1]=} should be less than {FLA_CHUNK_SIZE}"
backends = model_runner.linear_attn_backends backends = model_runner.linear_attn_backends
self.linear_attn_backends = backends
self.kernel_dispatcher = GDNKernelDispatcher( self.kernel_dispatcher = GDNKernelDispatcher(
backends.decode, backends.prefill, backends.verify backends.decode, backends.prefill, backends.verify
) )
@@ -21,6 +21,7 @@ class LinearAttnKernelBackend(Enum):
NVIDIA_KDA = "nvidia_kda" NVIDIA_KDA = "nvidia_kda"
PTX_KDA = "ptx_kda" PTX_KDA = "ptx_kda"
HELION = "helion" HELION = "helion"
INTEL_XPU = "intel_xpu"
CUSTOM = "custom" CUSTOM = "custom"
@classmethod @classmethod
@@ -51,6 +52,9 @@ class LinearAttnKernelBackend(Enum):
def is_helion(self): def is_helion(self):
return self == LinearAttnKernelBackend.HELION return self == LinearAttnKernelBackend.HELION
def is_intel_xpu(self):
return self == LinearAttnKernelBackend.INTEL_XPU
def is_custom(self): def is_custom(self):
return self == LinearAttnKernelBackend.CUSTOM return self == LinearAttnKernelBackend.CUSTOM
+45
View File
@@ -136,6 +136,7 @@ _qknorm_use_alt_stream = _is_cuda or (
get_bool_env_var("SGLANG_QK_NORM_ALT_STREAM", "False") and _hip_use_alt_stream get_bool_env_var("SGLANG_QK_NORM_ALT_STREAM", "False") and _hip_use_alt_stream
) )
_is_amx_available = cpu_has_amx_support() _is_amx_available = cpu_has_amx_support()
_is_xpu = is_xpu()
# Head-group ratios (num_v_heads // num_k_heads) served by the fused # Head-group ratios (num_v_heads // num_k_heads) served by the fused
# split/reshape/cat Triton kernel. On AMD/aiter the ratio-8 layout is also # split/reshape/cat Triton kernel. On AMD/aiter the ratio-8 layout is also
@@ -670,6 +671,40 @@ class Qwen3_5GatedDeltaNet(nn.Module):
projected_states_ba, _ = self.in_proj_ba(hs_bf16) projected_states_ba, _ = self.in_proj_ba(hs_bf16)
return projected_states_qkvz, projected_states_ba return projected_states_qkvz, projected_states_ba
def _forward_xpu(
self,
backend: object,
projected_states_qkvz: torch.Tensor,
projected_states_ba: torch.Tensor,
forward_batch: ForwardBatch,
):
core_attn_out, z = backend.forward_fused_gdn(
self.attn,
forward_batch,
projected_states_qkvz,
projected_states_ba,
)
assert core_attn_out is not None, "XPU backend must support fused GDN"
z_shape_og = z.shape
# reshape input data into 2D tensor
core_attn_out = core_attn_out.reshape(-1, core_attn_out.shape[-1])
z = z.reshape(-1, z.shape[-1])
# Add padding for DP-Attn
if core_attn_out.shape != z.shape:
core_attn_out_pad = torch.zeros_like(z)
core_attn_out_pad[: core_attn_out.shape[0], :] = core_attn_out
core_attn_out = core_attn_out_pad
core_attn_out = self.norm(core_attn_out, z)
core_attn_out = core_attn_out.reshape(z_shape_og)
core_attn_out = core_attn_out.reshape(*core_attn_out.shape[:-2], -1)
output, _ = self.out_proj(core_attn_out)
return output
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -685,6 +720,16 @@ class Qwen3_5GatedDeltaNet(nn.Module):
hidden_states hidden_states
) )
if _is_xpu and get_exec().mamba.linear_attn_backend == "intel_xpu":
from sglang.srt.model_executor.forward_context import get_attn_backend
backend = get_attn_backend()
backend = getattr(backend, "linear_attn_backend", backend)
if backend.supports_fused_gdn(self.attn, forward_batch):
return self._forward_xpu(
backend, projected_states_qkvz, projected_states_ba, forward_batch
)
if ( if (
self.num_v_heads // self.num_k_heads in _GDN_FUSED_QKVZBA_RATIOS self.num_v_heads // self.num_k_heads in _GDN_FUSED_QKVZBA_RATIOS
and not _is_npu and not _is_npu
+1
View File
@@ -404,6 +404,7 @@ LINEAR_ATTN_KERNEL_BACKEND_CHOICES = [
"nvidia_kda", "nvidia_kda",
"ptx_kda", "ptx_kda",
"helion", "helion",
"intel_xpu",
] ]
@@ -1,6 +1,8 @@
"""Qwen3.5-9B GSM8K accuracy on Intel XPU (TP=4). """Qwen3.5-9B GSM8K accuracy on Intel XPU (TP=4).
Scored by ``simple_eval_gsm8k.GSM8KEval``. Scored by ``simple_eval_gsm8k.GSM8KEval``. Covers both the opt-in fused GDN
SYCL kernel path (``--linear-attn-backend intel_xpu``) and the default Triton
GDN path (``triton``, unchanged from other platforms).
""" """
import unittest import unittest
@@ -18,9 +20,9 @@ register_xpu_ci(est_time=2400, suite="nightly-xpu-4-gpu", nightly=True)
torch.xpu.is_available(), torch.xpu.is_available(),
"Intel XPU not available (torch.xpu.is_available() returned False)", "Intel XPU not available (torch.xpu.is_available() returned False)",
) )
class TestQwen3_5_9BXPU(SimpleEvalGSM8KXPUMixin, CustomTestCase): class Qwen3_5_9BXPUBase(SimpleEvalGSM8KXPUMixin, CustomTestCase):
model = "Qwen/Qwen3.5-9B" model = "Qwen/Qwen3.5-9B"
tp_size = 4 tp_size = 1
accuracy = 0.90 accuracy = 0.90
# max_tokens=8192 lets the GSM8K CoT complete under num_threads=4. # max_tokens=8192 lets the GSM8K CoT complete under num_threads=4.
num_examples = 50 num_examples = 50
@@ -37,5 +39,23 @@ class TestQwen3_5_9BXPU(SimpleEvalGSM8KXPUMixin, CustomTestCase):
] ]
class TestQwen3_5_9BXPUDefault(Qwen3_5_9BXPUBase):
"""Default path: Triton GDN kernels (unchanged from other platforms)."""
class TestQwen3_5_9BXPUFusedGDN(Qwen3_5_9BXPUBase):
"""Opt-in fused SYCL GDN kernel path (``--linear-attn-backend intel_xpu``).
Small ``num_examples`` since this is a smoke check of the fused-kernel
dispatch, not a full accuracy regression test (already covered by the
default Triton path above); accuracy threshold is left at 0 accordingly.
"""
other_args = Qwen3_5_9BXPUBase.other_args + [
"--linear-attn-backend",
"intel_xpu",
]
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -0,0 +1,76 @@
"""Intel XPU linear-attn (GDN) backend dispatch: intel_xpu is opt-in only
(default stays triton, like every other platform) and fails fast rather than
silently degrading when misconfigured. Pure dispatch-logic tests -- no XPU
device required -- kept under test/registered/xpu to separate Intel-XPU-only
behavior from the platform-agnostic linear-attn dispatch tests.
"""
import unittest
from unittest.mock import patch
from sglang.srt.layers.attention.linear import gdn_backend
from sglang.srt.layers.attention.linear.gdn_backend import GDNKernelDispatcher
from sglang.srt.layers.attention.linear.kda_backend import KDAKernelDispatcher
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
from sglang.srt.layers.attention.linear.utils import LinearAttnKernelBackend
from sglang.srt.server_args import LINEAR_ATTN_KERNEL_BACKEND_CHOICES
from sglang.test.ci.ci_register import register_xpu_ci
from sglang.test.test_utils import CustomTestCase
register_xpu_ci(est_time=5, suite="stage-b-test-1-gpu-xpu")
class TestIntelXpuGDNDispatch(CustomTestCase):
def test_intel_xpu_is_a_registered_backend_choice(self):
self.assertIn("intel_xpu", LINEAR_ATTN_KERNEL_BACKEND_CHOICES)
def test_intel_xpu_requires_xpu_hardware(self):
with patch.object(gdn_backend, "is_xpu", return_value=False):
with self.assertRaisesRegex(ValueError, "requires Intel XPU"):
GDNKernelDispatcher(
LinearAttnKernelBackend.INTEL_XPU,
LinearAttnKernelBackend.TRITON,
)
with self.assertRaisesRegex(ValueError, "requires Intel XPU"):
GDNKernelDispatcher(
LinearAttnKernelBackend.TRITON,
LinearAttnKernelBackend.INTEL_XPU,
)
def test_intel_xpu_uses_triton_as_the_dispatcher_fallback_kernel(self):
# The fused SYCL kernel is dispatched outside GDNKernelDispatcher (via
# XpuGDNAttnBackend.forward_fused_gdn); the dispatcher itself only
# needs a valid fallback kernel for requests that hook declines
# (e.g. verify), which is Triton.
with patch.object(gdn_backend, "is_xpu", return_value=True):
dispatcher = GDNKernelDispatcher(
LinearAttnKernelBackend.INTEL_XPU,
LinearAttnKernelBackend.INTEL_XPU,
)
self.assertIsInstance(dispatcher.decode_kernel, TritonGDNKernel)
self.assertIsInstance(dispatcher.extend_kernel, TritonGDNKernel)
self.assertIsInstance(dispatcher.verify_kernel, TritonGDNKernel)
class TestIntelXpuKDADispatch(CustomTestCase):
def test_intel_xpu_is_not_a_supported_kda_backend(self):
# Unlike GDN, KDA has no Intel XPU SYCL kernel: intel_xpu must not be
# silently treated as Triton, it should fail fast like any other
# backend KDA does not implement.
with self.assertRaisesRegex(ValueError, "Unsupported KDA decode backend"):
KDAKernelDispatcher(
decode_backend=LinearAttnKernelBackend.INTEL_XPU,
prefill_backend=LinearAttnKernelBackend.TRITON,
verify_backend=LinearAttnKernelBackend.TRITON,
)
with self.assertRaisesRegex(ValueError, "Unsupported KDA prefill backend"):
KDAKernelDispatcher(
decode_backend=LinearAttnKernelBackend.TRITON,
prefill_backend=LinearAttnKernelBackend.INTEL_XPU,
verify_backend=LinearAttnKernelBackend.TRITON,
)
if __name__ == "__main__":
unittest.main()