submission 646222
sizezheng_94252 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 265 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-646222?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:4b29be9f65bb122cbb80b18c71a3637c364f5a13d8d4c401998eb60fad7a977f
license declaredunknown
license concludedunknown
authorssizezheng_94252
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
sc = tl.dot(q_1, tl.trans(kv1f)) + tl.dot(q_2, tl.trans(kv2f))online-softmax
m_new = tl.maximum(m_i, tl.max(sc, axis=1))persistent-kernel
batch_size = tl.num_programs(1) // NUM_HEAD_BLOCKSsplit-k
Split-KV flash attention with GQA (16 query heads, 1 KV head).Kernel source
submission.py265 lines
"""
Triton MLA decode kernel for MI355X.
Split-KV flash attention with GQA (16 query heads, 1 KV head).
FP8 KV, bf16 Q, online softmax, BLOCK_KV=32.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
@triton.jit
def _mla_splitkv_kernel(
Q, KV, kv_indptr, kv_scale_ptr,
Out, Partial_O, Partial_LSE,
sm_scale: tl.constexpr,
num_splits: tl.constexpr,
BLOCK_KV: tl.constexpr,
V_DIM: tl.constexpr,
QK_DIM: tl.constexpr,
HEADS_PER_BLOCK: tl.constexpr,
NUM_HEAD_BLOCKS: tl.constexpr,
TOTAL_HEADS: tl.constexpr,
SINGLE_SPLIT: tl.constexpr,
):
"""
Grid: (num_splits, batch_size * NUM_HEAD_BLOCKS)
Each block: HEADS_PER_BLOCK heads, one batch, one KV split.
"""
split_id = tl.program_id(0)
bh_id = tl.program_id(1)
batch_size = tl.num_programs(1) // NUM_HEAD_BLOCKS
batch_id = bh_id // NUM_HEAD_BLOCKS
hb_id = bh_id % NUM_HEAD_BLOCKS
h_off = hb_id * HEADS_PER_BLOCK
kv_start = tl.load(kv_indptr + batch_id)
kv_end = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_end - kv_start
tps = tl.cdiv(kv_len, num_splits)
my_start = kv_start + split_id * tps
my_end = tl.minimum(kv_start + (split_id + 1) * tps, kv_end)
my_len = my_end - my_start
kv_sc = tl.load(kv_scale_ptr).to(tl.float32)
q_base = batch_id * (TOTAL_HEADS * QK_DIM) + h_off * QK_DIM
h_ids = tl.arange(0, HEADS_PER_BLOCK)
d1 = tl.arange(0, 512)
d2 = 512 + tl.arange(0, 64)
q_1 = tl.load(Q + q_base + h_ids[:, None] * QK_DIM + d1[None, :])
q_2 = tl.load(Q + q_base + h_ids[:, None] * QK_DIM + d2[None, :])
m_i = tl.full([HEADS_PER_BLOCK], value=float("-inf"), dtype=tl.float32)
l_i = tl.zeros([HEADS_PER_BLOCK], dtype=tl.float32)
acc = tl.zeros([HEADS_PER_BLOCK, V_DIM], dtype=tl.float32)
v_ids = tl.arange(0, V_DIM)
for bi in range(tl.cdiv(my_len, BLOCK_KV)):
off = bi * BLOCK_KV + tl.arange(0, BLOCK_KV)
mask = off < my_len
g = my_start + off
kv1 = tl.load(KV + g[:, None] * QK_DIM + d1[None, :], mask=mask[:, None], other=0.0)
kv1f = (kv1.to(tl.float32) * kv_sc).to(tl.bfloat16)
kv2 = tl.load(KV + g[:, None] * QK_DIM + d2[None, :], mask=mask[:, None], other=0.0)
kv2f = (kv2.to(tl.float32) * kv_sc).to(tl.bfloat16)
sc = tl.dot(q_1, tl.trans(kv1f)) + tl.dot(q_2, tl.trans(kv2f))
sc = sc.to(tl.float32) * sm_scale
sc = tl.where(mask[None, :], sc, float("-inf"))
m_new = tl.maximum(m_i, tl.max(sc, axis=1))
alpha = tl.exp(m_i - m_new)
p = tl.exp(sc - m_new[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
acc = acc * alpha[:, None]
acc += tl.dot(p.to(tl.bfloat16), kv1f).to(tl.float32)
m_i = m_new
safe_l = tl.where(l_i > 0, l_i, 1.0)
acc = acc / safe_l[:, None]
if SINGLE_SPLIT == 1:
out_base = (batch_id * TOTAL_HEADS + h_off) * V_DIM
tl.store(Out + out_base + h_ids[:, None] * V_DIM + v_ids[None, :], acc.to(tl.bfloat16))
else:
po_base = (split_id * batch_size + batch_id) * (TOTAL_HEADS * V_DIM) + h_off * V_DIM
tl.store(Partial_O + po_base + h_ids[:, None] * V_DIM + v_ids[None, :], acc)
lse = m_i + tl.log(tl.where(l_i > 0, l_i, 1.0))
lse = tl.where(l_i > 0, lse, float("-inf"))
lse_base = (split_id * batch_size + batch_id) * TOTAL_HEADS + h_off
tl.store(Partial_LSE + lse_base + h_ids, lse)
@triton.jit
def _mla_reduce_kernel(
Partial_O, Partial_LSE, Output,
batch_size: tl.constexpr, num_splits: tl.constexpr,
V_DIM: tl.constexpr, TOTAL_HEADS: tl.constexpr,
):
pid = tl.program_id(0)
bid = pid // TOTAL_HEADS
hid = pid % TOTAL_HEADS
gm = float("-inf")
for s in tl.static_range(0, 64):
if s < num_splits:
gm = tl.maximum(gm, tl.load(Partial_LSE + (s * batch_size + bid) * TOTAL_HEADS + hid))
v_ids = tl.arange(0, V_DIM)
acc = tl.zeros([V_DIM], dtype=tl.float32)
tw = 0.0
for s in tl.static_range(0, 64):
if s < num_splits:
idx = (s * batch_size + bid) * TOTAL_HEADS + hid
w = tl.exp(tl.load(Partial_LSE + idx) - gm)
tw += w
po = (s * batch_size + bid) * (TOTAL_HEADS * V_DIM) + hid * V_DIM
acc += w * tl.load(Partial_O + po + v_ids)
tl.store(Output + (bid * TOTAL_HEADS + hid) * V_DIM + v_ids, (acc / tw).to(tl.bfloat16))
# ============================================================================
# Triton kernel caches
_triton_cache = {}
def _get_triton_cached(batch_size, num_splits):
key = (batch_size, num_splits)
if key in _triton_cache:
return _triton_cache[key]
if num_splits > 1:
po = torch.empty(num_splits * batch_size * NUM_HEADS * V_HEAD_DIM, dtype=torch.float32, device="cuda")
pl = torch.empty(num_splits * batch_size * NUM_HEADS, dtype=torch.float32, device="cuda")
else:
po = torch.empty(1, dtype=torch.float32, device="cuda")
pl = torch.empty(1, dtype=torch.float32, device="cuda")
out = torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
_triton_cache[key] = (po, pl, out)
return po, pl, out
# Triton split table (for small shapes where Triton beats aiter)
_TRITON_SPLIT_TABLE = {
(4, 1024): 8, (4, 8192): 32,
(32, 1024): 8, (32, 8192): 8,
(64, 1024): 4,
}
# Aiter path caches
_aiter_cache = {}
PAGE_SIZE = 1
FP8_DTYPE = None
def _aiter_mla(q, kv_fp8, kv_scale, kv_indptr, batch_size, kv_seq_len):
global FP8_DTYPE
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
import aiter as _aiter
if FP8_DTYPE is None:
FP8_DTYPE = aiter_dtypes.fp8
total_kv_len = batch_size * kv_seq_len
kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, 1, kv_fp8.shape[-1])
key = (batch_size, kv_seq_len)
if key not in _aiter_cache:
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
if batch_size >= 256:
num_kv_splits = 32
elif batch_size >= 64:
num_kv_splits = 16
else:
num_kv_splits = 8
q_fp8_buf = torch.empty((batch_size, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
o = torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
info = get_mla_metadata_info_v1(
batch_size, 1, NUM_HEADS, FP8_DTYPE, kv_fp8.dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=num_kv_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
_aiter_cache[key] = (kv_indices, num_kv_splits, work, o, q_fp8_buf, q_scale_buf)
kv_indices, num_kv_splits, work, o, q_fp8, q_scale = _aiter_cache[key]
(wm, wi, wis, ri, rfm, rpm) = work
_aiter.dynamic_per_tensor_quant(q_fp8, q, q_scale)
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
# Rebuild indptr-dependent metadata
from aiter import get_mla_metadata_v1
get_mla_metadata_v1(
torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda"), # qo_indptr
kv_indptr, kv_last_page_len,
NUM_HEADS, 1, True,
wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=False,
max_split_per_batch=num_kv_splits, intra_batch_mode=True,
dtype_q=FP8_DTYPE, dtype_kv=kv_fp8.dtype,
)
mla_decode_fwd(
q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d, o,
torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda"),
kv_indptr, kv_indices, kv_last_page_len, 1,
page_size=PAGE_SIZE, nhead_kv=1, sm_scale=SM_SCALE,
logit_cap=0.0, num_kv_splits=num_kv_splits,
q_scale=q_scale, kv_scale=kv_scale, intra_batch_mode=True,
work_meta_data=wm, work_indptr=wi, work_info_set=wis,
reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
)
return o
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
kv_fp8, kv_scale = kv_data["fp8"]
# Hybrid: Triton for small shapes (faster due to zero overhead),
# aiter assembly for large shapes (better bandwidth utilization)
key = (batch_size, kv_seq_len)
if key in _TRITON_SPLIT_TABLE:
# Triton path
kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
num_splits = _TRITON_SPLIT_TABLE[key]
po, pl, out = _get_triton_cached(batch_size, num_splits)
single_split = 1 if num_splits == 1 else 0
_mla_splitkv_kernel[(num_splits, batch_size)](
q, kv_flat, kv_indptr, kv_scale,
out, po, pl,
sm_scale=SM_SCALE, num_splits=num_splits, BLOCK_KV=32,
V_DIM=V_HEAD_DIM, QK_DIM=QK_HEAD_DIM,
HEADS_PER_BLOCK=16, NUM_HEAD_BLOCKS=1,
TOTAL_HEADS=NUM_HEADS, SINGLE_SPLIT=single_split,
)
if num_splits > 1:
_mla_reduce_kernel[(batch_size * NUM_HEADS,)](
po, pl, out,
batch_size=batch_size, num_splits=num_splits,
V_DIM=V_HEAD_DIM, TOTAL_HEADS=NUM_HEADS,
)
return out
else:
# Aiter assembly path for large shapes
return _aiter_mla(q, kv_fp8, kv_scale, kv_indptr, batch_size, kv_seq_len)
scrolls · 265 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON