Fused QK GemmaRMSNorm + RoPE + gate kernel for Qwen3.5 (#28320)
This commit is contained in:
@@ -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
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user