[GDN] Optimize prefill QKV split dispatch (#26206)
Co-authored-by: BBuf <xiaoyu.zhang@radixark.net>
This commit is contained in:
@@ -308,3 +308,100 @@ def fused_qkvzba_split_reshape_cat_contiguous(
|
||||
num_stages=3,
|
||||
)
|
||||
return mixed_qkv, z, b, a
|
||||
|
||||
|
||||
@triton.jit
|
||||
def fused_qkv_split_gdn_prefill_kernel(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mixed_qkv,
|
||||
MIXED_QKV_STRIDE_T: tl.constexpr,
|
||||
MIXED_QKV_STRIDE_D: tl.constexpr,
|
||||
NUM_Q_HEADS: tl.constexpr,
|
||||
NUM_K_HEADS: tl.constexpr,
|
||||
NUM_V_HEADS: tl.constexpr,
|
||||
HEAD_Q: tl.constexpr,
|
||||
HEAD_K: tl.constexpr,
|
||||
HEAD_V: tl.constexpr,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
i_t = tl.program_id(0)
|
||||
offsets = tl.arange(0, BLOCK_SIZE)
|
||||
|
||||
q_dim: tl.constexpr = NUM_Q_HEADS * HEAD_Q
|
||||
k_dim: tl.constexpr = NUM_K_HEADS * HEAD_K
|
||||
v_dim: tl.constexpr = NUM_V_HEADS * HEAD_V
|
||||
qk_dim: tl.constexpr = q_dim + k_dim
|
||||
qkv_dim: tl.constexpr = qk_dim + v_dim
|
||||
|
||||
mask = offsets < qkv_dim
|
||||
values = tl.load(
|
||||
mixed_qkv + i_t * MIXED_QKV_STRIDE_T + offsets * MIXED_QKV_STRIDE_D,
|
||||
mask=mask,
|
||||
)
|
||||
|
||||
q_mask = offsets < q_dim
|
||||
tl.store(q + i_t * q_dim + offsets, values, mask=q_mask)
|
||||
|
||||
k_offsets = offsets - q_dim
|
||||
k_mask = (offsets >= q_dim) & (offsets < qk_dim)
|
||||
tl.store(k + i_t * k_dim + k_offsets, values, mask=k_mask)
|
||||
|
||||
v_offsets = offsets - qk_dim
|
||||
v_mask = (offsets >= qk_dim) & (offsets < qkv_dim)
|
||||
tl.store(v + i_t * v_dim + v_offsets, values, mask=v_mask)
|
||||
|
||||
|
||||
def fused_qkv_split_gdn_prefill(
|
||||
mixed_qkv: torch.Tensor,
|
||||
num_q_heads: int,
|
||||
num_k_heads: int,
|
||||
num_v_heads: int,
|
||||
head_q: int,
|
||||
head_k: int,
|
||||
head_v: int,
|
||||
):
|
||||
"""Split packed post-conv GDN QKV into contiguous FLA prefill tensors.
|
||||
|
||||
`mixed_qkv` is laid out per token as `[all_q | all_k | all_v]`. The FLA
|
||||
chunk kernels consume separate contiguous `[1, T, H, D]` tensors, so this
|
||||
fused split replaces three independent `aten::copy_` kernels from the
|
||||
generic FLA input guard. `mixed_qkv` may be a strided `[T, qkv_dim]` view.
|
||||
"""
|
||||
seq_len = mixed_qkv.shape[0]
|
||||
q = torch.empty(
|
||||
(1, seq_len, num_q_heads, head_q),
|
||||
dtype=mixed_qkv.dtype,
|
||||
device=mixed_qkv.device,
|
||||
)
|
||||
k = torch.empty(
|
||||
(1, seq_len, num_k_heads, head_k),
|
||||
dtype=mixed_qkv.dtype,
|
||||
device=mixed_qkv.device,
|
||||
)
|
||||
v = torch.empty(
|
||||
(1, seq_len, num_v_heads, head_v),
|
||||
dtype=mixed_qkv.dtype,
|
||||
device=mixed_qkv.device,
|
||||
)
|
||||
|
||||
qkv_dim = num_q_heads * head_q + num_k_heads * head_k + num_v_heads * head_v
|
||||
fused_qkv_split_gdn_prefill_kernel[(seq_len,)](
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mixed_qkv,
|
||||
mixed_qkv.stride(0),
|
||||
mixed_qkv.stride(1),
|
||||
num_q_heads,
|
||||
num_k_heads,
|
||||
num_v_heads,
|
||||
head_q,
|
||||
head_k,
|
||||
head_v,
|
||||
BLOCK_SIZE=triton.next_power_of_2(qkv_dim),
|
||||
num_warps=8,
|
||||
num_stages=3,
|
||||
)
|
||||
return q, k, v
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
import os
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -20,6 +21,9 @@ from sglang.srt.layers.attention.fla.utils import (
|
||||
|
||||
NUM_WARPS = [2, 4] if is_nvidia_hopper else [2, 4, 8, 16]
|
||||
CHUNK_SIZE = 64
|
||||
GDN_CHUNK_H_BV = int(os.getenv("SGLANG_GDN_CHUNK_H_BV", "32"))
|
||||
GDN_CHUNK_H_NUM_WARPS = int(os.getenv("SGLANG_GDN_CHUNK_H_NUM_WARPS", "4"))
|
||||
GDN_CHUNK_H_NUM_STAGES = int(os.getenv("SGLANG_GDN_CHUNK_H_NUM_STAGES", "2"))
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
@@ -32,8 +36,16 @@ CHUNK_SIZE = 64
|
||||
# because cloning the cache pool for each benchmark exceeds available memory.
|
||||
# NT_BUCKET is kept in the autotune key for forward-compatibility (allows
|
||||
# future per-bucket configs once the kernel is refactored to write final
|
||||
# state to a separate output buffer).
|
||||
configs=[triton.Config({"BV": 32}, num_warps=4, num_stages=2)],
|
||||
# state to a separate output buffer). The env knobs keep this single-config
|
||||
# property while allowing model/hardware-local validation of the selected
|
||||
# tile without corrupting the state pool through multi-config autotune.
|
||||
configs=[
|
||||
triton.Config(
|
||||
{"BV": GDN_CHUNK_H_BV},
|
||||
num_warps=GDN_CHUNK_H_NUM_WARPS,
|
||||
num_stages=GDN_CHUNK_H_NUM_STAGES,
|
||||
)
|
||||
],
|
||||
key=["H", "K", "V", "BT", "USE_GK", "NT_BUCKET"],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
|
||||
@@ -26,6 +26,11 @@ if not is_cpu():
|
||||
CHUNK_SIZE as FLA_CHUNK_SIZE,
|
||||
)
|
||||
|
||||
if is_cuda():
|
||||
from sglang.jit_kernel.triton.gdn_fused_proj import fused_qkv_split_gdn_prefill
|
||||
|
||||
MAX_FUSED_QKV_SPLIT_DIM = 8192
|
||||
|
||||
if is_cuda():
|
||||
from sglang.srt.layers.attention.mamba.causal_conv1d import (
|
||||
causal_conv1d_fn as causal_conv1d_fn_cuda,
|
||||
@@ -444,16 +449,27 @@ class GDNAttnBackend(MambaAttnBackendBase):
|
||||
seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
|
||||
).transpose(0, 1)[:seq_len]
|
||||
|
||||
query, key, value = torch.split(
|
||||
mixed_qkv,
|
||||
[layer.q_dim, layer.k_dim, layer.v_dim],
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
actual_seq_len = query.shape[0]
|
||||
query = query.view(1, actual_seq_len, layer.num_q_heads, layer.head_q_dim)
|
||||
key = key.view(1, actual_seq_len, layer.num_k_heads, layer.head_k_dim)
|
||||
value = value.view(1, actual_seq_len, layer.num_v_heads, layer.head_v_dim)
|
||||
actual_seq_len = mixed_qkv.shape[0]
|
||||
qkv_dim = layer.q_dim + layer.k_dim + layer.v_dim
|
||||
if is_cuda() and qkv_dim <= MAX_FUSED_QKV_SPLIT_DIM:
|
||||
query, key, value = fused_qkv_split_gdn_prefill(
|
||||
mixed_qkv,
|
||||
layer.num_q_heads,
|
||||
layer.num_k_heads,
|
||||
layer.num_v_heads,
|
||||
layer.head_q_dim,
|
||||
layer.head_k_dim,
|
||||
layer.head_v_dim,
|
||||
)
|
||||
else:
|
||||
query, key, value = torch.split(
|
||||
mixed_qkv,
|
||||
[layer.q_dim, layer.k_dim, layer.v_dim],
|
||||
dim=-1,
|
||||
)
|
||||
query = query.view(1, actual_seq_len, layer.num_q_heads, layer.head_q_dim)
|
||||
key = key.view(1, actual_seq_len, layer.num_k_heads, layer.head_k_dim)
|
||||
value = value.view(1, actual_seq_len, layer.num_v_heads, layer.head_v_dim)
|
||||
|
||||
if is_target_verify:
|
||||
core_attn_out = self.kernel_dispatcher.target_verify(
|
||||
|
||||
Reference in New Issue
Block a user