submission 754850
Shaw · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 121 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754850?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:4ce8a30cbb7bdedfe0367280ef7db6e9638446648ee9871aed0ad0854cacdecf
license declaredunknown
license concludedunknown
authorsShaw
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
- a16w16 persistent bf16_kv for bs32/kv1024 (new: was a16w8)Kernel source
submission.py121 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V07-SAFE: Only proven bs<=32 improvements, baseline bs>=64.
- NP a16w16 for bs<=4 (proven on LB, same as baseline v258)
- a16w16 persistent bf16_kv for bs32/kv1024 (new: was a16w8)
- bs32/kv8192: same as baseline (already a16w16+bf16_kv)
- ALL bs>=64: EXACTLY same as baseline v258 (passes all secret seeds)
Risk: LOW. Only 1 shape differs from proven baseline.
Expected: small improvement from bs32/kv1024 a16w16 switch.
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
FP8 = aiter_dtypes.fp8
NH, NKH, QKD, VD = 16, 1, 576, 512
SM = 1.0 / (QKD ** 0.5)
_cc = {}
_SHAPES = [(4,1024),(4,8192),(32,1024),(32,8192),(64,1024),(64,8192),(256,1024),(256,8192)]
def _cfg(bs, kv):
# bs<=4: NP a16w16 (same as baseline)
if bs <= 4:
if kv <= 1024: return 'a16w16', 1, True, True, None, True, False
return 'a16w16', 8, True, True, None, True, False
# bs32/kv1024: NEW — a16w16+bf16_kv (EXP-151 proven for kv8192, trying for kv1024)
if bs <= 32:
if kv <= 1024: return 'a16w16', 2, False, True, 4, False, True # NEW: a16w16+bf16_kv
return 'a16w16', 8, False, True, 32, False, True # SAME as baseline
# === ALL BELOW: EXACTLY baseline v258 ===
# bs64/kv1024: pg=2 ib=False (fixes seed 1357)
if kv <= 1024 and bs <= 64:
return 'a16w8', 2, False, False, 16, False, False
# bs256/kv1024
if kv <= 1024:
return 'a16w8', 2, True, True, 4, False, False
# bs64/kv8192
if bs <= 64:
return 'a16w8', 8, False, True, 16, False, False
# bs256/kv8192
return 'a16w8', 8, True, True, 16, False, False
def _build(bs, qs, kv, qo, kvi):
key = (bs, kv)
if key in _cc: return _cc[key]
mode, pg, fm, ib, ns, use_np, bf16_kv = _cfg(bs, kv)
tq, tkv = bs * qs, bs * kv
tp = tkv // pg
kidx = torch.arange(tp, dtype=torch.int32, device="cuda")
kpi = (kvi // pg).to(torch.int32)
klp = torch.full((bs,), pg, dtype=torch.int32, device="cuda")
o = torch.empty(tq, NH, VD, dtype=torch.bfloat16, device="cuda")
if use_np:
_cc[key] = dict(mode=mode, kidx=kidx, klp=klp, o=o, tkv=tkv, kpi=kpi, pg=pg,
use_np=True, bf16_kv=False)
return _cc[key]
qdm = torch.bfloat16
kvdm = torch.bfloat16 if bf16_kv else FP8
info = get_mla_metadata_info_v1(bs, qs, NH, qdm, kvdm,
is_sparse=False, fast_mode=fm, num_kv_splits=ns, intra_batch_mode=ib)
meta = {k: torch.empty(info[i][0], dtype=info[i][1], device="cuda")
for i, k in enumerate(["work_meta_data","work_indptr","work_info_set",
"reduce_indptr","reduce_final_map","reduce_partial_map"])}
get_mla_metadata_v1(qo, kpi, klp, NH // NKH, NKH, True,
meta["work_meta_data"], meta["work_info_set"], meta["work_indptr"],
meta["reduce_indptr"], meta["reduce_final_map"], meta["reduce_partial_map"],
page_size=pg, kv_granularity=max(pg, 16), max_seqlen_qo=qs, uni_seqlen_qo=qs,
fast_mode=fm, max_split_per_batch=ns, intra_batch_mode=ib,
dtype_q=qdm, dtype_kv=kvdm)
_cc[key] = dict(mode=mode, kidx=kidx, klp=klp, meta=meta, o=o, tkv=tkv,
ns=ns, kpi=kpi, pg=pg, ib=ib, use_np=False, bf16_kv=bf16_kv)
return _cc[key]
for bs, kv in _SHAPES:
qo = torch.arange(bs + 1, dtype=torch.int32, device="cuda")
kvi = torch.arange(bs + 1, dtype=torch.int32, device="cuda") * kv
_build(bs, 1, kv, qo, kvi)
def _run(q, kv_data, qo_indptr, c):
tp = c['tkv'] // c['pg']
if c['use_np']:
kd = kv_data["bf16"].view(tp, c['pg'], NKH, QKD)
mla_decode_fwd(q.view(-1, NH, QKD), kd, c['o'],
qo_indptr, c['kpi'], c['kidx'], c['klp'],
1, page_size=c['pg'], nhead_kv=NKH,
sm_scale=SM, logit_cap=0.0,
num_kv_splits=None, q_scale=None, kv_scale=None)
elif c['bf16_kv']:
kd = kv_data["bf16"].view(tp, c['pg'], NKH, QKD)
mla_decode_fwd(q.view(-1, NH, QKD), kd, c['o'],
qo_indptr, c['kpi'], c['kidx'], c['klp'],
1, page_size=c['pg'], nhead_kv=NKH,
sm_scale=SM, logit_cap=0.0,
num_kv_splits=c['ns'],
q_scale=None, kv_scale=None,
intra_batch_mode=c['ib'], **c['meta'])
else:
kv_fp8, kv_sc = kv_data["fp8"]
kd = kv_fp8.view(tp, c['pg'], NKH, QKD)
mla_decode_fwd(q.view(-1, NH, QKD), kd, c['o'],
qo_indptr, c['kpi'], c['kidx'], c['klp'],
1, page_size=c['pg'], nhead_kv=NKH,
sm_scale=SM, logit_cap=0.0,
num_kv_splits=c['ns'],
q_scale=None, kv_scale=kv_sc,
intra_batch_mode=c['ib'], **c['meta'])
return c['o']
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs, qs, kv = config["batch_size"], config["q_seq_len"], config["kv_seq_len"]
c = _build(bs, qs, kv, qo_indptr, kv_indptr)
return _run(q, kv_data, qo_indptr, c)
scrolls · 121 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