submission 720383
div22 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 240 lines, June 9 Researcher Reciprocity License v1.0.
solution_373.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-720383?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:c1742990b53892ca2e57649afb82f211ccde6f1a2192eeba4c4d644f48a97dc0
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores = tl.dot(Q1, tl.trans(K1)) + tl.dot(Q2, tl.trans(K2))online-softmax
m_ij = tl.max(scores, axis=1); m_new = tl.maximum(m_i, m_ij)persistent-kernel
def _aiter_persistent(q, kv_data, qo_indptr, config, ps, nks=None, kg=16, fast_mode=False, ibm_override=None):tile-n = 32
BM = triton.next_power_of_2(nh); BN = 32; sq = nh * 576; sh = 576Kernel source
solution_373.py240 lines
"""
S370: S369 + Triton nw=4/nst=2 for bs>=32/kv=1K (s365 config).
s365 used nw=4/nst=2 for bs=32 and bs=64 → 24.6/37.2 vs default 26.7/40.7us.
"""
import math
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
NUM_KV_HEADS = 1
SM_SCALE = 1.0 / (576 ** 0.5)
KV_GRANULARITY = 16
CU_NUM = 256
FP8_DTYPE = torch.float8_e4m3fn
FP8_MAX = torch.finfo(FP8_DTYPE).max
def _auto_nks(bs, kv_seq_len):
avg_kv = kv_seq_len
best_eff = -1.0
best_nks = 1
for i in range(1, 17):
waves = math.ceil(bs * i / CU_NUM)
eff = (bs * i / waves) * CU_NUM * avg_kv / (avg_kv + 84.1 * i)
if eff > best_eff:
best_eff = eff
best_nks = i
return best_nks
_meta_cache = {}
_out_cache = {}
_indices_cache = {}
_kv_lpl_cache = {}
_kv_indptr_page_cache = {}
def _ci(n):
if n not in _indices_cache:
_indices_cache[n] = torch.arange(n, dtype=torch.int32, device="cuda")
return _indices_cache[n]
def _cl(bs, ps):
k = (bs, ps)
if k not in _kv_lpl_cache:
_kv_lpl_cache[k] = torch.full((bs,), ps, dtype=torch.int32, device="cuda")
return _kv_lpl_cache[k]
def _cp(bs, kvsl, ps):
k = (bs, kvsl, ps)
if k not in _kv_indptr_page_cache:
ppb = kvsl // ps
_kv_indptr_page_cache[k] = torch.arange(0, (bs + 1) * ppb, ppb, dtype=torch.int32, device="cuda")
return _kv_indptr_page_cache[k]
def _co(tq, nh, tag=""):
k = (tq, nh, tag)
if k not in _out_cache:
_out_cache[k] = torch.empty(tq, nh, 512, dtype=torch.bfloat16, device="cuda")
return _out_cache[k]
def _get_cached_meta(batch_size, q_seq, nhead, q_dtype, kv_dtype, page_size,
qo_indptr, kv_indptr_page, kv_lpl, nks, kv_seq_len, kg=16, fast_mode=False, ibm_override=None):
ibm = ibm_override if ibm_override is not None else (not fast_mode)
key = (batch_size, q_seq, nhead, q_dtype, kv_dtype, page_size, nks, kv_seq_len, kg, fast_mode, ibm)
if key not in _meta_cache:
info = get_mla_metadata_info_v1(
batch_size, q_seq, nhead, q_dtype, kv_dtype,
is_sparse=False, fast_mode=fast_mode,
num_kv_splits=nks, intra_batch_mode=ibm,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(
qo_indptr, kv_indptr_page, kv_lpl,
nhead // NUM_KV_HEADS, NUM_KV_HEADS, True,
wm, wis, wi, ri, rfm, rpm,
page_size=page_size, kv_granularity=max(page_size, kg),
max_seqlen_qo=q_seq, uni_seqlen_qo=q_seq,
fast_mode=fast_mode, max_split_per_batch=nks,
intra_batch_mode=ibm, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
_meta_cache[key] = work
return _meta_cache[key]
def _aiter_persistent(q, kv_data, qo_indptr, config, ps, nks=None, kg=16, fast_mode=False, ibm_override=None):
"""Persistent mode bf16 Q + fp8 KV."""
bs, nh, qsl, kvsl = config["batch_size"], config["num_heads"], config["q_seq_len"], config["kv_seq_len"]
tq, tkv = bs * qsl, bs * kvsl
if nks is None:
nks = _auto_nks(bs, kvsl)
np_ = tkv // ps
kv_fp8, kv_s = kv_data["fp8"]
kv_4d = kv_fp8.view(np_, ps, NUM_KV_HEADS, 576)
output = _co(tq, nh, "p")
ibm = ibm_override if ibm_override is not None else (not fast_mode)
wm, wi, wis, ri, rfm, rpm = _get_cached_meta(
bs, qsl, nh, q.dtype, kv_fp8.dtype, ps,
qo_indptr, _cp(bs, kvsl, ps), _cl(bs, ps), nks, kvsl, kg, fast_mode, ibm_override)
mla_decode_fwd(
q.view(-1, nh, 576), kv_4d, output,
qo_indptr, _cp(bs, kvsl, ps), _ci(np_), _cl(bs, ps), qsl,
page_size=ps, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,
num_kv_splits=nks, q_scale=None, kv_scale=kv_s, intra_batch_mode=ibm,
work_meta_data=wm, work_indptr=wi, work_info_set=wis,
reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
)
return output
# ============ TRITON (gfx950: BLOCK_N=32, nw=8, ns=3) ============
@triton.jit
def _mla_stage1(
Q_ptr, KV_ptr, PO_ptr, PM_ptr, PS_ptr, output_ptr, kv_indptr_ptr,
stride_q_seq, stride_q_head, total_q, num_heads: tl.constexpr,
q_seq_len, num_splits, sm_scale,
IS_FUSED: tl.constexpr, IS_FP8: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
pid_split = tl.program_id(0); qi = tl.program_id(1); bi = qi // q_seq_len
kv_s = tl.load(kv_indptr_ptr + bi); kv_e = tl.load(kv_indptr_ptr + bi + 1)
if IS_FUSED: cs = kv_s; ce = kv_e
else:
kv_len = kv_e - kv_s; chunk = tl.cdiv(kv_len, num_splits)
cs = kv_s + pid_split * chunk; ce = tl.minimum(cs + chunk, kv_e)
heads = tl.arange(0, BLOCK_M); d512 = tl.arange(0, 512); d64 = tl.arange(0, 64)
mask_h = heads < num_heads; q_base = qi * stride_q_seq
Q1 = tl.load(Q_ptr + q_base + heads[:, None] * stride_q_head + d512[None, :], mask=mask_h[:, None], other=0.0)
Q2 = tl.load(Q_ptr + q_base + heads[:, None] * stride_q_head + 512 + d64[None, :], mask=mask_h[:, None], other=0.0)
m_i = tl.full([BLOCK_M], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([BLOCK_M], dtype=tl.float32); acc = tl.zeros([BLOCK_M, 512], dtype=tl.float32)
for start_n in range(cs, ce, BLOCK_N):
n_offs = start_n + tl.arange(0, BLOCK_N); mask_n = n_offs < ce
if IS_FP8:
K1 = tl.load(KV_ptr + n_offs[:, None] * 576 + d512[None, :], mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
K2 = tl.load(KV_ptr + n_offs[:, None] * 576 + 512 + d64[None, :], mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
else:
K1 = tl.load(KV_ptr + n_offs[:, None] * 576 + d512[None, :], mask=mask_n[:, None], other=0.0)
K2 = tl.load(KV_ptr + n_offs[:, None] * 576 + 512 + d64[None, :], mask=mask_n[:, None], other=0.0)
scores = tl.dot(Q1, tl.trans(K1)) + tl.dot(Q2, tl.trans(K2))
scores = scores * sm_scale; scores = tl.where(mask_n[None, :], scores, float("-inf"))
m_ij = tl.max(scores, axis=1); m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2((m_i - m_new) * 1.44269504)
p = tl.math.exp2((scores - m_new[:, None]) * 1.44269504)
l_i = l_i * alpha + tl.sum(p, axis=1)
acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), K1); m_i = m_new
if IS_FUSED:
inv = 1.0 / tl.maximum(l_i, 1e-12); result = (acc * inv[:, None]).to(tl.bfloat16)
out_base = qi * num_heads
tl.store(output_ptr + (out_base + heads[:, None]) * 512 + d512[None, :], result, mask=mask_h[:, None])
else:
total_qh = total_q * num_heads; out_base = pid_split * total_qh + qi * num_heads
tl.store(PM_ptr + out_base + heads, m_i, mask=mask_h)
tl.store(PS_ptr + out_base + heads, l_i, mask=mask_h)
tl.store(PO_ptr + (out_base + heads[:, None]) * 512 + d512[None, :], acc, mask=mask_h[:, None])
@triton.jit
def _mla_reduce(PO_ptr, PM_ptr, PS_ptr, output_ptr, total_qh, out_scale, num_splits: tl.constexpr):
qh = tl.program_id(0); d512 = tl.arange(0, 512); gmax = float("-inf")
for s in range(num_splits): gmax = tl.maximum(gmax, tl.load(PM_ptr + s * total_qh + qh))
acc = tl.zeros([512], dtype=tl.float32); total_sum = 0.0
for s in range(num_splits):
idx = s * total_qh + qh
rs = tl.math.exp2((tl.load(PM_ptr + idx) - gmax) * 1.44269504)
total_sum += tl.load(PS_ptr + idx) * rs
acc += tl.load(PO_ptr + idx * 512 + d512) * rs
inv = 1.0 / tl.maximum(total_sum, 1e-12)
tl.store(output_ptr + qh * 512 + d512, (acc * inv * out_scale).to(tl.bfloat16))
_ws_cache = {}
def _triton_kernel(q, kv_data, qo_indptr, kv_indptr, config, nw=8, nst=3):
nh, qsl, kvsl, sm, bs = config["num_heads"], config["q_seq_len"], config["kv_seq_len"], config["sm_scale"], config["batch_size"]
tq = bs * qsl; tqh = tq * nh; tkv = bs * kvsl; dev = q.device
BM = triton.next_power_of_2(nh); BN = 32; sq = nh * 576; sh = 576
ns = max(1, min(32, kvsl // 64, -(-CU_NUM // tq)))
use_fp8 = tkv > 100000
if use_fp8:
kf, ks = kv_data["fp8"]; kv = kf.view(-1, 576); kvs = ks.item(); es = sm * kvs; os_val = kvs
else:
kv = kv_data["bf16"].view(-1, 576); es = sm; os_val = 1.0
output = _co(tq, nh, "t")
if ns == 1:
_mla_stage1[(1, tq)](q, kv, None, None, None, output, kv_indptr,
sq, sh, tq, nh, qsl, 1, es, True, use_fp8, BM, BN, num_warps=nw, num_stages=nst)
else:
wk = (ns, tqh)
if wk not in _ws_cache:
n = ns * tqh
_ws_cache[wk] = (torch.empty(n, 512, dtype=torch.float32, device=dev),
torch.empty(n, dtype=torch.float32, device=dev),
torch.empty(n, dtype=torch.float32, device=dev))
po, pm, ps = _ws_cache[wk]
_mla_stage1[(ns, tq)](q, kv, po, pm, ps, None, kv_indptr,
sq, sh, tq, nh, qsl, ns, es, False, use_fp8, BM, BN, num_warps=nw, num_stages=nst)
_mla_reduce[(tqh,)](po, pm, ps, output, tqh, os_val, ns)
return output
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_seq = config["kv_seq_len"]
# bs=4/kv≤1K: Triton nw=8 nst=3 (default, proven best for small batch)
if bs <= 4 and kv_seq <= 1024:
return _triton_kernel(q, kv_data, qo_indptr, kv_indptr, config)
# bs=32,64/kv≤1K: Triton nw=4 nst=2 (s365: 24.6/37.2 vs default 26.7/40.7)
if bs <= 64 and kv_seq <= 1024:
return _triton_kernel(q, kv_data, qo_indptr, kv_indptr, config, nw=4, nst=2)
# bs=256/kv=1K: persistent ps=2 nks=1
if bs >= 256 and kv_seq <= 1024:
return _aiter_persistent(q, kv_data, qo_indptr, config, ps=2, nks=1)
# bs=4/kv=8K: fast_mode=True (S289: -6us improvement)
if bs <= 4:
return _aiter_persistent(q, kv_data, qo_indptr, config, ps=8, fast_mode=True)
# bs=32/kv=8K: ibm=False (s351: -2.8us improvement)
if bs <= 32:
return _aiter_persistent(q, kv_data, qo_indptr, config, ps=8, ibm_override=False)
# kv>=8K: persistent ps=8 (ibm=True default - ibm=False HURTS bs>=64)
return _aiter_persistent(q, kv_data, qo_indptr, config, ps=8)
scrolls · 240 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