submission 755138
shaw061434 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 133 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-755138?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:afa4fec46b566377409e1e4ed9f0b3177ed6325adf9902028585d0714a12fdd0
license declaredunknown
license concludedunknown
authorsshaw061434
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
Teammate wikty EXP-151: a16w16 persistent -7.3% for bs32/kv8192Kernel source
submission.py133 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""EXP-258: Combined config from teammate findings + our kv_gran=16.
Teammate wikty EXP-151: a16w16 persistent -7.3% for bs32/kv8192
Teammate EXP-150: fm=False better for bs32/bs64/bs256
Our findings: kv_gran=16, ib=False fixes seed 1357
Combined: best of both configs.
"""
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):
# Returns: (mode, pg, fm, ib, ns, use_np, bf16_kv)
# EXP-169: cherry-pick safe changes from EXP-168 + NP for bs32/kv1024
# bs≤4: NP a16w16 (proven)
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 — NP a16w16 (skip metadata+reduce for medium batch)
# bs32/kv8192: lower ns (32→16, safe from EXP-168 data)
if bs <= 32:
if kv <= 1024: return 'a16w16', 2, True, True, None, True, False # NP!
return 'a16w16', 8, False, True, 16, False, True # ns:32→16
# bs64/kv1024: KEEP ns=16 (ns=8 broke correctness in EXP-168)
if kv <= 1024 and bs <= 64:
return 'a16w8', 2, False, False, 16, False, False
# bs256/kv1024: keep ns=4 + try fm=False (EXP-150: -0.9%)
if kv <= 1024:
return 'a16w8', 2, False, True, 4, False, False # fm=False
# bs64/kv8192: ns 16→8 (proven -2.3% in EXP-168)
if bs <= 64:
return 'a16w8', 8, False, True, 8, False, False
# bs256/kv8192: ns 16→8 + fm=False
return 'a16w8', 8, False, True, 8, 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 · 133 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