submission 641982
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 134 lines, June 9 Researcher Reciprocity License v1.0.
v86.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-641982?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:9858edadd7518acac7ef5fe342242ec25312687524d0d150fa16308fe65645b9
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
"""MLA v86 — Hybrid: mla_decode_fwd for NP paths, direct stage1+reduce for persistent.Kernel source
v86.py134 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v86 — Hybrid: mla_decode_fwd for NP paths, direct stage1+reduce for persistent.
Bypasses mla_decode_fwd ONLY for a8w8 persistent paths where torch.empty
overhead is highest (2x allocations per call). Keeps mla_decode_fwd for BF16 NP
and FP8 NP where overhead is lower.
"""
from task import input_t, output_t
import torch
import aiter as _aiter
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
_mc = {}
_tc = {}
_ic = {}
_oc = {}
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):
key = (bs, tot, nh, ns, ps)
v = _mc.get(key)
if v is not None: return v
qd, kd = aiter_dtypes.fp8, aiter_dtypes.fp8
info = get_mla_metadata_info_v1(bs, 1, nh, qd, kd, 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=kd)
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 _gi(key, fn):
v = _ic.get(key)
if v is None: v = fn(); _ic[key] = v
return v
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]
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
kd, ks = kv_data["fp8"]
tot = kd.shape[0]
# ---- BF16 non-persistent via mla_decode_fwd (bs<=4, kv<=1024) ----
if kvl <= 1024 and bs <= 4:
kv_bf16 = kv_data["bf16"]
t = _gt(('bf', bs, tot), lambda: {
'ki': torch.arange(tot, device=dev, dtype=torch.int32),
'kl': torch.ones(bs, device=dev, dtype=torch.int32),
})
mla_decode_fwd(q=q, kv_buffer=kv_bf16.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)
return o
# ---- FP8 NP splits=1 via mla_decode_fwd (bs=64, kv<=1024) ----
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
# ---- a8w8 persistent: DIRECT stage1_asm + reduce_v1 (bypass mla_decode_fwd) ----
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)
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)
# Pre-allocate intermediates (the key optimization)
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,
)
_aiter.mla_reduce_v1(
inter['logits'], inter['attn_lse'],
m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],
1, o, None,
)
return o
scrolls · 134 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 625424.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """MLA v64 — v60 dispatch + v61 split tuning. Safe for secret runner.+ """MLA v86 — Hybrid: mla_decode_fwd for NP paths, direct stage1+reduce for persistent.- v61's bs=64/kv=1024 a8w8 ps=2 FAILED secret runner (mismatch >5%).- Revert to v60's FP8 NP for that shape. Keep splits=8 for kv=8192.+ Bypasses mla_decode_fwd ONLY for a8w8 persistent paths where torch.empty+ overhead is highest (2x allocations per call). Keeps mla_decode_fwd for BF16 NP+ and FP8 NP where overhead is lower."""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- _meta_cache = {}- _tensor_cache = {}+ _mc = {}+ _tc = {}+ _ic = {}+ _oc = {}+ def _gt(key, fn):+ v = _tc.get(key)+ if v is None: v = fn(); _tc[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 _gm(bs, tot, nh, ns, ps, qi, ki, kl, dev):+ key = (bs, tot, nh, ns, ps)+ v = _mc.get(key)+ if v is not None: return v+ qd, kd = aiter_dtypes.fp8, aiter_dtypes.fp8+ info = get_mla_metadata_info_v1(bs, 1, nh, qd, kd, 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=kd)+ 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 _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- 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+ 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]- 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,- )+ 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- 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+ kd, ks = kv_data["fp8"]+ tot = kd.shape[0]+ # ---- BF16 non-persistent via mla_decode_fwd (bs<=4, kv<=1024) ----+ if kvl <= 1024 and bs <= 4:+ kv_bf16 = kv_data["bf16"]+ t = _gt(('bf', bs, tot), lambda: {+ 'ki': torch.arange(tot, device=dev, dtype=torch.int32),+ 'kl': torch.ones(bs, device=dev, dtype=torch.int32),+ })+ mla_decode_fwd(q=q, kv_buffer=kv_bf16.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)+ 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)+ # ---- FP8 NP splits=1 via mla_decode_fwd (bs=64, kv<=1024) ----+ 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- 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),- }- )+ # ---- a8w8 persistent: DIRECT stage1_asm + reduce_v1 (bypass mla_decode_fwd) ----+ 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)- 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'],- )+ 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)+ # Pre-allocate intermediates (the key optimization)+ 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),+ })- def _run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config):- batch_size = config['batch_size']- total_kv = kv_fp8_data.shape[0]-- q_fp8 = q.to(torch.float8_e4m3fn)- kv_buffer = kv_fp8_data.unsqueeze(1)-- tensors = _get_cached_tensors(- ('fp8np', batch_size, total_kv),- lambda: {- 'q_scale': torch.ones(1, dtype=torch.float32, device=q.device),- '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),- 'num_kv_splits_indptr': torch.arange(batch_size + 1, dtype=torch.int32, device=q.device),- }+ _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,)- mla_decode_fwd(- q=q_fp8, 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'],- num_kv_splits=1, num_kv_splits_indptr=tensors['num_kv_splits_indptr'],- q_scale=tensors['q_scale'], kv_scale=kv_fp8_scale,+ _aiter.mla_reduce_v1(+ inter['logits'], inter['attn_lse'],+ m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],+ 1, o, None,)--- 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."""- 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 and batch_size == 64:- # FP8 NP splits=1 — safe for secret runner (ps=2 fails at bs=64)- kv_fp8_data, kv_fp8_scale = kv_data["fp8"]- _run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config)- elif kv_seq_len <= 1024:- # a8w8 ps=2 for bs=32, bs=256- kv_fp8_data, kv_fp8_scale = kv_data["fp8"]- splits = 16 if batch_size <= 32 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 — splits=8 for all (proven optimal)- kv_fp8_data, kv_fp8_scale = kv_data["fp8"]- if batch_size <= 4:- splits = 16- else:- splits = 8- _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)-- return output+ return o
scrolls · 281 diff lines total
Best evidence level for this revision: reported
JSON