submission 739598
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 148 lines, June 9 Researcher Reciprocity License v1.0.
v61.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-739598?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:2c944642d6ae0e570be29d1eec9b04d5deeea38fc802e259ebd774cb925126a6
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
"""a8w8 persistent: FP8 Q + FP8 KV — uses higher-throughput FP8×FP8 MFMA."""Kernel source
v61.py148 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v61 — v60 + split tuning + a8w8 for bs=64/kv=1024.
Changes from v60:
1. bs=64/kv=1024: FP8 NP splits=1 → a8w8 ps=2 splits=16
2. bs=64/kv=8192: splits 4→8
3. bs=256/kv=8192: splits 4→8
"""
from task import input_t, output_t
import torch
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
_meta_cache = {}
_tensor_cache = {}
def _get_cached_tensors(key, create_fn):
cached = _tensor_cache.get(key)
if cached is None:
cached = create_fn()
_tensor_cache[key] = cached
return cached
def _get_or_make_metadata(batch_size, total_kv, num_heads, nhead_kv, num_splits, page_size,
q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_lens, device):
key = (batch_size, total_kv, num_heads, num_splits, page_size, str(q_dtype), str(kv_dtype))
cached = _meta_cache.get(key)
if cached is not None:
return cached
info = get_mla_metadata_info_v1(
batch_size, 1, num_heads, q_dtype, kv_dtype,
is_sparse=False, fast_mode=True,
num_kv_splits=num_splits, intra_batch_mode=False,
)
work = [torch.empty(s, dtype=t, device=device) for s, t in info]
(wmd, wi, wis, ri, rfm, rpm) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_lens,
num_heads // nhead_kv, nhead_kv, False,
wmd, wis, wi, ri, rfm, rpm,
page_size=page_size, kv_granularity=max(page_size, 16),
max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,
max_split_per_batch=num_splits, intra_batch_mode=False,
dtype_q=q_dtype, dtype_kv=kv_dtype,
)
result = dict(
work_meta_data=wmd, work_indptr=wi, work_info_set=wis,
reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
)
_meta_cache[key] = result
return result
def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):
batch_size = config['batch_size']
total_kv = kv_bf16.shape[0]
kv_buffer = kv_bf16.unsqueeze(1)
tensors = _get_cached_tensors(
('bf16', batch_size, total_kv),
lambda: {
'kv_indices': torch.arange(total_kv, device=q.device, dtype=torch.int32),
'kv_last_page_lens': torch.ones(batch_size, device=q.device, dtype=torch.int32),
}
)
mla_decode_fwd(
q=q, kv_buffer=kv_buffer, o=output,
qo_indptr=qo_indptr, kv_indptr=kv_indptr,
kv_indices=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],
max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
)
def _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):
"""a8w8 persistent: FP8 Q + FP8 KV — uses higher-throughput FP8×FP8 MFMA."""
batch_size = config['batch_size']
num_heads = config['num_heads']
total_kv = kv_fp8_data.shape[0]
q_fp8 = q.to(torch.float8_e4m3fn)
num_pages = total_kv // page_size
kv_buffer = kv_fp8_data.view(num_pages, page_size, 1, 576)
tensors = _get_cached_tensors(
('a8w8', batch_size, total_kv, page_size),
lambda: {
'q_scale': torch.ones(1, dtype=torch.float32, device=q.device),
'kv_indices': torch.arange(num_pages, device=q.device, dtype=torch.int32),
'kv_indptr_pages': kv_indptr // page_size,
'kv_last_page_lens': torch.full((batch_size,), page_size, device=q.device, dtype=torch.int32),
}
)
meta = _get_or_make_metadata(
batch_size, total_kv, num_heads, 1, num_splits, page_size,
aiter_dtypes.fp8, aiter_dtypes.fp8,
qo_indptr, tensors['kv_indptr_pages'], tensors['kv_last_page_lens'], q.device,
)
mla_decode_fwd(
q_fp8, kv_buffer, output,
qo_indptr, tensors['kv_indptr_pages'], tensors['kv_indices'], tensors['kv_last_page_lens'],
1, page_size=page_size, nhead_kv=1, sm_scale=config['sm_scale'],
logit_cap=0.0, num_kv_splits=num_splits,
q_scale=tensors['q_scale'], kv_scale=kv_fp8_scale,
intra_batch_mode=False, **meta,
)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config['batch_size']
num_heads = config['num_heads']
v_head_dim = config['v_head_dim']
kv_seq_len = config['kv_seq_len']
output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)
if kv_seq_len <= 1024 and batch_size <= 4:
# BF16 non-persistent — fastest for small batch, exact
_run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
elif kv_seq_len <= 1024:
# a8w8 ps=2 for all kv=1024 shapes (including bs=64)
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
splits = 16 if batch_size <= 64 else 8
_run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=2, num_splits=splits)
else:
# a8w8 ps=8 for kv=8192 — increased splits for larger batches
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
if batch_size <= 4:
splits = 16
elif batch_size <= 64:
splits = 8
else:
splits = 8 # was 4, now 8 for bs=256
_run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)
return output
scrolls · 148 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 677054.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """MLA v90 — v89 with bs=64/kv=1024 reverted to FP8 NP (BF16p failed leaderboard).+ """MLA v61 — v60 + split tuning + a8w8 for bs=64/kv=1024.- Changes from v86:- - bs=4/kv=1024: BF16 NP → BF16 persistent (-1.2µs) ← KEPT- - bs=64/kv=1024: stays FP8 NP splits=1 (BF16p fails secret runner) ← REVERTED+ Changes from v60:+ 1. bs=64/kv=1024: FP8 NP splits=1 → a8w8 ps=2 splits=16+ 2. bs=64/kv=8192: splits 4→8+ 3. bs=256/kv=8192: splits 4→8"""from task import input_t, output_timport torch- import aiter as _aiterfrom aiter.mla import mla_decode_fwdfrom aiter import dtypes as aiter_dtypesfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1- _mc = {}- _tc = {}- _ic = {}- _oc = {}+ _meta_cache = {}+ _tensor_cache = {}- def _gt(key, fn):- v = _tc.get(key)- if v is None: v = fn(); _tc[key] = v- return v- def _gm(bs, tot, nh, ns, ps, qi, ki, kl, dev, qd, kdd):- key = (bs, tot, nh, ns, ps, str(qd), str(kdd))- v = _mc.get(key)- if v is not None: return v- info = get_mla_metadata_info_v1(bs, 1, nh, qd, kdd, is_sparse=False, fast_mode=True, num_kv_splits=ns, intra_batch_mode=False)- work = [torch.empty(s, dtype=t, device=dev) for s, t in info]- wmd, wi, wis, ri, rfm, rpm = work- get_mla_metadata_v1(qi, ki, kl, nh, 1, False, wmd, wis, wi, ri, rfm, rpm,- page_size=ps, kv_granularity=max(ps,16), max_seqlen_qo=1, uni_seqlen_qo=1,- fast_mode=True, max_split_per_batch=ns, intra_batch_mode=False, dtype_q=qd, dtype_kv=kdd)- v = dict(work_meta_data=wmd, work_indptr=wi, work_info_set=wis, reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm)- _mc[key] = v; return v+ def _get_cached_tensors(key, create_fn):+ cached = _tensor_cache.get(key)+ if cached is None:+ cached = create_fn()+ _tensor_cache[key] = cached+ return cached- def _gi(key, fn):- v = _ic.get(key)- if v is None: v = fn(); _ic[key] = v- return v+ def _get_or_make_metadata(batch_size, total_kv, num_heads, nhead_kv, num_splits, page_size,+ q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_lens, device):+ key = (batch_size, total_kv, num_heads, num_splits, page_size, str(q_dtype), str(kv_dtype))+ cached = _meta_cache.get(key)+ if cached is not None:+ return cached- 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']- vd = config['v_head_dim']- kvl = config['kv_seq_len']- sms = config['sm_scale']- dev = q.device- tq = q.shape[0]+ info = get_mla_metadata_info_v1(+ batch_size, 1, num_heads, q_dtype, kv_dtype,+ is_sparse=False, fast_mode=True,+ num_kv_splits=num_splits, intra_batch_mode=False,+ )+ work = [torch.empty(s, dtype=t, device=device) for s, t in info]+ (wmd, wi, wis, ri, rfm, rpm) = work- ok = (tq, nh, vd)- o = _oc.get(ok)- if o is None:- o = torch.empty((tq, nh, vd), dtype=q.dtype, device=dev)- _oc[ok] = o+ get_mla_metadata_v1(+ qo_indptr, kv_indptr, kv_last_page_lens,+ num_heads // nhead_kv, nhead_kv, False,+ wmd, wis, wi, ri, rfm, rpm,+ page_size=page_size, kv_granularity=max(page_size, 16),+ max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,+ max_split_per_batch=num_splits, intra_batch_mode=False,+ dtype_q=q_dtype, dtype_kv=kv_dtype,+ )- kd, ks = kv_data["fp8"]- tot = kd.shape[0]+ result = dict(+ work_meta_data=wmd, work_indptr=wi, work_info_set=wis,+ reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,+ )+ _meta_cache[key] = result+ return result- # ---- BF16 persistent for bs<=4/kv<=1024 only (bs=64 FAILS leaderboard) ----- if kvl <= 1024 and bs <= 4:- kv_bf16 = kv_data["bf16"]- ps, ns = 2, 16- np = tot // ps- t = _gt(('bfp', bs, tot, ps), lambda: {- 'ki': torch.arange(np, device=dev, dtype=torch.int32),- 'kip': kv_indptr // ps,- 'kl': torch.full((bs,), ps, device=dev, dtype=torch.int32),- })- m = _gm(bs, tot, nh, ns, ps, qo_indptr, t['kip'], t['kl'], dev, torch.bfloat16, torch.bfloat16)- mla_decode_fwd(q, kv_bf16.view(np, ps, 1, 576), o,- qo_indptr, t['kip'], t['ki'], t['kl'],- 1, page_size=ps, nhead_kv=1, sm_scale=sms,- logit_cap=0.0, num_kv_splits=ns,- intra_batch_mode=False, **m)- return o- # ---- FP8 NP splits=1 for bs=64/kv=1024 (BF16p fails leaderboard) ----- if kvl <= 1024 and bs == 64:- q8 = q.to(torch.float8_e4m3fn)- t = _gt(('fn', bs, tot), lambda: {- 'qs': torch.ones(1, dtype=torch.float32, device=dev),- 'ki': torch.arange(tot, device=dev, dtype=torch.int32),- 'kl': torch.ones(bs, device=dev, dtype=torch.int32),- 'si': torch.arange(bs+1, dtype=torch.int32, device=dev),- })- mla_decode_fwd(q=q8, kv_buffer=kd.unsqueeze(1), o=o,- qo_indptr=qo_indptr, kv_indptr=kv_indptr,- kv_indices=t['ki'], kv_last_page_lens=t['kl'],- max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=sms,- num_kv_splits=1, num_kv_splits_indptr=t['si'],- q_scale=t['qs'], kv_scale=ks)- return o+ def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):+ batch_size = config['batch_size']+ total_kv = kv_bf16.shape[0]+ kv_buffer = kv_bf16.unsqueeze(1)- # ---- a8w8 persistent (direct calls) for all other shapes ----- q8 = q.to(torch.float8_e4m3fn)- if kvl <= 1024:- ps, ns = 2, (16 if bs <= 32 else 8)- else:- ps, ns = 8, (16 if bs <= 4 else 8)+ tensors = _get_cached_tensors(+ ('bf16', batch_size, total_kv),+ lambda: {+ 'kv_indices': torch.arange(total_kv, device=q.device, dtype=torch.int32),+ 'kv_last_page_lens': torch.ones(batch_size, device=q.device, dtype=torch.int32),+ }+ )- np = tot // ps- t = _gt(('a8', bs, tot, ps), lambda: {- 'qs': torch.ones(1, dtype=torch.float32, device=dev),- 'ki': torch.arange(np, device=dev, dtype=torch.int32),- 'kip': kv_indptr // ps,- 'kl': torch.full((bs,), ps, device=dev, dtype=torch.int32),- })- m = _gm(bs, tot, nh, ns, ps, qo_indptr, t['kip'], t['kl'], dev, aiter_dtypes.fp8, aiter_dtypes.fp8)+ mla_decode_fwd(+ q=q, kv_buffer=kv_buffer, o=output,+ qo_indptr=qo_indptr, kv_indptr=kv_indptr,+ kv_indices=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],+ max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],+ )- rpm_sz = m['reduce_partial_map'].size(0)- inter = _gi(('a8_i', rpm_sz, nh, vd), lambda: {- 'logits': torch.empty((rpm_sz, 1, nh, vd), dtype=torch.float32, device=dev),- 'attn_lse': torch.empty((rpm_sz, 1, nh, 1), dtype=torch.float32, device=dev),- })- _aiter.mla_decode_stage1_asm_fwd(- q8, kd.view(np, ps, 1, 576), qo_indptr, t['kip'],- t['ki'], t['kl'],- None,- m['work_meta_data'], m['work_indptr'], m['work_info_set'],- 1, ps, 1, sms,- inter['logits'], inter['attn_lse'], o,- t['qs'], ks,+ def _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):+ """a8w8 persistent: FP8 Q + FP8 KV — uses higher-throughput FP8×FP8 MFMA."""+ batch_size = config['batch_size']+ num_heads = config['num_heads']+ total_kv = kv_fp8_data.shape[0]++ q_fp8 = q.to(torch.float8_e4m3fn)+ num_pages = total_kv // page_size+ kv_buffer = kv_fp8_data.view(num_pages, page_size, 1, 576)++ tensors = _get_cached_tensors(+ ('a8w8', batch_size, total_kv, page_size),+ lambda: {+ 'q_scale': torch.ones(1, dtype=torch.float32, device=q.device),+ 'kv_indices': torch.arange(num_pages, device=q.device, dtype=torch.int32),+ 'kv_indptr_pages': kv_indptr // page_size,+ 'kv_last_page_lens': torch.full((batch_size,), page_size, device=q.device, dtype=torch.int32),+ })- _aiter.mla_reduce_v1(- inter['logits'], inter['attn_lse'],- m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],- 1, o, None,+ meta = _get_or_make_metadata(+ batch_size, total_kv, num_heads, 1, num_splits, page_size,+ aiter_dtypes.fp8, aiter_dtypes.fp8,+ qo_indptr, tensors['kv_indptr_pages'], tensors['kv_last_page_lens'], q.device,)- return o++ mla_decode_fwd(+ q_fp8, kv_buffer, output,+ qo_indptr, tensors['kv_indptr_pages'], tensors['kv_indices'], tensors['kv_last_page_lens'],+ 1, page_size=page_size, nhead_kv=1, sm_scale=config['sm_scale'],+ logit_cap=0.0, num_kv_splits=num_splits,+ q_scale=tensors['q_scale'], kv_scale=kv_fp8_scale,+ intra_batch_mode=False, **meta,+ )+++ def custom_kernel(data: input_t) -> output_t:+ q, kv_data, qo_indptr, kv_indptr, config = data++ batch_size = config['batch_size']+ num_heads = config['num_heads']+ v_head_dim = config['v_head_dim']+ kv_seq_len = config['kv_seq_len']++ output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)++ if kv_seq_len <= 1024 and batch_size <= 4:+ # BF16 non-persistent — fastest for small batch, exact+ _run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)+ elif kv_seq_len <= 1024:+ # a8w8 ps=2 for all kv=1024 shapes (including bs=64)+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]+ splits = 16 if batch_size <= 64 else 8+ _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=2, num_splits=splits)+ else:+ # a8w8 ps=8 for kv=8192 — increased splits for larger batches+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]+ if batch_size <= 4:+ splits = 16+ elif batch_size <= 64:+ splits = 8+ else:+ splits = 8 # was 4, now 8 for bs=256+ _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)++ return output
scrolls · 257 diff lines total
Best evidence level for this revision: reported
JSON