[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_npu,
is_sm120_supported,
is_xpu,
)
if not is_npu():
@@ -389,6 +390,11 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
GDNAttnBackend,
flashinfer_gdn_prefill_default,
)
if is_xpu():
from sglang.srt.hardware_backend.xpu.attention.xpu_gdn_backend import (
XpuGDNAttnBackend as GDNAttnBackend,
)
else:
from sglang.srt.hardware_backend.npu.attention.ascend_gdn_backend import (
AscendGDNAttnBackend as GDNAttnBackend,
@@ -142,6 +142,13 @@ class GDNKernelDispatcher:
cutedsl_kernel = None
if decode_backend.is_triton():
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():
if not is_cuda():
raise ValueError("GDN CuTe DSL backend requires CUDA")
@@ -169,6 +176,12 @@ class GDNKernelDispatcher:
if prefill_backend.is_triton():
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():
if not is_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}"
backends = model_runner.linear_attn_backends
self.linear_attn_backends = backends
self.kernel_dispatcher = GDNKernelDispatcher(
backends.decode, backends.prefill, backends.verify
)
@@ -21,6 +21,7 @@ class LinearAttnKernelBackend(Enum):
NVIDIA_KDA = "nvidia_kda"
PTX_KDA = "ptx_kda"
HELION = "helion"
INTEL_XPU = "intel_xpu"
CUSTOM = "custom"
@classmethod
@@ -51,6 +52,9 @@ class LinearAttnKernelBackend(Enum):
def is_helion(self):
return self == LinearAttnKernelBackend.HELION
def is_intel_xpu(self):
return self == LinearAttnKernelBackend.INTEL_XPU
def is_custom(self):
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
)
_is_amx_available = cpu_has_amx_support()
_is_xpu = is_xpu()
# 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
@@ -670,6 +671,40 @@ class Qwen3_5GatedDeltaNet(nn.Module):
projected_states_ba, _ = self.in_proj_ba(hs_bf16)
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(
self,
hidden_states: torch.Tensor,
@@ -685,6 +720,16 @@ class Qwen3_5GatedDeltaNet(nn.Module):
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 (
self.num_v_heads // self.num_k_heads in _GDN_FUSED_QKVZBA_RATIOS
and not _is_npu
+1
View File
@@ -404,6 +404,7 @@ LINEAR_ATTN_KERNEL_BACKEND_CHOICES = [
"nvidia_kda",
"ptx_kda",
"helion",
"intel_xpu",
]