submission 645036
John Hahn · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 177 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-645036?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:259286088521b4e682d83bbcb619a19ec942e05152b0ca99fceaf60d58334dad
license declaredunknown
license concludedunknown
authorsJohn Hahn
imported2026-08-15
Kernel source
submission.py177 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
v21_best_combo: Cherry-pick best config per shape from v12 + v19.
kv=1024: keep v12 splits (32 for bs=4, 16 for rest) — fast_mode for bs=4.
kv=8192: use v19 splits=16 everywhere (was 32 for bs=4,32 in v12).
"""
import torch
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
QK_DIM = 576
V_DIM = 512
SM_SCALE = 1.0 / (QK_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
_TUNE = {
(4, 1024): (1, 32, True),
(4, 8192): (8, 16, True),
(32, 1024): (1, 16, False),
(32, 8192): (8, 16, False),
(64, 1024): (1, 16, False),
(64, 8192): (8, 16, False),
(256, 1024): (1, 16, False),
(256, 8192): (8, 16, False),
}
_cache = {}
_q_scale = None
def _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, dev):
global _q_scale
if cfg_key in _cache:
return _cache[cfg_key]
if _q_scale is None:
_q_scale = torch.ones(1, dtype=torch.float32, device=dev)
ps, num_splits, fast_mode = _TUNE.get((bs, kvsl), (1, 32, bs <= 4))
nkv = 1
effective_bs = bs * qsl
if qsl > 1:
eff_kv_indptr = torch.zeros(effective_bs + 1, dtype=torch.int32, device=dev)
for i in range(bs):
kv_len = kv_indptr[i + 1].item() - kv_indptr[i].item()
for j in range(qsl):
idx = i * qsl + j
eff_kv_indptr[idx + 1] = eff_kv_indptr[idx] + kv_len
else:
eff_kv_indptr = kv_indptr
eff_qo_indptr = torch.arange(effective_bs + 1, dtype=torch.int32, device=dev)
total_eff_kv = int(eff_kv_indptr[-1].item())
if ps > 1:
pages_per_seq = (kvsl + ps - 1) // ps
kv_last_page_len = torch.full((effective_bs,),
kvsl % ps if kvsl % ps != 0 else ps,
dtype=torch.int32, device=dev)
kv_indices = torch.arange(effective_bs * pages_per_seq, dtype=torch.int32, device=dev)
paged_kv_indptr = torch.arange(effective_bs + 1, dtype=torch.int32, device=dev) * pages_per_seq
else:
kv_last_page_len = torch.full((effective_bs,), ps, dtype=torch.int32, device=dev)
kv_indices = torch.arange(total_eff_kv, dtype=torch.int32, device=dev)
paged_kv_indptr = eff_kv_indptr
info = get_mla_metadata_info_v1(
effective_bs, 1, nh, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=fast_mode,
num_kv_splits=num_splits, intra_batch_mode=(not fast_mode),
)
work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
eff_qo_indptr, paged_kv_indptr, kv_last_page_len,
nh // nkv, nkv, True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=ps,
kv_granularity=max(ps, 16),
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=fast_mode,
max_split_per_batch=num_splits,
intra_batch_mode=(not fast_mode),
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
total_q = bs * qsl
o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)
q_fp8_buf = torch.empty((total_q, nh, QK_DIM), dtype=FP8_DTYPE, device=dev)
entry = {
'ps': ps, 'num_splits': num_splits, 'fast_mode': fast_mode,
'eff_qo_indptr': eff_qo_indptr,
'paged_kv_indptr': paged_kv_indptr,
'kv_indices': kv_indices,
'kv_last_page_len': kv_last_page_len,
'meta': {
'work_meta_data': work_metadata,
'work_indptr': work_indptr,
'work_info_set': work_info_set,
'reduce_indptr': reduce_indptr,
'reduce_final_map': reduce_final_map,
'reduce_partial_map': reduce_partial_map,
},
'o': o,
'q_fp8_buf': q_fp8_buf,
'qsl': qsl, 'bs': bs, 'nh': nh, 'nkv': nkv,
'effective_bs': effective_bs,
'kvsl': kvsl,
}
_cache[cfg_key] = entry
return entry
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
nh = config["num_heads"]
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
cfg_key = (bs, qsl, kvsl, nh)
c = _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, q.device)
q_fp8 = c['q_fp8_buf']
q_fp8.copy_(q.view(bs * qsl, nh, QK_DIM))
kv_fp8_raw, kv_fp8_scale = kv_data["fp8"]
kv_scale = kv_fp8_scale.view(1) if kv_fp8_scale.numel() == 1 else kv_fp8_scale
ps = c['ps']
total_kv = bs * kvsl
if ps > 1:
pages_per_seq = (kvsl + ps - 1) // ps
effective_bs = c['effective_bs']
kv_buf = kv_fp8_raw.view(bs, kvsl, 1, QK_DIM)
if qsl > 1:
kv_buf = kv_buf.repeat_interleave(qsl, dim=0)
effective_bs = bs * qsl
if kvsl % ps == 0:
kv_4d = kv_buf.reshape(effective_bs * pages_per_seq, ps, 1, QK_DIM)
else:
pad_len = pages_per_seq * ps - kvsl
kv_buf = torch.nn.functional.pad(kv_buf, (0, 0, 0, 0, 0, pad_len))
kv_4d = kv_buf.reshape(effective_bs * pages_per_seq, ps, 1, QK_DIM)
else:
if qsl > 1:
kv_buf = kv_fp8_raw.view(bs, kvsl, 1, QK_DIM)
kv_4d = kv_buf.repeat(qsl, 1, 1, 1).reshape(bs * qsl * kvsl, 1, 1, QK_DIM)
else:
kv_4d = kv_fp8_raw.view(total_kv, 1, 1, QK_DIM)
o = c['o']
mla_decode_fwd(
q_fp8, kv_4d, o,
c['eff_qo_indptr'], c['paged_kv_indptr'],
c['kv_indices'], c['kv_last_page_len'],
1, page_size=ps, nhead_kv=1,
sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=c['num_splits'],
q_scale=_q_scale, kv_scale=kv_scale,
intra_batch_mode=(not c['fast_mode']),
**c['meta'],
)
return o
scrolls · 177 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 644617.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- v12_safe: ps=1 for kv=1024 (zero mismatch), ps=8 for kv=8192.- All cases: qsl=1, nh=16.- Avoids mismatch warnings that could fail leaderboard secret tests.+ v21_best_combo: Cherry-pick best config per shape from v12 + v19.+ kv=1024: keep v12 splits (32 for bs=4, 16 for rest) — fast_mode for bs=4.+ kv=8192: use v19 splits=16 everywhere (was 32 for bs=4,32 in v12)."""import torchfrom task import input_t, output_t⋯ 6 unchanged linesSM_SCALE = 1.0 / (QK_DIM ** 0.5)FP8_DTYPE = aiter_dtypes.fp8- # (bs, kvsl) -> (page_size, num_kv_splits, fast_mode)_TUNE = {(4, 1024): (1, 32, True),- (4, 8192): (8, 32, True),+ (4, 8192): (8, 16, True),(32, 1024): (1, 16, False),- (32, 8192): (8, 32, False),+ (32, 8192): (8, 16, False),(64, 1024): (1, 16, False),(64, 8192): (8, 16, False),(256, 1024): (1, 16, False),
scrolls · 28 diff lines total
Best evidence level for this revision: reported
JSON