[Qwen3.5] Fuse split/reshape/cat ops in GDN projection with Triton kernel (#21019)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-03-23 23:17:01 +08:00
committed by GitHub
co-authored by luoyuan.luo
parent 8662ba7db4
commit 5bdc07d974
3 changed files with 597 additions and 202 deletions
@@ -0,0 +1,310 @@
from __future__ import annotations
import torch
import triton
import triton.language as tl
# =============================================================================
# Fused kernel — reads INTERLEAVED input format
# Used by Qwen3-Next whose checkpoint stores fused in_proj_qkvz weights
# in per-head-group interleaved layout:
# [g0_q, g0_k, g0_v, g0_z, g1_q, g1_k, g1_v, g1_z, ...]
# =============================================================================
@triton.jit
def fused_qkvzba_split_reshape_cat_kernel(
mixed_qkv,
z,
b,
a,
mixed_qkvz,
mixed_ba,
NUM_HEADS_QK: tl.constexpr,
NUM_HEADS_V: tl.constexpr,
HEAD_QK: tl.constexpr,
HEAD_V: tl.constexpr,
):
i_bs, i_qk = tl.program_id(0), tl.program_id(1)
QKVZ_DIM_T: tl.constexpr = HEAD_QK * 2 + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V * 2
BA_DIM_T: tl.constexpr = NUM_HEADS_V // NUM_HEADS_QK * 2
QKV_DIM_T: tl.constexpr = HEAD_QK * 2 + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V
q_end: tl.constexpr = HEAD_QK
blk_q_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(0, q_end)
)
k_end: tl.constexpr = q_end + HEAD_QK
blk_k_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(q_end, k_end)
)
v_end: tl.constexpr = k_end + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V
blk_v_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(k_end, v_end)
)
z_end: tl.constexpr = v_end + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V
blk_z_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(v_end, z_end)
)
blk_q_st_ptr = (
mixed_qkv
+ i_bs * NUM_HEADS_QK * QKV_DIM_T
+ i_qk * HEAD_QK
+ tl.arange(0, HEAD_QK)
)
blk_k_st_ptr = (
mixed_qkv
+ i_bs * NUM_HEADS_QK * QKV_DIM_T
+ NUM_HEADS_QK * HEAD_QK
+ i_qk * HEAD_QK
+ tl.arange(0, HEAD_QK)
)
blk_v_st_ptr = (
mixed_qkv
+ i_bs * NUM_HEADS_QK * QKV_DIM_T
+ NUM_HEADS_QK * HEAD_QK * 2
+ i_qk * HEAD_V * NUM_HEADS_V // NUM_HEADS_QK
+ tl.arange(0, HEAD_V * NUM_HEADS_V // NUM_HEADS_QK)
)
blk_z_st_ptr = (
z
+ i_bs * NUM_HEADS_V * HEAD_V
+ i_qk * HEAD_V * NUM_HEADS_V // NUM_HEADS_QK
+ tl.arange(0, HEAD_V * NUM_HEADS_V // NUM_HEADS_QK)
)
tl.store(blk_q_st_ptr, tl.load(blk_q_ptr))
tl.store(blk_k_st_ptr, tl.load(blk_k_ptr))
tl.store(blk_v_st_ptr, tl.load(blk_v_ptr))
tl.store(blk_z_st_ptr, tl.load(blk_z_ptr))
b_end: tl.constexpr = NUM_HEADS_V // NUM_HEADS_QK
a_end: tl.constexpr = b_end + NUM_HEADS_V // NUM_HEADS_QK
for i in tl.static_range(b_end):
blk_b_ptr = mixed_ba + i_bs * NUM_HEADS_QK * BA_DIM_T + i_qk * BA_DIM_T + i
blk_b_st_ptr = b + i_bs * NUM_HEADS_V + i_qk * NUM_HEADS_V // NUM_HEADS_QK + i
tl.store(blk_b_st_ptr, tl.load(blk_b_ptr))
for i in tl.static_range(b_end, a_end):
blk_a_ptr = mixed_ba + i_bs * NUM_HEADS_QK * BA_DIM_T + i_qk * BA_DIM_T + i
blk_a_st_ptr = (
a + i_bs * NUM_HEADS_V + i_qk * NUM_HEADS_V // NUM_HEADS_QK + (i - b_end)
)
tl.store(blk_a_st_ptr, tl.load(blk_a_ptr))
def fused_qkvzba_split_reshape_cat(
mixed_qkvz,
mixed_ba,
num_heads_qk,
num_heads_v,
head_qk,
head_v,
):
batch, seq_len = mixed_qkvz.shape[0], 1
qkv_dim_t = num_heads_qk * head_qk * 2 + num_heads_v * head_v
mixed_qkv = torch.empty(
[batch * seq_len, qkv_dim_t],
dtype=mixed_qkvz.dtype,
device=mixed_qkvz.device,
)
z = torch.empty(
[batch * seq_len, num_heads_v, head_v],
dtype=mixed_qkvz.dtype,
device=mixed_qkvz.device,
)
b = torch.empty(
[batch * seq_len, num_heads_v],
dtype=mixed_ba.dtype,
device=mixed_ba.device,
)
a = torch.empty_like(b)
grid = (batch * seq_len, num_heads_qk)
fused_qkvzba_split_reshape_cat_kernel[grid](
mixed_qkv,
z,
b,
a,
mixed_qkvz,
mixed_ba,
num_heads_qk,
num_heads_v,
head_qk,
head_v,
num_warps=1,
num_stages=3,
)
return mixed_qkv, z, b, a
# =============================================================================
# Fused kernel — reads CONTIGUOUS input format
# Used by Qwen3.5 whose checkpoint stores in_proj_qkv and in_proj_z separately.
# After MergedColumnParallelLinear loads them, the matmul output is contiguous:
# mixed_qkvz: [all_q | all_k | all_v | all_z]
# mixed_ba: [all_b | all_a]
#
# Output format is identical to the interleaved kernel (same downstream consumer).
# =============================================================================
@triton.jit
def fused_qkvzba_split_reshape_cat_contiguous_kernel(
mixed_qkv,
z,
b,
a,
mixed_qkvz,
mixed_ba,
NUM_HEADS_QK: tl.constexpr,
NUM_HEADS_V: tl.constexpr,
HEAD_QK: tl.constexpr,
HEAD_V: tl.constexpr,
):
i_bs, i_qk = tl.program_id(0), tl.program_id(1)
V_PER_GROUP: tl.constexpr = NUM_HEADS_V // NUM_HEADS_QK
# ── Input dimensions (contiguous layout) ──
TOTAL_Q: tl.constexpr = NUM_HEADS_QK * HEAD_QK
TOTAL_K: tl.constexpr = NUM_HEADS_QK * HEAD_QK
TOTAL_V: tl.constexpr = NUM_HEADS_V * HEAD_V
TOTAL_QKVZ: tl.constexpr = TOTAL_Q + TOTAL_K + TOTAL_V + TOTAL_V
TOTAL_BA: tl.constexpr = NUM_HEADS_V * 2
# ── Output dimensions ──
QKV_DIM_T: tl.constexpr = TOTAL_Q + TOTAL_K + TOTAL_V
# ── Read from contiguous input ──
# q for head group i_qk: in the all_q region, offset i_qk * HEAD_QK
blk_q_ptr = mixed_qkvz + i_bs * TOTAL_QKVZ + i_qk * HEAD_QK + tl.arange(0, HEAD_QK)
# k for head group i_qk: in the all_k region
blk_k_ptr = (
mixed_qkvz
+ i_bs * TOTAL_QKVZ
+ TOTAL_Q
+ i_qk * HEAD_QK
+ tl.arange(0, HEAD_QK)
)
# v for head group i_qk: in the all_v region
blk_v_ptr = (
mixed_qkvz
+ i_bs * TOTAL_QKVZ
+ TOTAL_Q
+ TOTAL_K
+ i_qk * V_PER_GROUP * HEAD_V
+ tl.arange(0, V_PER_GROUP * HEAD_V)
)
# z for head group i_qk: in the all_z region
blk_z_ptr = (
mixed_qkvz
+ i_bs * TOTAL_QKVZ
+ TOTAL_Q
+ TOTAL_K
+ TOTAL_V
+ i_qk * V_PER_GROUP * HEAD_V
+ tl.arange(0, V_PER_GROUP * HEAD_V)
)
# ── Write to output (identical layout to the interleaved kernel) ──
blk_q_st_ptr = mixed_qkv + i_bs * QKV_DIM_T + i_qk * HEAD_QK + tl.arange(0, HEAD_QK)
blk_k_st_ptr = (
mixed_qkv
+ i_bs * QKV_DIM_T
+ NUM_HEADS_QK * HEAD_QK
+ i_qk * HEAD_QK
+ tl.arange(0, HEAD_QK)
)
blk_v_st_ptr = (
mixed_qkv
+ i_bs * QKV_DIM_T
+ NUM_HEADS_QK * HEAD_QK * 2
+ i_qk * V_PER_GROUP * HEAD_V
+ tl.arange(0, V_PER_GROUP * HEAD_V)
)
blk_z_st_ptr = (
z
+ i_bs * NUM_HEADS_V * HEAD_V
+ i_qk * V_PER_GROUP * HEAD_V
+ tl.arange(0, V_PER_GROUP * HEAD_V)
)
tl.store(blk_q_st_ptr, tl.load(blk_q_ptr))
tl.store(blk_k_st_ptr, tl.load(blk_k_ptr))
tl.store(blk_v_st_ptr, tl.load(blk_v_ptr))
tl.store(blk_z_st_ptr, tl.load(blk_z_ptr))
# ── b and a from contiguous [all_b | all_a] ──
for i in tl.static_range(V_PER_GROUP):
blk_b_ptr = mixed_ba + i_bs * TOTAL_BA + i_qk * V_PER_GROUP + i
blk_b_st_ptr = b + i_bs * NUM_HEADS_V + i_qk * V_PER_GROUP + i
tl.store(blk_b_st_ptr, tl.load(blk_b_ptr))
for i in tl.static_range(V_PER_GROUP):
blk_a_ptr = mixed_ba + i_bs * TOTAL_BA + NUM_HEADS_V + i_qk * V_PER_GROUP + i
blk_a_st_ptr = a + i_bs * NUM_HEADS_V + i_qk * V_PER_GROUP + i
tl.store(blk_a_st_ptr, tl.load(blk_a_ptr))
def fused_qkvzba_split_reshape_cat_contiguous(
mixed_qkvz,
mixed_ba,
num_heads_qk,
num_heads_v,
head_qk,
head_v,
):
"""Fused split/reshape/cat for CONTIGUOUS input format (Qwen3.5).
Input layout:
mixed_qkvz: [all_q | all_k | all_v | all_z]
mixed_ba: [all_b | all_a]
Output layout (same as fused_qkvzba_split_reshape_cat):
mixed_qkv: [all_q | all_k | all_v] (z stripped)
z: [num_v_heads, head_v]
b: [num_v_heads]
a: [num_v_heads]
"""
batch, seq_len = mixed_qkvz.shape[0], 1
qkv_dim_t = num_heads_qk * head_qk * 2 + num_heads_v * head_v
mixed_qkv = torch.empty(
[batch * seq_len, qkv_dim_t],
dtype=mixed_qkvz.dtype,
device=mixed_qkvz.device,
)
z = torch.empty(
[batch * seq_len, num_heads_v, head_v],
dtype=mixed_qkvz.dtype,
device=mixed_qkvz.device,
)
b = torch.empty(
[batch * seq_len, num_heads_v],
dtype=mixed_ba.dtype,
device=mixed_ba.device,
)
a = torch.empty_like(b)
grid = (batch * seq_len, num_heads_qk)
fused_qkvzba_split_reshape_cat_contiguous_kernel[grid](
mixed_qkv,
z,
b,
a,
mixed_qkvz,
mixed_ba,
num_heads_qk,
num_heads_v,
head_qk,
head_v,
num_warps=1,
num_stages=3,
)
return mixed_qkv, z, b, a
+278 -58
View File
@@ -20,6 +20,11 @@ from typing import Iterable, Optional, Set, Tuple, Union
import torch import torch
import torch.nn as nn import torch.nn as nn
import triton
from sglang.jit_kernel.triton.gdn_fused_proj import (
fused_qkvzba_split_reshape_cat_contiguous,
)
# Configs # Configs
from sglang.srt.configs.qwen3_5 import ( from sglang.srt.configs.qwen3_5 import (
@@ -54,6 +59,10 @@ from sglang.srt.layers.linear import (
RowParallelLinear, RowParallelLinear,
) )
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.parameter import (
BlockQuantScaleParameter,
PerTensorScaleParameter,
)
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
@@ -70,11 +79,14 @@ from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
# Models # Models
from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration
from sglang.srt.server_args import get_global_server_args
# Utils # Utils
from sglang.srt.utils import ( from sglang.srt.utils import (
LazyValue, LazyValue,
add_prefix, add_prefix,
cpu_has_amx_support,
is_cpu,
is_cuda, is_cuda,
is_npu, is_npu,
make_layers, make_layers,
@@ -85,6 +97,9 @@ from sglang.srt.utils.hf_transformers_utils import get_processor, get_rope_confi
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_is_cuda = is_cuda() _is_cuda = is_cuda()
_is_npu = is_npu() _is_npu = is_npu()
_is_cpu = is_cpu()
_is_amx_available = cpu_has_amx_support()
cached_get_processor = lru_cache(get_processor) cached_get_processor = lru_cache(get_processor)
@@ -129,54 +144,39 @@ class Qwen3_5GatedDeltaNet(nn.Module):
) )
self.conv1d.weight.data = self.conv1d.weight.data.unsqueeze(1) self.conv1d.weight.data = self.conv1d.weight.data.unsqueeze(1)
# Split projection layers (following vLLM's implementation) # projection of the input hidden states
# Instead of fused in_proj_qkvz and in_proj_ba, use separate layers self.in_proj_qkvz = self.create_qkvz_proj(
self.in_proj_qkv = MergedColumnParallelLinear( hidden_size=self.hidden_size,
input_size=self.hidden_size, key_dim=self.key_dim,
output_sizes=[self.key_dim, self.key_dim, self.value_dim], value_dim=self.value_dim,
bias=False,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("in_proj_qkvz", prefix),
tp_rank=self.attn_tp_rank, tp_rank=self.attn_tp_rank,
tp_size=self.attn_tp_size, tp_size=self.attn_tp_size,
prefix=add_prefix("in_proj_qkv", prefix),
) )
self.in_proj_z = ColumnParallelLinear(
input_size=self.hidden_size, self.in_proj_ba = self.create_ba_proj(
output_size=self.value_dim, hidden_size=self.hidden_size,
bias=False, num_v_heads=self.num_v_heads,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("in_proj_ba", prefix),
tp_rank=self.attn_tp_rank, tp_rank=self.attn_tp_rank,
tp_size=self.attn_tp_size, tp_size=self.attn_tp_size,
prefix=add_prefix("in_proj_z", prefix),
)
self.in_proj_b = ColumnParallelLinear(
input_size=self.hidden_size,
output_size=self.num_v_heads,
bias=False,
quant_config=quant_config,
tp_rank=self.attn_tp_rank,
tp_size=self.attn_tp_size,
prefix=add_prefix("in_proj_b", prefix),
)
self.in_proj_a = ColumnParallelLinear(
input_size=self.hidden_size,
output_size=self.num_v_heads,
bias=False,
quant_config=quant_config,
tp_rank=self.attn_tp_rank,
tp_size=self.attn_tp_size,
prefix=add_prefix("in_proj_a", prefix),
) )
# Override weight loaders for packed checkpoint format.
# Important: for FP8, this must cover not only `.weight` but also
# `weight_scale_inv` / `weight_scale` / `input_scale` if present.
self._bind_packed_weight_loaders(self.in_proj_qkvz)
self._bind_packed_weight_loaders(self.in_proj_ba)
# Conv1d weight loader setup # Conv1d weight loader setup
query_key_settings = (self.key_dim, 0, False) query_key_settings = (self.key_dim, 0, False)
value_settings = (self.value_dim, 0, False) value_settings = (self.value_dim, 0, False)
delattr(self.conv1d.weight, "weight_loader") self._override_weight_loader(
set_weight_attrs(
self.conv1d.weight, self.conv1d.weight,
{ mamba_v2_sharded_weight_loader(
"weight_loader": mamba_v2_sharded_weight_loader(
[ [
query_key_settings, query_key_settings,
query_key_settings, query_key_settings,
@@ -184,8 +184,7 @@ class Qwen3_5GatedDeltaNet(nn.Module):
], ],
self.attn_tp_size, self.attn_tp_size,
self.attn_tp_rank, self.attn_tp_rank,
) ),
},
) )
# State parameters # State parameters
@@ -202,7 +201,6 @@ class Qwen3_5GatedDeltaNet(nn.Module):
conv_weights = self.conv1d.weight.view( conv_weights = self.conv1d.weight.view(
self.conv1d.weight.size(0), self.conv1d.weight.size(2) self.conv1d.weight.size(0), self.conv1d.weight.size(2)
) )
# RadixLinearAttention layer
self.attn = RadixLinearAttention( self.attn = RadixLinearAttention(
layer_id=layer_id, layer_id=layer_id,
num_q_heads=self.num_k_heads // self.attn_tp_size, num_q_heads=self.num_k_heads // self.attn_tp_size,
@@ -218,7 +216,6 @@ class Qwen3_5GatedDeltaNet(nn.Module):
dt_bias=self.dt_bias, dt_bias=self.dt_bias,
) )
# Normalization layer
self.norm = RMSNormGated( self.norm = RMSNormGated(
self.head_v_dim, self.head_v_dim,
eps=self.layer_norm_epsilon, eps=self.layer_norm_epsilon,
@@ -228,7 +225,6 @@ class Qwen3_5GatedDeltaNet(nn.Module):
dtype=config.torch_dtype, dtype=config.torch_dtype,
) )
# Output projection
self.out_proj = RowParallelLinear( self.out_proj = RowParallelLinear(
self.value_dim, self.value_dim,
self.hidden_size, self.hidden_size,
@@ -241,16 +237,190 @@ class Qwen3_5GatedDeltaNet(nn.Module):
prefix=add_prefix("out_proj", prefix), prefix=add_prefix("out_proj", prefix),
) )
@staticmethod
def _override_weight_loader(param, loader):
"""Robustly override loader for:
1) BasevLLMParameter subclasses: real storage is `_weight_loader`
2) regular Parameters that already have mutable `weight_loader`
3) regular Parameters without `weight_loader` yet
"""
if hasattr(param, "_weight_loader"):
# FP8 / quantized BasevLLMParameter path
param._weight_loader = loader
return
if hasattr(param, "weight_loader"):
# Regular parameter/tensor that already has a mutable attr.
# Do NOT call set_weight_attrs here, because it asserts when
# overwriting an existing attribute.
param.weight_loader = loader
return
# Fresh attribute on a normal tensor/Parameter
set_weight_attrs(param, {"weight_loader": loader})
def _bind_packed_weight_loaders(self, module):
"""Bind packed-checkpoint-aware loaders to all relevant params of a merged module."""
for attr_name in ("weight", "weight_scale_inv", "weight_scale", "input_scale"):
param = getattr(module, attr_name, None)
if param is None:
continue
original_loader = getattr(param, "weight_loader", None)
if original_loader is None:
continue
wrapped_loader = self._make_packed_weight_loader(module, original_loader)
self._override_weight_loader(param, wrapped_loader)
@staticmethod
def _get_split_sizes_for_param(module, param, loaded_shard_id):
"""Return checkpoint-side split sizes for this param type."""
if isinstance(param, BlockQuantScaleParameter):
# Split by output blocks, not raw output sizes.
block_n, _ = module.quant_method.quant_config.weight_block_size
block_n = 1 if getattr(param, "format_ue8m0", False) else block_n
return [
(module.output_sizes[idx] + block_n - 1) // block_n
for idx in loaded_shard_id
]
if isinstance(param, PerTensorScaleParameter):
# One logical scale per logical shard.
return [1 for _ in loaded_shard_id]
# Normal weight / non-block quant tensor
return [module.output_sizes[idx] for idx in loaded_shard_id]
@classmethod
def _make_packed_weight_loader(cls, module, original_weight_loader):
"""Wrap the param's original loader so split checkpoints:
- in_proj_qkv + in_proj_z -> merged in_proj_qkvz
- in_proj_b + in_proj_a -> merged in_proj_ba
can load correctly for both normal and FP8 params.
"""
def weight_loader(param, loaded_weight, loaded_shard_id=None):
# Only intercept split-checkpoint tuple shards.
# int shard_id and None should preserve original behavior.
if isinstance(loaded_shard_id, tuple):
split_sizes = cls._get_split_sizes_for_param(
module, param, loaded_shard_id
)
if len(loaded_weight.shape) == 0:
# Scalar only makes sense for a single logical shard.
assert len(split_sizes) == 1 and split_sizes[0] == 1, (
f"Unexpected scalar for tuple shard load: "
f"{loaded_shard_id=}, {split_sizes=}"
)
chunks = [loaded_weight.reshape(1)]
else:
split_dim = getattr(param, "output_dim", 0)
chunks = loaded_weight.split(split_sizes, dim=split_dim)
assert len(chunks) == len(loaded_shard_id), (
f"Chunk/shard mismatch: {len(chunks)=}, "
f"{len(loaded_shard_id)=}, {split_sizes=}"
)
for idx, chunk in zip(loaded_shard_id, chunks):
# Delegate each chunk to the param's original int-shard loader.
original_weight_loader(param, chunk, idx)
return
return original_weight_loader(param, loaded_weight, loaded_shard_id)
return weight_loader
def create_qkvz_proj(
self,
hidden_size: int,
key_dim: int,
value_dim: int,
quant_config: QuantizationConfig | None,
prefix: str,
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
) -> MergedColumnParallelLinear:
return MergedColumnParallelLinear(
input_size=hidden_size,
output_sizes=[key_dim, key_dim, value_dim, value_dim],
bias=False,
quant_config=quant_config,
prefix=prefix,
tp_rank=tp_rank,
tp_size=tp_size,
)
def create_ba_proj(
self,
hidden_size: int,
num_v_heads: int,
quant_config: QuantizationConfig | None,
prefix: str,
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
) -> MergedColumnParallelLinear:
# Qwen3.5 has separate in_proj_b and in_proj_a weights in the
# checkpoint, which are loaded into the fused in_proj_ba parameter
# via stacked_params_mapping with shard_id 0 and 1 respectively.
return MergedColumnParallelLinear(
input_size=hidden_size,
output_sizes=[num_v_heads, num_v_heads],
bias=False,
quant_config=quant_config,
prefix=prefix,
tp_rank=tp_rank,
tp_size=tp_size,
)
def fix_query_key_value_ordering( def fix_query_key_value_ordering(
self, self,
mixed_qkv, mixed_qkvz: torch.Tensor,
z, mixed_ba: torch.Tensor,
b,
a,
): ):
raise NotImplementedError( """
"Qwen3.5 Series dont need to fix query key value ordering" Derives `query`, `key` and `value` tensors from `mixed_qkvzba`.
) """
k_tp = self.key_dim // self.attn_tp_size
v_tp = self.value_dim // self.attn_tp_size
nv_tp = self.num_v_heads // self.attn_tp_size
# Directly split, no head group reshape
query, key, value, z = mixed_qkvz.split([k_tp, k_tp, v_tp, v_tp], dim=-1)
b, a = mixed_ba.split([nv_tp, nv_tp], dim=-1)
# value / z reshape to (seq, num_v_heads/tp, head_v_dim)
value = value.reshape(value.size(0), -1, self.head_v_dim)
z = z.reshape(z.size(0), -1, self.head_v_dim)
return query, key, value, z, b, a
def _forward_input_proj(self, hidden_states: torch.Tensor):
if (
_is_cpu
or _is_npu
or not get_global_server_args().disable_piecewise_cuda_graph
):
DUAL_STREAM_TOKEN_THRESHOLD = 0
else:
DUAL_STREAM_TOKEN_THRESHOLD = 1024
seq_len, _ = hidden_states.shape
if (
self.alt_stream is not None
and get_is_capture_mode()
and seq_len < DUAL_STREAM_TOKEN_THRESHOLD
):
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
projected_states_qkvz, _ = self.in_proj_qkvz(hidden_states)
with torch.cuda.stream(self.alt_stream):
projected_states_ba, _ = self.in_proj_ba(hidden_states)
current_stream.wait_stream(self.alt_stream)
else:
projected_states_qkvz, _ = self.in_proj_qkvz(hidden_states)
projected_states_ba, _ = self.in_proj_ba(hidden_states)
return projected_states_qkvz, projected_states_ba
def forward( def forward(
self, self,
@@ -263,30 +433,60 @@ class Qwen3_5GatedDeltaNet(nn.Module):
2. Core attention (custom op) 2. Core attention (custom op)
3. Output projection 3. Output projection
""" """
seq_len, _ = hidden_states.shape projected_states_qkvz, projected_states_ba = self._forward_input_proj(
hidden_states
mixed_qkv, _ = self.in_proj_qkv(hidden_states) )
z, _ = self.in_proj_z(hidden_states)
z = z.reshape(z.size(0), -1, self.head_v_dim)
b, _ = self.in_proj_b(hidden_states)
a, _ = self.in_proj_a(hidden_states)
b = b.contiguous()
a = a.contiguous()
if self.num_v_heads // self.num_k_heads in [1, 2, 4] and not _is_cpu:
mixed_qkv, z, b, a = fused_qkvzba_split_reshape_cat_contiguous(
projected_states_qkvz,
projected_states_ba,
triton.cdiv(self.num_k_heads, self.attn_tp_size),
triton.cdiv(self.num_v_heads, self.attn_tp_size),
self.head_k_dim,
self.head_v_dim,
)
elif _is_cpu and _is_amx_available:
mixed_qkv, z, b, a = (
torch.ops.sgl_kernel.fused_qkvzba_split_reshape_cat_cpu(
projected_states_qkvz,
projected_states_ba,
self.num_k_heads // self.attn_tp_size,
self.num_v_heads // self.attn_tp_size,
self.head_k_dim,
self.head_v_dim,
)
)
else:
query, key, value, z, b, a = self.fix_query_key_value_ordering(
projected_states_qkvz, projected_states_ba
)
query, key, value = map(
lambda x: x.reshape(x.shape[0], -1), (query, key, value)
)
mixed_qkv = torch.cat((query, key, value), dim=-1)
core_attn_out = self.attn( core_attn_out = self.attn(
forward_batch=forward_batch, forward_batch,
mixed_qkv=mixed_qkv, mixed_qkv=mixed_qkv,
a=a, a=a,
b=b, b=b,
) )
z_shape_og = z.shape z_shape_og = z.shape
# reshape input data into 2D tensor
core_attn_out = core_attn_out.reshape(-1, core_attn_out.shape[-1]) core_attn_out = core_attn_out.reshape(-1, core_attn_out.shape[-1])
z = z.reshape(-1, z.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 = self.norm(core_attn_out, z)
core_attn_out = core_attn_out.reshape(z_shape_og) core_attn_out = core_attn_out.reshape(z_shape_og)
core_attn_out = core_attn_out.flatten(-2) # ... h d -> ... (h d) core_attn_out = core_attn_out.reshape(*core_attn_out.shape[:-2], -1)
output, _ = self.out_proj(core_attn_out) output, _ = self.out_proj(core_attn_out)
return output return output
@@ -818,6 +1018,11 @@ class Qwen3_5ForCausalLM(nn.Module):
("qkv_proj", "v_proj", "v"), ("qkv_proj", "v_proj", "v"),
("gate_up_proj", "gate_proj", 0), ("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1), ("gate_up_proj", "up_proj", 1),
# GDN
("in_proj_qkvz.", "in_proj_qkv.", (0, 1, 2)),
("in_proj_qkvz.", "in_proj_z.", 3),
("in_proj_ba.", "in_proj_b.", 0),
("in_proj_ba.", "in_proj_a.", 1),
] ]
loaded_params: Set[str] = set() loaded_params: Set[str] = set()
@@ -894,6 +1099,11 @@ class Qwen3_5MoeForCausalLM(Qwen3_5ForCausalLM):
("qkv_proj", "v_proj", "v"), ("qkv_proj", "v_proj", "v"),
("gate_up_proj", "gate_proj", 0), ("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1), ("gate_up_proj", "up_proj", 1),
# GDN
("in_proj_qkvz.", "in_proj_qkv.", (0, 1, 2)),
("in_proj_qkvz.", "in_proj_z.", 3),
("in_proj_ba.", "in_proj_b.", 0),
("in_proj_ba.", "in_proj_a.", 1),
] ]
# Params for weights, fp8 weight scales, fp8 activation scales # Params for weights, fp8 weight scales, fp8 activation scales
@@ -1127,6 +1337,11 @@ class Qwen3_5ForConditionalGeneration(Qwen3VLForConditionalGeneration):
("qkv_proj", "v_proj", "v"), ("qkv_proj", "v_proj", "v"),
("gate_up_proj", "gate_proj", 0), ("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1), ("gate_up_proj", "up_proj", 1),
# GDN fused projections
("in_proj_qkvz.", "in_proj_qkv.", (0, 1, 2)),
("in_proj_qkvz.", "in_proj_z.", 3),
("in_proj_ba.", "in_proj_b.", 0),
("in_proj_ba.", "in_proj_a.", 1),
] ]
loaded_params: Set[str] = set() loaded_params: Set[str] = set()
@@ -1223,6 +1438,11 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
("qkv_proj", "v_proj", "v"), ("qkv_proj", "v_proj", "v"),
("gate_up_proj", "gate_proj", 0), ("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1), ("gate_up_proj", "up_proj", 1),
# GDN fused projections
("in_proj_qkvz.", "in_proj_qkv.", (0, 1, 2)),
("in_proj_qkvz.", "in_proj_z.", 3),
("in_proj_ba.", "in_proj_b.", 0),
("in_proj_ba.", "in_proj_a.", 1),
] ]
# Params for weights, fp8 weight scales, fp8 activation scales # Params for weights, fp8 weight scales, fp8 activation scales
+2 -137
View File
@@ -3,6 +3,7 @@ import logging
from typing import Any, Iterable, Optional, Set, Tuple from typing import Any, Iterable, Optional, Set, Tuple
import torch import torch
import triton
from torch import nn from torch import nn
from sglang.srt.configs.qwen3_next import Qwen3NextConfig from sglang.srt.configs.qwen3_next import Qwen3NextConfig
@@ -55,6 +56,7 @@ from sglang.srt.utils import (
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from sglang.jit_kernel.triton.gdn_fused_proj import fused_qkvzba_split_reshape_cat
from sglang.srt.layers.attention.fla.fused_norm_gate import FusedRMSNormGated from sglang.srt.layers.attention.fla.fused_norm_gate import FusedRMSNormGated
_is_cuda = is_cuda() _is_cuda = is_cuda()
@@ -63,143 +65,6 @@ _is_cpu = is_cpu()
_is_amx_available = cpu_has_amx_support() _is_amx_available = cpu_has_amx_support()
import triton
import triton.language as tl
@triton.jit
def fused_qkvzba_split_reshape_cat_kernel(
mixed_qkv,
z,
b,
a,
mixed_qkvz,
mixed_ba,
NUM_HEADS_QK: tl.constexpr,
NUM_HEADS_V: tl.constexpr,
HEAD_QK: tl.constexpr,
HEAD_V: tl.constexpr,
):
i_bs, i_qk = tl.program_id(0), tl.program_id(1)
QKVZ_DIM_T: tl.constexpr = HEAD_QK * 2 + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V * 2
BA_DIM_T: tl.constexpr = NUM_HEADS_V // NUM_HEADS_QK * 2
QKV_DIM_T: tl.constexpr = HEAD_QK * 2 + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V
q_end: tl.constexpr = HEAD_QK
blk_q_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(0, q_end)
)
k_end: tl.constexpr = q_end + HEAD_QK
blk_k_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(q_end, k_end)
)
v_end: tl.constexpr = k_end + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V
blk_v_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(k_end, v_end)
)
z_end: tl.constexpr = v_end + NUM_HEADS_V // NUM_HEADS_QK * HEAD_V
blk_z_ptr = (
mixed_qkvz
+ i_bs * NUM_HEADS_QK * QKVZ_DIM_T
+ i_qk * QKVZ_DIM_T
+ tl.arange(v_end, z_end)
)
blk_q_st_ptr = (
mixed_qkv
+ i_bs * NUM_HEADS_QK * QKV_DIM_T
+ i_qk * HEAD_QK
+ tl.arange(0, HEAD_QK)
)
blk_k_st_ptr = (
mixed_qkv
+ i_bs * NUM_HEADS_QK * QKV_DIM_T
+ NUM_HEADS_QK * HEAD_QK
+ i_qk * HEAD_QK
+ tl.arange(0, HEAD_QK)
)
blk_v_st_ptr = (
mixed_qkv
+ i_bs * NUM_HEADS_QK * QKV_DIM_T
+ NUM_HEADS_QK * HEAD_QK * 2
+ i_qk * HEAD_V * NUM_HEADS_V // NUM_HEADS_QK
+ tl.arange(0, HEAD_V * NUM_HEADS_V // NUM_HEADS_QK)
)
blk_z_st_ptr = (
z
+ i_bs * NUM_HEADS_V * HEAD_V
+ i_qk * HEAD_V * NUM_HEADS_V // NUM_HEADS_QK
+ tl.arange(0, HEAD_V * NUM_HEADS_V // NUM_HEADS_QK)
)
tl.store(blk_q_st_ptr, tl.load(blk_q_ptr))
tl.store(blk_k_st_ptr, tl.load(blk_k_ptr))
tl.store(blk_v_st_ptr, tl.load(blk_v_ptr))
tl.store(blk_z_st_ptr, tl.load(blk_z_ptr))
b_end: tl.constexpr = NUM_HEADS_V // NUM_HEADS_QK
a_end: tl.constexpr = b_end + NUM_HEADS_V // NUM_HEADS_QK
for i in tl.static_range(b_end):
blk_b_ptr = mixed_ba + i_bs * NUM_HEADS_QK * BA_DIM_T + i_qk * BA_DIM_T + i
blk_b_st_ptr = b + i_bs * NUM_HEADS_V + i_qk * NUM_HEADS_V // NUM_HEADS_QK + i
tl.store(blk_b_st_ptr, tl.load(blk_b_ptr))
for i in tl.static_range(b_end, a_end):
blk_a_ptr = mixed_ba + i_bs * NUM_HEADS_QK * BA_DIM_T + i_qk * BA_DIM_T + i
blk_a_st_ptr = (
a + i_bs * NUM_HEADS_V + i_qk * NUM_HEADS_V // NUM_HEADS_QK + (i - b_end)
)
tl.store(blk_a_st_ptr, tl.load(blk_a_ptr))
def fused_qkvzba_split_reshape_cat(
mixed_qkvz,
mixed_ba,
num_heads_qk,
num_heads_v,
head_qk,
head_v,
):
batch, seq_len = mixed_qkvz.shape[0], 1
qkv_dim_t = num_heads_qk * head_qk * 2 + num_heads_v * head_v
mixed_qkv = torch.empty(
[batch * seq_len, qkv_dim_t],
dtype=mixed_qkvz.dtype,
device=mixed_qkvz.device,
)
z = torch.empty(
[batch * seq_len, num_heads_v, head_v],
dtype=mixed_qkvz.dtype,
device=mixed_qkvz.device,
)
b = torch.empty(
[batch * seq_len, num_heads_v],
dtype=mixed_ba.dtype,
device=mixed_ba.device,
)
a = torch.empty_like(b)
grid = (batch * seq_len, num_heads_qk)
fused_qkvzba_split_reshape_cat_kernel[grid](
mixed_qkv,
z,
b,
a,
mixed_qkvz,
mixed_ba,
num_heads_qk,
num_heads_v,
head_qk,
head_v,
num_warps=1,
num_stages=3,
)
return mixed_qkv, z, b, a
class Qwen3GatedDeltaNet(nn.Module): class Qwen3GatedDeltaNet(nn.Module):
def __init__( def __init__(
self, self,