submission 716430
somethingobscurefordevstuff · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 131 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-716430?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:3af2846a6f2f8add6b226670bea510754973649832eef374a1ee88a986c273e1
license declaredunknown
license concludedunknown
authorssomethingobscurefordevstuff
imported2026-08-15
Kernel source
submission.py131 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# v382 plus mla_decode_fwd only on the medium-short bf16 region (48<=bs<128, kv<=2048).
import aiter
import torch
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
from task import input_t, output_t
NH = 16; NKV = 1; DQ = 576; DV = 512; SM = DQ**-0.5
FP8 = torch.float8_e4m3fn; BF16 = torch.bfloat16
_s1 = aiter.mla_decode_stage1_asm_fwd; _rd = aiter.mla_reduce_v1
_c = {}
def _cfg(bs, kvl):
large_kv = kvl > 2048
med_bs = 16 < bs < 128
medium_short = 48 <= bs < 128 and kvl <= 2048
use_fp8 = (large_kv and med_bs) or (bs >= 128)
kv_dt = FP8 if use_fp8 else BF16
if use_fp8 and large_kv and med_bs:
fm = False
ibm = bs >= 48
else:
fm = True
ibm = True
if bs <= 8 and kvl <= 1024:
ps = 1
elif kvl <= 1024:
ps = 2
else:
ps = 8
if bs >= 128:
ns = 1
elif bs >= 48:
ns = 4
elif bs >= 16:
ns = 8
else:
ns = 16
return kv_dt, BF16, fm, ibm, ps, ns, medium_short
def _setup(bs, kvl, dev):
k = (bs, kvl)
if k in _c:
return _c[k]
if bs <= 8 and kvl <= 1024:
total = bs * kvl
qo = torch.arange(bs + 1, dtype=torch.int32, device=dev)
kvi = torch.arange(bs + 1, dtype=torch.int32, device=dev) * kvl
klp = torch.full((bs,), 1, dtype=torch.int32, device=dev)
ki = torch.arange(total, dtype=torch.int32, device=dev)
o = torch.empty((bs, NH, DV), dtype=BF16, device=dev)
_c[k] = ("nonpersist", o, qo, kvi, klp, ki)
return _c[k]
kv_dt, q_dt, fm, ibm, ps, ns, medium_short = _cfg(bs, kvl)
pp = kvl // ps
qo = torch.arange(bs + 1, dtype=torch.int32, device=dev)
kvi = torch.arange(bs + 1, dtype=torch.int32, device=dev) * pp
klp = torch.full((bs,), ps, dtype=torch.int32, device=dev)
ki = torch.arange(bs * pp, dtype=torch.int32, device=dev)
info = get_mla_metadata_info_v1(
bs, 1, NH, q_dt, kv_dt,
is_sparse=False, fast_mode=fm, num_kv_splits=ns, intra_batch_mode=ibm,
)
wm, wi, ws, ri, rf, rp = [torch.empty(s, dtype=t, device=dev) for s, t in info]
kvg = 8 if (bs >= 128 or (bs >= 64 and kvl > 2048)) else max(ps, 16)
get_mla_metadata_v1(
qo, kvi, klp, NH, NKV, False, wm, ws, wi, ri, rf, rp,
page_size=ps, kv_granularity=kvg, max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=fm, max_split_per_batch=ns, intra_batch_mode=ibm, dtype_q=q_dt, dtype_kv=kv_dt,
)
o = torch.empty((bs, NH, DV), dtype=BF16, device=dev)
po = torch.empty((bs, ns, NH, DV), dtype=torch.float32, device=dev)
pl = torch.empty((bs, ns, NH, 1), dtype=torch.float32, device=dev)
use_bf16 = (kv_dt == BF16)
use_fwd_bf16 = medium_short and use_bf16
_c[k] = ("persist", o, po, pl, qo, kvi, klp, ki, wm, wi, ws, ri, rf, rp, ps, ns, ibm, use_bf16, use_fwd_bf16)
return _c[k]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, _, _, _ = data
bs = q.shape[0]
kv_fp8, kv_scale = kv_data["fp8"]
kv_bf16 = kv_data["bf16"]
kvl = kv_fp8.shape[0] // bs
entry = _setup(bs, kvl, q.device)
if entry[0] == "nonpersist":
_, o, qo, kvi, klp, ki = entry
kv_buf = kv_bf16.view(bs * kvl, 1, NKV, DQ)
mla_decode_fwd(
q=q, kv_buffer=kv_buf, o=o,
qo_indptr=qo, kv_indptr=kvi, kv_indices=ki, kv_last_page_lens=klp,
max_seqlen_q=1, page_size=1, nhead_kv=NKV, sm_scale=SM,
)
return o
_, o, po, pl, qo, kvi, klp, ki, wm, wi, ws, ri, rf, rp, ps, ns, ibm, use_bf16, use_fwd_bf16 = entry
if use_fwd_bf16:
kv4 = kv_bf16.view(kv_bf16.shape[0] // ps, ps, NKV, DQ)
mla_decode_fwd(
q, kv4, o, qo, kvi, ki, klp, 1,
page_size=ps, nhead_kv=NKV, sm_scale=SM,
num_kv_splits=ns, q_scale=None, kv_scale=None,
intra_batch_mode=ibm,
work_meta_data=wm, work_indptr=wi, work_info_set=ws,
reduce_indptr=ri, reduce_final_map=rf, reduce_partial_map=rp,
)
return o
if use_bf16:
kv4 = kv_bf16.view(kv_bf16.shape[0] // ps, ps, NKV, DQ)
_s1(q, kv4, qo, kvi, ki, klp, None, wm, wi, ws, 1, ps, NKV, SM, po, pl, o, None, None)
else:
kv4 = kv_fp8.view(kv_fp8.shape[0] // ps, ps, NKV, DQ)
_s1(q, kv4, qo, kvi, ki, klp, None, wm, wi, ws, 1, ps, NKV, SM, po, pl, o, None, kv_scale)
if ns > 1:
_rd(po, pl, ri, rf, rp, 1, o, None)
return o
scrolls · 131 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