submission 754469
Leon · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 136 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754469?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:c6e910e45543e2dd48c6b4e6e8bae7d6cc64dee6d92a54306f0cad8e0b539ce5
license declaredunknown
license concludedunknown
authorsLeon
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
1. Bypass mla_decode_fwd() Python overhead for persistent shapesKernel source
submission.py136 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""EXP-261: Direct API calls + pre-allocated buffers.
Source-level optimization (MOE teammate methodology):
1. Bypass mla_decode_fwd() Python overhead for persistent shapes
2. Pre-allocate logits/attn_lse buffers (eliminate torch.empty per call)
3. Direct aiter.mla_decode_stage1_asm_fwd + aiter.mla_reduce_v1
4. Keep 1-warmup for JIT, no self-timing loop
5. All dispatch params match EXP-259 (proven config)
"""
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)
if bs <= 4:
if kv <= 1024: return 'a16w16', 1, True, True, None, True, False
return 'a16w16', 8, True, True, None, True, False
if bs <= 32:
if kv <= 1024: return 'a16w8', 2, True, True, 8, False, False
return 'a16w16', 8, False, True, 32, False, True
if kv <= 1024 and bs <= 64:
return 'a16w8', 2, False, False, 16, False, False
if kv <= 1024:
return 'a16w8', 2, True, True, 4, False, False
if bs <= 64:
return 'a16w8', 8, False, True, 16, False, False
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)
# Pre-allocate intermediate buffers for direct API calls
rpm = meta["reduce_partial_map"].size(0)
logits = torch.empty((rpm, 1, NH, VD), dtype=torch.float32, device="cuda")
attn_lse = torch.empty((rpm, 1, NH, 1), dtype=torch.float32, device="cuda")
_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,
logits=logits, attn_lse=attn_lse)
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)
else:
# Direct bottom-level API: skip mla_decode_fwd Python overhead + torch.empty
m = c['meta']
if c['bf16_kv']:
kd = kv_data["bf16"].view(tp, c['pg'], NKH, QKD)
kv_sc = None
else:
kv_fp8, kv_sc = kv_data["fp8"]
kd = kv_fp8.view(tp, c['pg'], NKH, QKD)
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, NH, QKD), kd,
qo_indptr, c['kpi'], c['kidx'], c['klp'],
None, # num_kv_splits_indptr (persistent mode)
m['work_meta_data'], m['work_indptr'], m['work_info_set'],
1, c['pg'], NKH, SM,
c['logits'], c['attn_lse'], c['o'],
None, kv_sc) # q_scale, kv_scale (runner has 19 params, no lse)
aiter.mla_reduce_v1(
c['logits'], c['attn_lse'],
m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],
1, c['o'], None)
return c['o']
_warmed = set()
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)
key = (bs, kv)
if key not in _warmed:
_warmed.add(key)
_run(q, kv_data, qo_indptr, c)
torch.cuda.synchronize()
return _run(q, kv_data, qo_indptr, c)
scrolls · 136 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