94057c3d3e
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
176 lines
5.7 KiB
Python
176 lines
5.7 KiB
Python
"""CuTe DSL kernels for GDN (Gated Delta Network) linear attention.
|
|
|
|
Decode path uses the existing ``cutedsl_fused_sigmoid_gating_delta_rule_update``
|
|
(works on SM90+).
|
|
|
|
Prefill (extend) path uses the ported vLLM SM100 chunkwise kernel
|
|
(``chunk_gated_delta_rule_cutedsl``). Requires SM100+ and ``head_k_dim == 128``.
|
|
"""
|
|
|
|
import logging
|
|
from typing import Optional
|
|
|
|
import torch
|
|
|
|
from sglang.jit_kernel.cutedsl_gdn import cutedsl_fused_sigmoid_gating_delta_rule_update
|
|
from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
|
|
LinearAttnKernelBase,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _is_blackwell() -> bool:
|
|
"""True iff running on SM100+ (Blackwell) where the ported kernel is valid."""
|
|
if not torch.cuda.is_available():
|
|
return False
|
|
major, _ = torch.cuda.get_device_capability()
|
|
return major >= 10
|
|
|
|
|
|
class CuteDSLGDNKernel(LinearAttnKernelBase):
|
|
"""CuTe DSL kernel for GDN.
|
|
|
|
Decode: ``cutedsl_fused_sigmoid_gating_delta_rule_update`` (SM90+).
|
|
Extend (prefill): chunkwise ``chunk_gated_delta_rule_cutedsl``
|
|
(SM100+ only, ``head_k_dim`` must be 128). On SM90 the prefill path is
|
|
unsupported; callers should query :attr:`supports_prefill` and fall back
|
|
to another backend (e.g. Triton).
|
|
"""
|
|
|
|
def __init__(self):
|
|
# The Blackwell extend kernel uses tcgen05/TMA-bulk-swizzle features
|
|
# that don't exist on SM90. The decode kernel does work on SM90+.
|
|
self.supports_prefill = _is_blackwell()
|
|
|
|
# Heavy CuteDSL imports are deferred to extend() so SM90 boxes can
|
|
# still construct the kernel just for decode.
|
|
self._extend_fn: Optional[callable] = None
|
|
self._prepare_meta_fn: Optional[callable] = None
|
|
self._l2norm_fn: Optional[callable] = None
|
|
|
|
def _ensure_extend_loaded(self, head_k_dim: int) -> None:
|
|
if self._extend_fn is not None:
|
|
return
|
|
if not self.supports_prefill:
|
|
major = (
|
|
torch.cuda.get_device_capability()[0]
|
|
if torch.cuda.is_available()
|
|
else -1
|
|
)
|
|
raise RuntimeError(
|
|
f"CuTe DSL GDN prefill requires SM100+ (Blackwell); got SM{major}."
|
|
)
|
|
if head_k_dim != 128:
|
|
raise RuntimeError(
|
|
f"CuTe DSL GDN prefill requires head_k_dim=128, got {head_k_dim}."
|
|
)
|
|
from sglang.srt.layers.attention.fla.l2norm import l2norm_fwd
|
|
from sglang.srt.layers.attention.linear.kernels.gdn_blackwell import (
|
|
chunk_gated_delta_rule_cutedsl,
|
|
prepare_metadata_cutedsl,
|
|
)
|
|
|
|
self._extend_fn = chunk_gated_delta_rule_cutedsl
|
|
self._prepare_meta_fn = prepare_metadata_cutedsl
|
|
self._l2norm_fn = l2norm_fwd
|
|
logger.info("Using CuTe DSL GDN prefill (Blackwell)")
|
|
|
|
def decode(
|
|
self,
|
|
q: torch.Tensor,
|
|
k: torch.Tensor,
|
|
v: torch.Tensor,
|
|
a: torch.Tensor,
|
|
b: torch.Tensor,
|
|
*,
|
|
A_log: torch.Tensor,
|
|
dt_bias: torch.Tensor,
|
|
ssm_states: torch.Tensor,
|
|
cache_indices: torch.Tensor,
|
|
query_start_loc: torch.Tensor,
|
|
**kwargs,
|
|
) -> torch.Tensor:
|
|
return cutedsl_fused_sigmoid_gating_delta_rule_update(
|
|
A_log=A_log,
|
|
dt_bias=dt_bias,
|
|
q=q,
|
|
k=k,
|
|
v=v,
|
|
a=a,
|
|
b=b,
|
|
initial_state_source=ssm_states,
|
|
initial_state_indices=cache_indices,
|
|
cu_seqlens=query_start_loc,
|
|
use_qk_l2norm_in_kernel=True,
|
|
softplus_beta=1.0,
|
|
softplus_threshold=20.0,
|
|
)
|
|
|
|
def extend(
|
|
self,
|
|
q: torch.Tensor,
|
|
k: torch.Tensor,
|
|
v: torch.Tensor,
|
|
g: torch.Tensor,
|
|
beta: torch.Tensor,
|
|
*,
|
|
ssm_states: torch.Tensor,
|
|
cache_indices: torch.Tensor,
|
|
query_start_loc: torch.Tensor,
|
|
**kwargs,
|
|
) -> tuple:
|
|
head_k_dim = k.shape[-1]
|
|
self._ensure_extend_loaded(head_k_dim)
|
|
|
|
total_seq_len = q.shape[1]
|
|
num_v_heads = v.shape[2]
|
|
head_v_dim = v.shape[3]
|
|
|
|
# L2 norm Q/K outside the kernel (same as flashinfer path).
|
|
q_norm = self._l2norm_fn(q[0].contiguous()).unsqueeze(0)
|
|
k_norm = self._l2norm_fn(k[0].contiguous()).unsqueeze(0)
|
|
v_in = v[0].contiguous().unsqueeze(0)
|
|
# Kernel expects log-space float32 gate per (token, v-head).
|
|
g_in = g[0].to(torch.float32).unsqueeze(0)
|
|
beta_in = beta[0].to(torch.float32).unsqueeze(0)
|
|
|
|
cu_seqlens = query_start_loc.to(torch.int32)
|
|
|
|
# Pool gather: remap padding (-1) to the last (sentinel) slot.
|
|
ssm_cache_indices = torch.where(
|
|
cache_indices >= 0,
|
|
cache_indices,
|
|
ssm_states.shape[0] - 1,
|
|
).to(torch.long)
|
|
initial_state = ssm_states[ssm_cache_indices].contiguous()
|
|
|
|
chunk_indices, chunk_offsets = self._prepare_meta_fn(
|
|
cu_seqlens, total_seq_len, chunk_size=64
|
|
)
|
|
|
|
output, final_state = self._extend_fn(
|
|
q=q_norm,
|
|
k=k_norm,
|
|
v=v_in,
|
|
g=g_in,
|
|
beta=beta_in,
|
|
initial_state=initial_state,
|
|
cu_seqlens=cu_seqlens,
|
|
chunk_indices=chunk_indices,
|
|
chunk_offsets=chunk_offsets,
|
|
)
|
|
|
|
ssm_states.index_copy_(
|
|
0,
|
|
ssm_cache_indices,
|
|
final_state.to(ssm_states.dtype),
|
|
)
|
|
|
|
# Match Triton extend interface: (output, last_recurrent_state, h).
|
|
# We've already written state back, so no need to return it.
|
|
return output, None, None
|
|
|
|
def target_verify(self, *args, **kwargs):
|
|
raise NotImplementedError("CuteDSLGDNKernel does not support target_verify")
|