[AMD] Add K3 verified mla kernel for DSpark on triton backend (#33981)

This commit is contained in:
Thomas Wang
2026-08-07 22:59:26 -07:00
committed by GitHub
parent b61a06921e
commit 6679d9b60c
2 changed files with 671 additions and 3 deletions
@@ -0,0 +1,647 @@
"""
MLA split-KV attention for EAGLE/DSpark speculative *verify* (topk==1).
Following the pattern of ``python/sglang/kernels/ops/attention/verify_splitkv.py``.
Grid is ``(bs, n_head_blocks, split)``; each program handles ``BLOCK_H`` query
heads x ALL ``L_EXT`` draft queries.
Correctness matches ``extend_attention_fwd`` for the topk==1 causal verify case.
score(h,i,t) = q_nope[i,h] · c_KV[t] + q_pe[i,h] · k_pe[t] # 512 dot + 64 dot
out(h,i) = Σ_t softmax_t · c_KV[t] # V = c_KV
"""
import torch
import triton
import triton.language as tl
from sglang.kernels.ops.attention.verify_splitkv import (
_AMD_LAUNCH_KWARGS,
can_handle,
)
MAX_N_SPLITS = 32 # Grid split dim upper bound
TARGET_PROGRAMS = 512 # Target total stage-1 programs
DEFAULT_BLOCK_H = (
4 # BLOCK_H must be a power of 2 (tl.arange); heads beyond H_Q are masked.
)
DEFAULT_BLOCK_N = 64
DEFAULT_NUM_WARPS = 8
_BLOCK_CONFIG = {
# head_dim: (BLOCK_H, BLOCK_N, num_warps)
576: (4, 64, 8), # K3 MLA (kv_lora_rank 512 + qk_rope 64)
}
def block_config(head_dim):
"""
Return (BLOCK_H, BLOCK_N, num_warps) for a head_dim; default for untuned
dims. BLOCK_H must be a power of 2 (heads beyond H_Q are masked).
"""
return _BLOCK_CONFIG.get(
head_dim, (DEFAULT_BLOCK_H, DEFAULT_BLOCK_N, DEFAULT_NUM_WARPS)
)
@triton.jit
def _active_splits(seqlen, num_splits, BLOCK_N: tl.constexpr):
"""
The launched split count, capped by ``seqlen // BLOCK_N``.
Floor partitioning keeps every active split non-empty.
"""
return tl.maximum(1, tl.minimum(num_splits, seqlen // BLOCK_N))
@triton.jit
def _verify_mla_prefix_stage1(
Q,
K_Buffer,
V_Buffer,
sm_scale,
k_scale,
v_scale,
qo_indptr,
kv_indptr,
kv_indices,
Att_Out, # [BS, H_Q, MAX_N_SPLITS, L_EXT, Dv] bf16
Att_Lse, # [BS, H_Q, MAX_N_SPLITS, L_EXT] fp32
num_splits, # launched split count
stride_qbs,
stride_qh,
stride_buf_kbs,
stride_buf_vbs,
stride_ob,
stride_oh,
stride_os,
stride_ol,
stride_lb,
stride_lh,
stride_ls,
H_Q: tl.constexpr,
L_EXT: tl.constexpr,
BLOCK_H: tl.constexpr,
NOPE_DIM: tl.constexpr,
PE_DIM: tl.constexpr,
V_HEAD_DIM: tl.constexpr,
BLOCK_DNOPE: tl.constexpr,
BLOCK_DPE: tl.constexpr,
BLOCK_DV: tl.constexpr,
BLOCK_N: tl.constexpr,
):
cur_batch = tl.program_id(0)
head_block = tl.program_id(1)
split_kv_id = tl.program_id(2)
# row tile size for each workgroup
R: tl.constexpr = BLOCK_H * L_EXT
cur_batch_kv_start_idx = tl.load(kv_indptr + cur_batch)
cur_batch_seq_len = tl.load(kv_indptr + cur_batch + 1) - cur_batch_kv_start_idx
active = _active_splits(cur_batch_seq_len, num_splits, BLOCK_N)
# skip idle workgroups
if split_kv_id < active:
head_start = head_block * BLOCK_H
offs_h = head_start + tl.arange(0, BLOCK_H)
offs_l = tl.arange(0, L_EXT)
offs_dn = tl.arange(0, BLOCK_DNOPE)
offs_dp = tl.arange(0, BLOCK_DPE)
offs_dv = tl.arange(0, BLOCK_DV)
cur_q_start = tl.load(qo_indptr + cur_batch)
l_ext = tl.load(qo_indptr + cur_batch + 1) - cur_q_start
row_mask = tl.reshape((offs_h[:, None] < H_Q) & (offs_l[None, :] < l_ext), (R,))
# the last kv split completes the remaining
kv_len_per_split = cur_batch_seq_len // active
split_start = kv_len_per_split * split_kv_id
split_end = tl.where(
split_kv_id == active - 1, cur_batch_seq_len, split_start + kv_len_per_split
)
# load q_nope and q_pe
q_row = tl.reshape(
(cur_q_start + offs_l)[None, :] * stride_qbs + offs_h[:, None] * stride_qh,
(R,),
)
q_nope = tl.load(
Q + q_row[:, None] + offs_dn[None, :],
mask=row_mask[:, None] & (offs_dn[None, :] < NOPE_DIM),
other=0.0,
).to(K_Buffer.dtype.element_ty)
q_pe = tl.load(
Q + q_row[:, None] + (NOPE_DIM + offs_dp)[None, :],
mask=row_mask[:, None] & (offs_dp[None, :] < PE_DIM),
other=0.0,
).to(K_Buffer.dtype.element_ty)
e_max = tl.zeros([R], dtype=tl.float32) - float("inf")
e_sum = tl.zeros([R], dtype=tl.float32)
acc = tl.zeros([R, BLOCK_DV], dtype=tl.float32)
# calculate attention scores for each kv split
for start_n in tl.range(split_start, split_end, BLOCK_N):
offs_n = start_n + tl.arange(0, BLOCK_N)
n_mask = offs_n < split_end
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
k_nope = tl.load(
K_Buffer + base + offs_dn[:, None],
mask=(offs_dn[:, None] < NOPE_DIM) & n_mask[None, :],
other=0.0,
)
k_pe = tl.load(
K_Buffer + base + (NOPE_DIM + offs_dp)[:, None],
mask=(offs_dp[:, None] < PE_DIM) & n_mask[None, :],
other=0.0,
)
qk = tl.dot(q_nope, k_nope) + tl.dot(q_pe, k_pe)
qk *= sm_scale * k_scale
qk = tl.where(n_mask[None, :], qk, float("-inf"))
# V is the same as k_nope but transposed; tl.trans is slow, so re-load
# V instead of reusing k_nope.
v = tl.load(
V_Buffer + kv_loc[:, None] * stride_buf_vbs + offs_dv[None, :],
mask=n_mask[:, None] & (offs_dv[None, :] < V_HEAD_DIM),
other=0.0,
)
n_e_max = tl.maximum(tl.max(qk, 1), e_max)
re_scale = tl.exp(e_max - n_e_max)
p = tl.exp(qk - n_e_max[:, None])
acc *= re_scale[:, None]
acc += tl.dot(p.to(v.dtype), v)
e_sum = e_sum * re_scale + tl.sum(p, 1)
e_max = n_e_max
# fp8 dequant of prefix V: scale the accumulated (pre-normalised) output.
acc *= v_scale
o_row = tl.reshape(
cur_batch * stride_ob
+ offs_h[:, None] * stride_oh
+ split_kv_id * stride_os
+ offs_l[None, :] * stride_ol,
(R,),
)
tl.store(
Att_Out + o_row[:, None] + offs_dv[None, :],
(acc / e_sum[:, None]).to(Att_Out.dtype.element_ty),
mask=row_mask[:, None] & (offs_dv[None, :] < V_HEAD_DIM),
)
lse_row = tl.reshape(
cur_batch * stride_lb
+ offs_h[:, None] * stride_lh
+ split_kv_id * stride_ls
+ offs_l[None, :],
(R,),
)
tl.store(Att_Lse + lse_row, e_max + tl.log(e_sum), mask=row_mask)
@triton.jit
def _verify_mla_combine_stage2(
Att_Out,
Att_Lse,
Q,
K_Extend,
V_Extend,
O_Out,
sm_scale,
qo_indptr,
kv_indptr,
num_splits,
stride_ob,
stride_oh,
stride_os,
stride_ol,
stride_lb,
stride_lh,
stride_ls,
stride_qbs,
stride_qh,
stride_kebs,
stride_vebs,
stride_oobs,
stride_ooh,
L_EXT: tl.constexpr,
HEAD_DIM: tl.constexpr,
V_HEAD_DIM: tl.constexpr,
BLOCK_DMODEL: tl.constexpr,
BLOCK_DV: tl.constexpr,
BLOCK_N: tl.constexpr,
):
cur_batch = tl.program_id(0)
cur_head = tl.program_id(1)
offs_d = tl.arange(0, BLOCK_DMODEL)
offs_dv = tl.arange(0, BLOCK_DV)
offs_l = tl.arange(0, L_EXT)
cur_q_start = tl.load(qo_indptr + cur_batch)
l_ext = tl.load(qo_indptr + cur_batch + 1) - cur_q_start
mask_l = offs_l < l_ext
# ---- (a) combine prefix splits (online logsumexp over active splits) ---
seqlen = tl.load(kv_indptr + cur_batch + 1) - tl.load(kv_indptr + cur_batch)
active = _active_splits(seqlen, num_splits, BLOCK_N)
m = tl.zeros([L_EXT], dtype=tl.float32) - float("inf")
l_acc = tl.zeros([L_EXT], dtype=tl.float32)
acc = tl.zeros([L_EXT, BLOCK_DV], dtype=tl.float32)
for s in range(active):
lse_s = tl.load(
Att_Lse
+ cur_batch * stride_lb
+ cur_head * stride_lh
+ s * stride_ls
+ offs_l,
mask=mask_l,
other=float("-inf"),
)
o_s = tl.load(
Att_Out
+ cur_batch * stride_ob
+ cur_head * stride_oh
+ s * stride_os
+ offs_l[:, None] * stride_ol
+ offs_dv[None, :],
mask=mask_l[:, None] & (offs_dv[None, :] < V_HEAD_DIM),
other=0.0,
).to(tl.float32)
new_m = tl.maximum(m, lse_s)
alpha = tl.exp(m - new_m)
beta = tl.exp(lse_s - new_m)
acc = acc * alpha[:, None] + o_s * beta[:, None]
l_acc = l_acc * alpha + beta
m = new_m
o_prefix = acc / l_acc[:, None]
lse_prefix = m + tl.log(l_acc)
# ---- (b) draft-draft causal attention (L_EXT x L_EXT) -----------------
# load draft queries [L_EXT, D], draft K/V [L_EXT, D]/[L_EXT, Dv]
offs_q = (
(cur_q_start + offs_l)[:, None] * stride_qbs
+ cur_head * stride_qh
+ offs_d[None, :]
)
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, :]
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, :]
ve = tl.load(
V_Extend + offs_ve,
mask=mask_l[:, None] & (offs_dv[None, :] < V_HEAD_DIM),
other=0.0,
).to(tl.float32)
# 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"))
m_d = tl.max(qk, 1)
pd = tl.exp(qk - m_d[:, None])
denom_d = tl.sum(pd, 1)
o_draft = tl.sum(pd[:, :, None] * ve[None, :, :], 1) / denom_d[:, None]
lse_draft = m_d + tl.log(denom_d)
# ---- (c) final LSE merge (prefix vs draft) ----------------------------
mm = tl.maximum(lse_prefix, lse_draft)
wp = tl.exp(lse_prefix - mm)
wd = tl.exp(lse_draft - mm)
o = (o_prefix * wp[:, None] + o_draft * wd[:, None]) / (wp + wd)[:, None]
offs_oo = (
(cur_q_start + offs_l)[:, None] * stride_oobs
+ cur_head * stride_ooh
+ offs_dv[None, :]
)
tl.store(
O_Out + offs_oo,
o.to(O_Out.dtype.element_ty),
mask=mask_l[:, None] & (offs_dv[None, :] < V_HEAD_DIM),
)
class VerifyMLA:
def __init__(
self,
max_bs,
h_q,
head_dim,
v_head_dim,
l_ext,
device="cuda",
block_h=DEFAULT_BLOCK_H,
block_n=DEFAULT_BLOCK_N,
num_warps=DEFAULT_NUM_WARPS,
):
self.h_q = h_q
self.head_dim = head_dim
self.v_head_dim = v_head_dim
self.nope_dim = v_head_dim
self.pe_dim = head_dim - v_head_dim
self.l_ext = l_ext
self.l_pad = triton.next_power_of_2(l_ext)
self.device = device
self.block_h = block_h
self.block_n = block_n
self.num_warps = num_warps
self.n_head_blocks = triton.cdiv(h_q, block_h)
self._alloc(max_bs)
def _alloc(self, max_bs):
self.max_bs = max_bs
# bf16 partials (halves scratch traffic vs fp32); lse stays fp32.
self.att_out = torch.empty(
(max_bs, self.h_q, MAX_N_SPLITS, self.l_pad, self.v_head_dim),
dtype=torch.bfloat16,
device=self.device,
)
self.att_lse = torch.empty(
(max_bs, self.h_q, MAX_N_SPLITS, self.l_pad),
dtype=torch.float32,
device=self.device,
)
def grow_buffers(self, max_bs):
if max_bs > self.max_bs:
self._alloc(max_bs)
def _num_splits(self, bs):
budget = TARGET_PROGRAMS // max(1, bs * self.n_head_blocks)
return max(1, min(MAX_N_SPLITS, budget))
def _run_prefix_kernel(
self,
bs,
num_splits,
q_extend,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
sm_scale,
k_scale,
v_scale,
):
grid = (bs, self.n_head_blocks, num_splits)
_verify_mla_prefix_stage1[grid](
q_extend,
k_buffer,
v_buffer,
sm_scale,
k_scale,
v_scale,
qo_indptr,
kv_indptr,
kv_indices,
self.att_out,
self.att_lse,
num_splits,
q_extend.stride(0),
q_extend.stride(1),
k_buffer.stride(0),
v_buffer.stride(0),
self.att_out.stride(0),
self.att_out.stride(1),
self.att_out.stride(2),
self.att_out.stride(3),
self.att_lse.stride(0),
self.att_lse.stride(1),
self.att_lse.stride(2),
H_Q=self.h_q,
L_EXT=self.l_pad,
BLOCK_H=self.block_h,
NOPE_DIM=self.nope_dim,
PE_DIM=self.pe_dim,
V_HEAD_DIM=self.v_head_dim,
BLOCK_DNOPE=triton.next_power_of_2(self.nope_dim),
BLOCK_DPE=triton.next_power_of_2(self.pe_dim),
BLOCK_DV=triton.next_power_of_2(self.v_head_dim),
BLOCK_N=self.block_n,
num_warps=self.num_warps,
num_stages=1,
**_AMD_LAUNCH_KWARGS,
)
def _run_combine_kernel(
self,
bs,
num_splits,
q_extend,
k_extend,
v_extend,
o_out,
qo_indptr,
kv_indptr,
sm_scale,
):
grid = (bs, self.h_q)
_verify_mla_combine_stage2[grid](
self.att_out,
self.att_lse,
q_extend,
k_extend,
v_extend,
o_out,
sm_scale,
qo_indptr,
kv_indptr,
num_splits,
self.att_out.stride(0),
self.att_out.stride(1),
self.att_out.stride(2),
self.att_out.stride(3),
self.att_lse.stride(0),
self.att_lse.stride(1),
self.att_lse.stride(2),
q_extend.stride(0),
q_extend.stride(1),
k_extend.stride(0),
v_extend.stride(0),
o_out.stride(0),
o_out.stride(1),
L_EXT=self.l_pad,
HEAD_DIM=self.head_dim,
V_HEAD_DIM=self.v_head_dim,
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,
num_warps=4,
num_stages=1,
)
def __call__(
self,
q_extend,
k_extend,
v_extend,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
sm_scale,
o_out=None,
k_scale=1.0,
v_scale=1.0,
):
if o_out is None:
o_out = torch.empty(
(q_extend.shape[0], self.h_q, self.v_head_dim),
dtype=q_extend.dtype,
device=q_extend.device,
)
bs = qo_indptr.shape[0] - 1
# One split count for both stages (they must agree on the active count).
num_splits = self._num_splits(bs)
self._run_prefix_kernel(
bs,
num_splits,
q_extend,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
sm_scale,
k_scale,
v_scale,
)
self._run_combine_kernel(
bs,
num_splits,
q_extend,
k_extend,
v_extend,
o_out,
qo_indptr,
kv_indptr,
sm_scale,
)
return o_out
_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))
vk = _VMLA_CACHE.get(key)
if vk is None:
block_h, block_n, num_warps = block_config(head_dim)
vk = VerifyMLA(
max_bs,
h_q,
head_dim,
v_head_dim,
l_ext,
device=device,
block_h=block_h,
block_n=block_n,
num_warps=num_warps,
)
_VMLA_CACHE[key] = vk
else:
vk.grow_buffers(max_bs)
return vk
def verify_mla_fwd(
q_extend,
k_extend,
v_extend,
o_extend,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
custom_mask,
is_causal,
mask_indptr,
max_len_extend,
k_scale,
v_scale,
sm_scale=None,
logit_cap=0.0,
skip_prefix_custom_mask=True,
sliding_window_size=-1,
sinks=None,
window_kv_offsets=None,
xai_temperature_len=-1,
max_bs=None,
):
"""
MLA-native drop-in for extend_attention_fwd on the EAGLE target-verify
(topk==1) shape. Returns True if it ran (o_extend written), False if unsupported
(caller falls back). Requires h_kv == 1 (MLA single latent).
"""
if not can_handle(
q_extend,
k_extend,
v_extend,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
custom_mask,
is_causal,
mask_indptr,
max_len_extend,
sliding_window_size=sliding_window_size,
sinks=sinks,
logit_cap=logit_cap,
xai_temperature_len=xai_temperature_len,
):
return False
if k_extend.shape[1] != 1: # MLA: single shared latent head
return False
bs = qo_indptr.shape[0] - 1
h_q = q_extend.shape[1]
head_dim = q_extend.shape[2]
v_head_dim = v_extend.shape[2]
l_ext = int(max_len_extend)
if sm_scale is None:
sm_scale = 1.0 / (head_dim**0.5)
try:
k_scale = float(k_scale)
except (TypeError, ValueError):
k_scale = 1.0
try:
v_scale = float(v_scale)
except (TypeError, ValueError):
v_scale = 1.0
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(
q_extend,
k_extend.contiguous(),
v_extend.contiguous(),
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
sm_scale,
o_out=o_extend,
k_scale=k_scale,
v_scale=v_scale,
)
return True
@@ -15,7 +15,7 @@ from sglang.srt.configs.hybrid_arch import (
kimi_linear_config,
linear_attn_model_spec,
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.configs.model_config import AttentionArch, is_kimi_k3
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
@@ -137,6 +137,9 @@ class TritonAttnBackend(AttentionBackend):
extend_attention_fwd,
extend_attention_fwd_unified,
)
from sglang.kernels.ops.attention.verify_mla import (
verify_mla_fwd,
)
from sglang.kernels.ops.attention.verify_splitkv import (
verify_splitkv_fwd,
)
@@ -151,6 +154,8 @@ class TritonAttnBackend(AttentionBackend):
self.build_unified_kv_indices = torch.compiler.disable(build_unified_kv_indices)
# Split-KV EAGLE-verify kernel; enabled below once topk is known (valid only at topk == 1).
self.verify_splitkv_fwd = torch.compiler.disable(verify_splitkv_fwd)
# MLA split-KV EAGLE-verify kernel; enabled below once topk is known (valid only at topk == 1).
self.verify_mla_fwd = torch.compiler.disable(verify_mla_fwd)
# Parse args
self.skip_prefill = skip_prefill
@@ -183,6 +188,14 @@ class TritonAttnBackend(AttentionBackend):
and self.topk == 1
)
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
# The MLA verify kernel (verify_mla_fwd) is tuned and validated for the
# Kimi-K3 absorbed-MLA shape; gate it on K3.
self.use_verify_mla = (
is_gfx95_supported()
and self.topk == 1
and self.use_mla
and is_kimi_k3(model_runner.model_config.hf_config)
)
self.dcp_size = get_parallel().attn_dcp_size
self.dcp_rank = get_parallel().attn_dcp_rank
self.num_head = (
@@ -1396,11 +1409,19 @@ class TritonAttnBackend(AttentionBackend):
# serve bit-equivalently (its can_handle() gates on non-causal / sinks /
# sliding-window / ragged / topk>1), so we fall through to
# extend_attention_fwd below. Correctness is never at risk.
# Route target-verify to the K3-tuned MLA kernel when eligible, else the
# per-head split-KV kernel.
if self.use_verify_mla:
verify_fwd = self.verify_mla_fwd
elif self.use_verify_splitkv:
verify_fwd = self.verify_splitkv_fwd
else:
verify_fwd = None
if (
self.use_verify_splitkv
verify_fwd is not None
and score_mod is None
and forward_batch.forward_mode.is_target_verify()
and self.verify_splitkv_fwd(
and verify_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
k.contiguous(),
v.contiguous(),