[AMD] Improve K3 dspark draft attn kernel perf (#35499)
This commit is contained in:
@@ -30,6 +30,7 @@ _BLOCK_CONFIG = {
|
||||
# head_dim: (BLOCK_H, BLOCK_N, num_warps)
|
||||
256: (4, 64, 8), # Qwen3.5 TP2 / TP4 / TP8
|
||||
576: (4, 64, 8), # K3 MLA (kv_lora_rank 512 + qk_rope 64)
|
||||
64: (4, 256, 4), # K3 GQA (dspark draft attention)
|
||||
}
|
||||
|
||||
|
||||
@@ -70,6 +71,8 @@ def _verify_mla_prefix_stage1(
|
||||
stride_qh,
|
||||
stride_buf_kbs,
|
||||
stride_buf_vbs,
|
||||
stride_buf_kh,
|
||||
stride_buf_vh,
|
||||
stride_ob,
|
||||
stride_oh,
|
||||
stride_os,
|
||||
@@ -80,6 +83,8 @@ def _verify_mla_prefix_stage1(
|
||||
H_Q: tl.constexpr,
|
||||
L_EXT: tl.constexpr,
|
||||
BLOCK_H: tl.constexpr,
|
||||
KV_GROUP_NUM: tl.constexpr,
|
||||
HAS_KV_HEADS: tl.constexpr,
|
||||
NOPE_DIM: tl.constexpr,
|
||||
PE_DIM: tl.constexpr,
|
||||
V_HEAD_DIM: tl.constexpr,
|
||||
@@ -103,6 +108,16 @@ def _verify_mla_prefix_stage1(
|
||||
if split_kv_id < active:
|
||||
head_start = head_block * BLOCK_H
|
||||
offs_h = head_start + tl.arange(0, BLOCK_H)
|
||||
# The caller guarantees KV_GROUP_NUM % BLOCK_H == 0, so every query head
|
||||
# in this block maps to the same KV head and the K/V tiles stay 2D loads.
|
||||
# MLA has a single latent head; keeping the offset out of that path
|
||||
# entirely preserves the original address arithmetic in the inner loop.
|
||||
if HAS_KV_HEADS:
|
||||
kv_head_off_k = (head_start // KV_GROUP_NUM) * stride_buf_kh
|
||||
kv_head_off_v = (head_start // KV_GROUP_NUM) * stride_buf_vh
|
||||
else:
|
||||
kv_head_off_k = 0
|
||||
kv_head_off_v = 0
|
||||
offs_l = tl.arange(0, L_EXT)
|
||||
offs_dn = tl.arange(0, BLOCK_DNOPE)
|
||||
offs_dp = tl.arange(0, BLOCK_DPE)
|
||||
@@ -149,7 +164,7 @@ def _verify_mla_prefix_stage1(
|
||||
kv_loc = tl.load(
|
||||
kv_indices + cur_batch_kv_start_idx + offs_n, mask=n_mask, other=0
|
||||
)
|
||||
base = kv_loc[None, :] * stride_buf_kbs
|
||||
base = kv_loc[None, :] * stride_buf_kbs + kv_head_off_k
|
||||
k_nope = tl.load(
|
||||
K_Buffer + base + offs_dn[:, None],
|
||||
mask=(offs_dn[:, None] < NOPE_DIM) & n_mask[None, :],
|
||||
@@ -169,7 +184,10 @@ def _verify_mla_prefix_stage1(
|
||||
# MLA exposes its latent V through V_Buffer; ordinary shared-KV
|
||||
# attention has an independent V cache. Both use this same load.
|
||||
v = tl.load(
|
||||
V_Buffer + kv_loc[:, None] * stride_buf_vbs + offs_dv[None, :],
|
||||
V_Buffer
|
||||
+ kv_loc[:, None] * stride_buf_vbs
|
||||
+ kv_head_off_v
|
||||
+ offs_dv[None, :],
|
||||
mask=n_mask[:, None] & (offs_dv[None, :] < V_HEAD_DIM),
|
||||
other=0.0,
|
||||
)
|
||||
@@ -228,6 +246,8 @@ def _verify_mla_combine_stage2(
|
||||
stride_qh,
|
||||
stride_kebs,
|
||||
stride_vebs,
|
||||
stride_keh,
|
||||
stride_veh,
|
||||
stride_oobs,
|
||||
stride_ooh,
|
||||
L_EXT: tl.constexpr,
|
||||
@@ -236,9 +256,18 @@ def _verify_mla_combine_stage2(
|
||||
BLOCK_DMODEL: tl.constexpr,
|
||||
BLOCK_DV: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
KV_GROUP_NUM: tl.constexpr,
|
||||
HAS_KV_HEADS: tl.constexpr,
|
||||
IS_CAUSAL: tl.constexpr,
|
||||
):
|
||||
cur_batch = tl.program_id(0)
|
||||
cur_head = tl.program_id(1)
|
||||
if HAS_KV_HEADS:
|
||||
kv_head_off_ke = (cur_head // KV_GROUP_NUM) * stride_keh
|
||||
kv_head_off_ve = (cur_head // KV_GROUP_NUM) * stride_veh
|
||||
else:
|
||||
kv_head_off_ke = 0
|
||||
kv_head_off_ve = 0
|
||||
|
||||
offs_d = tl.arange(0, BLOCK_DMODEL)
|
||||
offs_dv = tl.arange(0, BLOCK_DV)
|
||||
@@ -294,13 +323,19 @@ def _verify_mla_combine_stage2(
|
||||
q = tl.load(
|
||||
Q + offs_q, mask=mask_l[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0
|
||||
).to(tl.float32)
|
||||
offs_ke = (cur_q_start + offs_l)[:, None] * stride_kebs + offs_d[None, :]
|
||||
offs_ke = (
|
||||
(cur_q_start + offs_l)[:, None] * stride_kebs + kv_head_off_ke + offs_d[None, :]
|
||||
)
|
||||
ke = tl.load(
|
||||
K_Extend + offs_ke,
|
||||
mask=mask_l[:, None] & (offs_d[None, :] < HEAD_DIM),
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
offs_ve = (cur_q_start + offs_l)[:, None] * stride_vebs + offs_dv[None, :]
|
||||
offs_ve = (
|
||||
(cur_q_start + offs_l)[:, None] * stride_vebs
|
||||
+ kv_head_off_ve
|
||||
+ offs_dv[None, :]
|
||||
)
|
||||
ve = tl.load(
|
||||
V_Extend + offs_ve,
|
||||
mask=mask_l[:, None] & (offs_dv[None, :] < V_HEAD_DIM),
|
||||
@@ -309,9 +344,14 @@ def _verify_mla_combine_stage2(
|
||||
|
||||
# scores[i,j] = q_i . k_j (i query, j key) -> [L_EXT, L_EXT]
|
||||
qk = tl.sum(q[:, None, :] * ke[None, :, :], 2) * sm_scale
|
||||
# causal among drafts: query i sees key j iff j <= i, and both valid
|
||||
causal = (offs_l[None, :] <= offs_l[:, None]) & mask_l[None, :] & mask_l[:, None]
|
||||
qk = tl.where(causal, qk, float("-inf"))
|
||||
# causal: query i sees key j iff j <= i
|
||||
# non-causal: query i sees all keys
|
||||
both_valid = mask_l[None, :] & mask_l[:, None]
|
||||
if IS_CAUSAL:
|
||||
vis = (offs_l[None, :] <= offs_l[:, None]) & both_valid
|
||||
else:
|
||||
vis = both_valid
|
||||
qk = tl.where(vis, qk, float("-inf"))
|
||||
m_d = tl.max(qk, 1)
|
||||
pd = tl.exp(qk - m_d[:, None])
|
||||
denom_d = tl.sum(pd, 1)
|
||||
@@ -348,8 +388,16 @@ class VerifyMLA:
|
||||
block_h=DEFAULT_BLOCK_H,
|
||||
block_n=DEFAULT_BLOCK_N,
|
||||
num_warps=DEFAULT_NUM_WARPS,
|
||||
kv_group_num=None,
|
||||
):
|
||||
self.h_q = h_q
|
||||
# MLA is the h_kv == 1 case (kv_group_num == h_q); a GQA draft passes a
|
||||
# smaller group. block_h must divide it so a head block maps to one KV
|
||||
# head -- can_handle enforces that before this is constructed.
|
||||
self.kv_group_num = h_q if kv_group_num is None else kv_group_num
|
||||
# h_kv == 1 (MLA / MQA) needs no KV-head offset at all; making that a
|
||||
# constexpr keeps the inner-loop addressing identical to the original.
|
||||
self.has_kv_heads = self.kv_group_num < h_q
|
||||
self.head_dim = head_dim
|
||||
self.v_head_dim = v_head_dim
|
||||
self.nope_dim = v_head_dim
|
||||
@@ -419,6 +467,8 @@ class VerifyMLA:
|
||||
q_extend.stride(1),
|
||||
k_buffer.stride(0),
|
||||
v_buffer.stride(0),
|
||||
k_buffer.stride(1),
|
||||
v_buffer.stride(1),
|
||||
self.att_out.stride(0),
|
||||
self.att_out.stride(1),
|
||||
self.att_out.stride(2),
|
||||
@@ -429,6 +479,8 @@ class VerifyMLA:
|
||||
H_Q=self.h_q,
|
||||
L_EXT=self.l_pad,
|
||||
BLOCK_H=self.block_h,
|
||||
KV_GROUP_NUM=self.kv_group_num,
|
||||
HAS_KV_HEADS=self.has_kv_heads,
|
||||
NOPE_DIM=self.nope_dim,
|
||||
PE_DIM=self.pe_dim,
|
||||
V_HEAD_DIM=self.v_head_dim,
|
||||
@@ -452,6 +504,7 @@ class VerifyMLA:
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
sm_scale,
|
||||
is_causal=True,
|
||||
):
|
||||
grid = (bs, self.h_q)
|
||||
_verify_mla_combine_stage2[grid](
|
||||
@@ -476,6 +529,8 @@ class VerifyMLA:
|
||||
q_extend.stride(1),
|
||||
k_extend.stride(0),
|
||||
v_extend.stride(0),
|
||||
k_extend.stride(1),
|
||||
v_extend.stride(1),
|
||||
o_out.stride(0),
|
||||
o_out.stride(1),
|
||||
L_EXT=self.l_pad,
|
||||
@@ -484,6 +539,9 @@ class VerifyMLA:
|
||||
BLOCK_DMODEL=triton.next_power_of_2(self.head_dim),
|
||||
BLOCK_DV=triton.next_power_of_2(self.v_head_dim),
|
||||
BLOCK_N=self.block_n,
|
||||
KV_GROUP_NUM=self.kv_group_num,
|
||||
HAS_KV_HEADS=self.has_kv_heads,
|
||||
IS_CAUSAL=is_causal,
|
||||
num_warps=4,
|
||||
num_stages=1,
|
||||
)
|
||||
@@ -502,6 +560,7 @@ class VerifyMLA:
|
||||
o_out=None,
|
||||
k_scale=1.0,
|
||||
v_scale=1.0,
|
||||
is_causal=True,
|
||||
):
|
||||
if o_out is None:
|
||||
o_out = torch.empty(
|
||||
@@ -535,6 +594,7 @@ class VerifyMLA:
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
sm_scale,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
return o_out
|
||||
|
||||
@@ -542,11 +602,20 @@ class VerifyMLA:
|
||||
_VMLA_CACHE = {}
|
||||
|
||||
|
||||
def _get_vmla(max_bs, h_q, head_dim, v_head_dim, l_ext, device):
|
||||
key = (h_q, head_dim, v_head_dim, l_ext, str(device))
|
||||
def _get_vmla(max_bs, h_q, head_dim, v_head_dim, l_ext, device, kv_group_num=None):
|
||||
key = (h_q, head_dim, v_head_dim, l_ext, str(device), kv_group_num)
|
||||
vk = _VMLA_CACHE.get(key)
|
||||
if vk is None:
|
||||
block_h, block_n, num_warps = block_config(head_dim)
|
||||
if (
|
||||
kv_group_num is not None
|
||||
and kv_group_num < h_q # i.e. h_kv > 1, so the offset is live
|
||||
and kv_group_num % block_h != 0
|
||||
):
|
||||
# Shrink the head block so it stays inside one KV head. can_handle
|
||||
# only admits power-of-two groups in that case, so this stays valid
|
||||
# for tl.arange(0, BLOCK_H).
|
||||
block_h = kv_group_num
|
||||
vk = VerifyMLA(
|
||||
max_bs,
|
||||
h_q,
|
||||
@@ -557,6 +626,7 @@ def _get_vmla(max_bs, h_q, head_dim, v_head_dim, l_ext, device):
|
||||
block_h=block_h,
|
||||
block_n=block_n,
|
||||
num_warps=num_warps,
|
||||
kv_group_num=kv_group_num,
|
||||
)
|
||||
_VMLA_CACHE[key] = vk
|
||||
else:
|
||||
@@ -588,8 +658,9 @@ def can_handle(
|
||||
baseline.
|
||||
|
||||
IMPORTANT: ``custom_mask`` is intentionally NOT inspected (its values can't
|
||||
be read inside a captured HIP graph without a host sync). The kernel always
|
||||
computes pure-causal attention, which equals the tree mask ONLY at
|
||||
be read inside a captured HIP graph without a host sync). Every draft query
|
||||
sees the whole committed prefix and, when ``is_causal``, only its own
|
||||
predecessors among the draft tokens -- that is the tree mask ONLY at
|
||||
speculative topk == 1. The caller therefore MUST gate enablement on topk == 1
|
||||
(TritonAttnBackend does: ``use_verify_shared_kv = ... and self.topk == 1``).
|
||||
At topk > 1 the tree is not causal and this path must stay disabled."""
|
||||
@@ -602,7 +673,7 @@ def can_handle(
|
||||
return False
|
||||
if xai_temperature_len is not None and xai_temperature_len > 0:
|
||||
return False
|
||||
if not is_causal:
|
||||
if not is_causal and custom_mask is not None:
|
||||
return False
|
||||
# q layout must be [tokens, H_Q, D]; head dims handled by power-of-2 pad.
|
||||
if q_extend.dim() != 3 or k_extend.dim() != 3 or v_extend.dim() != 3:
|
||||
@@ -612,6 +683,9 @@ def can_handle(
|
||||
h_kv = k_extend.shape[1]
|
||||
if h_kv == 0 or h_q % h_kv != 0:
|
||||
return False
|
||||
kv_group_num = h_q // h_kv
|
||||
if h_kv > 1 and (kv_group_num & (kv_group_num - 1)) != 0:
|
||||
return False
|
||||
# head dims must match buffers.
|
||||
if k_buffer.shape[1] != h_kv or v_buffer.shape[1] != h_kv:
|
||||
return False
|
||||
@@ -676,7 +750,8 @@ def verify_shared_kv_fwd(
|
||||
"""
|
||||
Grouped-head drop-in for extend_attention_fwd on a topk==1 target-verify
|
||||
shape. Returns True if it ran (o_extend written), False if unsupported
|
||||
(caller falls back). Requires exactly one TP-local KV head.
|
||||
(caller falls back). A single TP-local KV head needs no offset at all;
|
||||
several are served when the group size is a power of two.
|
||||
"""
|
||||
if not can_handle(
|
||||
q_extend,
|
||||
@@ -697,8 +772,6 @@ def verify_shared_kv_fwd(
|
||||
xai_temperature_len=xai_temperature_len,
|
||||
):
|
||||
return False
|
||||
if k_extend.shape[1] != 1:
|
||||
return False
|
||||
if q_extend.shape[2] < v_extend.shape[2]:
|
||||
return False
|
||||
if kv_indices.numel() == 0:
|
||||
@@ -706,9 +779,11 @@ def verify_shared_kv_fwd(
|
||||
|
||||
bs = qo_indptr.shape[0] - 1
|
||||
h_q = q_extend.shape[1]
|
||||
h_kv = k_extend.shape[1]
|
||||
head_dim = q_extend.shape[2]
|
||||
v_head_dim = v_extend.shape[2]
|
||||
l_ext = int(max_len_extend)
|
||||
kv_group_num = h_q // h_kv
|
||||
|
||||
if sm_scale is None:
|
||||
sm_scale = 1.0 / (head_dim**0.5)
|
||||
@@ -723,7 +798,9 @@ def verify_shared_kv_fwd(
|
||||
|
||||
if max_bs is None or max_bs < bs:
|
||||
max_bs = bs
|
||||
vk = _get_vmla(max_bs, h_q, head_dim, v_head_dim, l_ext, q_extend.device)
|
||||
vk = _get_vmla(
|
||||
max_bs, h_q, head_dim, v_head_dim, l_ext, q_extend.device, kv_group_num
|
||||
)
|
||||
vk(
|
||||
q_extend,
|
||||
k_extend.contiguous(),
|
||||
@@ -737,5 +814,6 @@ def verify_shared_kv_fwd(
|
||||
o_out=o_extend,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
return True
|
||||
|
||||
@@ -125,6 +125,10 @@ def is_kimi_k3(config) -> bool:
|
||||
return _hf_arch(config) == "KimiK3ForConditionalGeneration"
|
||||
|
||||
|
||||
def is_dspark_draft(config) -> bool:
|
||||
return _hf_arch(config) == "DSparkDraftModel"
|
||||
|
||||
|
||||
def is_qwen3_5(config) -> bool:
|
||||
return _hf_arch(config) in (
|
||||
"Qwen3_5ForConditionalGeneration",
|
||||
|
||||
@@ -11,7 +11,12 @@ from sglang.kernels.ops.kvcache.kv_indices import (
|
||||
create_flashinfer_kv_indices_triton,
|
||||
)
|
||||
from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||
from sglang.srt.configs.model_config import AttentionArch, is_kimi_k3, is_qwen3_5
|
||||
from sglang.srt.configs.model_config import (
|
||||
AttentionArch,
|
||||
is_dspark_draft,
|
||||
is_kimi_k3,
|
||||
is_qwen3_5,
|
||||
)
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
)
|
||||
@@ -81,6 +86,10 @@ def _should_use_verify_shared_kv(model_config, topk, use_mla, use_verify_splitkv
|
||||
return False
|
||||
if use_mla:
|
||||
return is_kimi_k3(model_config.hf_config)
|
||||
if is_dspark_draft(model_config.hf_config):
|
||||
# Added for the K3 DSpark draft model, which is qwen3 type attention,
|
||||
# and using bidirectional (non-causal) mode.
|
||||
return use_verify_splitkv
|
||||
return (
|
||||
use_verify_splitkv
|
||||
and is_qwen3_5(model_config.hf_config)
|
||||
|
||||
Reference in New Issue
Block a user