[AMD] Improve K3 dspark draft attn kernel perf (#35499)

This commit is contained in:
Thomas Wang
2026-08-20 21:48:59 -07:00
committed by GitHub
parent bda9952377
commit 34180a0d35
3 changed files with 108 additions and 17 deletions
@@ -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)