Fused QK GemmaRMSNorm + RoPE + gate kernel for Qwen3.5 (#28320)

This commit is contained in:
Yuhao Yang
2026-06-25 15:58:54 +08:00
committed by GitHub
parent 2812a3c93a
commit 4a8200565e
2 changed files with 241 additions and 1 deletions
@@ -0,0 +1,201 @@
"""Fused Q/K GemmaRMSNorm + NeoX RoPE + gate deinterleave (Triton).
Single kernel launch fusing per-head GemmaRMSNorm, partial NeoX RoPE,
and gate deinterleave for Qwen3.5's interleaved Q+Gate layout.
2D grid (T, num_q_heads + num_kv_heads) — each program handles one
(token, head) pair. Q programs also copy the gate slice.
"""
from typing import Optional, Tuple
import torch
import triton
import triton.language as tl
def _pdl_supported() -> bool:
"""Check if Programmatic Dependent Launch is supported (NVIDIA SM >= 90)."""
if not torch.cuda.is_available():
return False
try:
major, _ = torch.cuda.get_device_capability()
return major >= 9
except Exception:
return False
_ENABLE_PDL = _pdl_supported()
@triton.jit
def _fused_qk_rmsnorm_rope_gate_kernel(
q_gate_ptr,
k_ptr,
q_out_ptr,
k_out_ptr,
gate_out_ptr,
q_weight_ptr,
k_weight_ptr,
cos_sin_cache_ptr,
positions_ptr,
stride_qg_t,
stride_k_t,
stride_qo_t,
stride_ko_t,
stride_gate_t,
stride_cos_t,
NUM_Q_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
ROTARY_DIM: tl.constexpr,
HALF_ROTARY: tl.constexpr,
HEAD_BLOCK: tl.constexpr,
ROT_HALF_BLOCK: tl.constexpr,
EPS: tl.constexpr,
FP16: tl.constexpr,
HAS_PASS: tl.constexpr,
HAS_GATE: tl.constexpr,
ENABLE_PDL: tl.constexpr,
):
token = tl.program_id(0)
head = tl.program_id(1)
is_k = head >= NUM_Q_HEADS
local_head = tl.where(is_k, head - NUM_Q_HEADS, head)
out_dtype = tl.float16 if FP16 else tl.bfloat16
if is_k:
in_base = k_ptr + token * stride_k_t + local_head * HEAD_DIM
w_ptr = k_weight_ptr
out_base = k_out_ptr + token * stride_ko_t + local_head * HEAD_DIM
else:
if HAS_GATE:
in_base = q_gate_ptr + token * stride_qg_t + local_head * 2 * HEAD_DIM
else:
in_base = q_gate_ptr + token * stride_qg_t + local_head * HEAD_DIM
w_ptr = q_weight_ptr
out_base = q_out_ptr + token * stride_qo_t + local_head * HEAD_DIM
# Full load -> RMSNorm variance
head_offs = tl.arange(0, HEAD_BLOCK)
head_mask = head_offs < HEAD_DIM
x = tl.load(in_base + head_offs, mask=head_mask, other=0.0).to(tl.float32)
w = tl.load(w_ptr + head_offs, mask=head_mask, other=0.0).to(tl.float32)
var = tl.sum(x * x, axis=0) / HEAD_DIM
inv_rms = tl.rsqrt(var + EPS)
x_norm = (x * inv_rms * (w + 1.0)).to(out_dtype).to(tl.float32)
# Pass-through tail [rotary_dim, head_dim)
if HAS_PASS:
pass_mask = head_mask & (head_offs >= ROTARY_DIM)
tl.store(out_base + head_offs, x_norm, mask=pass_mask)
# Reload rotary portion from L1 -> re-norm -> RoPE
rot_offs = tl.arange(0, ROT_HALF_BLOCK)
rot_mask = rot_offs < HALF_ROTARY
xr1 = tl.load(in_base + rot_offs, mask=rot_mask, other=0.0).to(tl.float32)
xr2 = tl.load(in_base + HALF_ROTARY + rot_offs, mask=rot_mask, other=0.0).to(
tl.float32
)
wr1 = tl.load(w_ptr + rot_offs, mask=rot_mask, other=0.0).to(tl.float32)
wr2 = tl.load(w_ptr + HALF_ROTARY + rot_offs, mask=rot_mask, other=0.0).to(
tl.float32
)
xr1 = (xr1 * inv_rms * (wr1 + 1.0)).to(out_dtype).to(tl.float32)
xr2 = (xr2 * inv_rms * (wr2 + 1.0)).to(out_dtype).to(tl.float32)
pos = tl.load(positions_ptr + token).to(tl.int64)
cache_off = pos * stride_cos_t
cos = tl.load(
cos_sin_cache_ptr + cache_off + rot_offs, mask=rot_mask, other=0.0
).to(tl.float32)
sin = tl.load(
cos_sin_cache_ptr + cache_off + HALF_ROTARY + rot_offs, mask=rot_mask, other=0.0
).to(tl.float32)
tl.store(out_base + rot_offs, (xr1 * cos - xr2 * sin), mask=rot_mask)
tl.store(out_base + HALF_ROTARY + rot_offs, (xr2 * cos + xr1 * sin), mask=rot_mask)
# Gate copy (Q heads only)
if HAS_GATE and not is_k:
gate_in = in_base + HEAD_DIM
gate_out = gate_out_ptr + token * stride_gate_t + local_head * HEAD_DIM
g = tl.load(gate_in + head_offs, mask=head_mask, other=0.0)
tl.store(gate_out + head_offs, g, mask=head_mask)
# PDL: signal dependent kernels (attention/allreduce) can start early.
# Only available on NVIDIA Hopper+ (sm_90+); guarded for AMD/other backends.
if ENABLE_PDL:
tl.extra.cuda.gdc_launch_dependents()
def fused_qk_gemma_rmsnorm_rope_gate(
q_gate: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
cos_sin_cache: torch.Tensor,
positions: torch.Tensor,
eps: float,
num_q_heads: int,
num_kv_heads: int,
head_dim: int,
rotary_dim: int,
has_gate: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
"""Fused QK GemmaRMSNorm + NeoX RoPE + gate deinterleave.
Args:
q_gate: [T, num_q_heads * (1 + has_gate) * head_dim] — interleaved Q+Gate if has_gate
k: [T, num_kv_heads * head_dim]
q_weight, k_weight: [head_dim] — raw GemmaRMSNorm weights (kernel adds +1.0)
cos_sin_cache: [max_seq_len, rotary_dim] — [cos..., sin...]
positions: [T] — token positions
"""
T = q_gate.shape[0]
q_size = num_q_heads * head_dim
kv_size = num_kv_heads * head_dim
q_out = torch.empty(T, q_size, dtype=q_gate.dtype, device=q_gate.device)
k_out = torch.empty(T, kv_size, dtype=k.dtype, device=k.device)
gate_out = (
torch.empty(T, num_q_heads, head_dim, dtype=q_gate.dtype, device=q_gate.device)
if has_gate
else q_out
)
half_rotary = rotary_dim // 2
head_block = triton.next_power_of_2(head_dim)
rot_half_block = triton.next_power_of_2(half_rotary)
grid = (T, num_q_heads + num_kv_heads)
_fused_qk_rmsnorm_rope_gate_kernel[grid](
q_gate,
k,
q_out,
k_out,
gate_out,
q_weight,
k_weight,
cos_sin_cache,
positions,
q_gate.stride(0),
k.stride(0),
q_out.stride(0),
k_out.stride(0),
gate_out.stride(0),
cos_sin_cache.stride(0),
NUM_Q_HEADS=num_q_heads,
NUM_KV_HEADS=num_kv_heads,
HEAD_DIM=head_dim,
ROTARY_DIM=rotary_dim,
HALF_ROTARY=half_rotary,
HEAD_BLOCK=head_block,
ROT_HALF_BLOCK=rot_half_block,
EPS=eps,
FP16=q_gate.dtype == torch.float16,
HAS_PASS=rotary_dim < head_dim,
HAS_GATE=has_gate,
ENABLE_PDL=_ENABLE_PDL,
)
return q_out, k_out, gate_out if has_gate else None
+40 -1
View File
@@ -137,6 +137,11 @@ def _disable_shared_experts_fusion() -> bool:
return get_global_server_args().disable_shared_experts_fusion
if _is_cuda:
from sglang.srt.layers.fused_qk_rmsnorm_rope_gate import (
fused_qk_gemma_rmsnorm_rope_gate,
)
if _is_npu:
from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope import (
split_qkvgate_gemma_rmsnorm_rope,
@@ -887,6 +892,35 @@ class Qwen3_5AttentionDecoderLayer(nn.Module):
k = k_by_head.view(k.shape)
return q, k
def forward_prepare_cuda_fused(self, positions, hidden_states):
"""Fused QK GemmaRMSNorm + NeoX RoPE + gate deinterleave."""
qkv, _ = self.qkv_proj(hidden_states)
if self.attn_output_gate:
q_gate, k, v = qkv.split(
[self.q_size * 2, self.kv_size, self.kv_size], dim=-1
)
else:
q_gate, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q_out, k_out, gate_out = fused_qk_gemma_rmsnorm_rope_gate(
q_gate,
k,
self.q_norm.weight.data,
self.k_norm.weight.data,
self.rotary_emb.cos_sin_cache,
positions,
self.q_norm.variance_epsilon,
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.rotary_emb.rotary_dim,
has_gate=self.attn_output_gate,
)
seq_len = hidden_states.shape[0]
q = q_out.view(seq_len, -1)
k = k_out.view(seq_len, -1)
gate = gate_out.view(seq_len, -1) if gate_out is not None else None
return q, k, v, gate
def forward_prepare_native(self, positions, hidden_states):
qkv, _ = self.qkv_proj(hidden_states)
if self.attn_output_gate:
@@ -960,7 +994,12 @@ class Qwen3_5AttentionDecoderLayer(nn.Module):
forward_batch: ForwardBatch,
) -> torch.Tensor:
"""Full attention forward pass."""
if (_is_hip or _is_xpu) and self.attn_output_gate:
if _is_cuda and self.attn_output_gate:
q, k, v, gate = self.forward_prepare_cuda_fused(
positions=positions,
hidden_states=hidden_states,
)
elif (_is_hip or _is_xpu) and self.attn_output_gate:
q, k, v, gate = self.forward_prepare_fused_gate(
positions=positions,
hidden_states=hidden_states,