submission 651367
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 155 lines, June 9 Researcher Reciprocity License v1.0.
v108.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-651367?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:e469b896e184c99fc52593e6dc8196107ba48578e4a9ca84e86652faee301695
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
v90 bypassed mla_decode_fwd only for a8w8 persistent. v108 also bypasses for:Kernel source
v108.py155 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v108 — Full bypass: ALL paths use direct stage1_asm (+ reduce where needed).
v90 bypassed mla_decode_fwd only for a8w8 persistent. v108 also bypasses for:
- BF16 persistent (bs<=4/kv<=1024): saves ~3µs from torch.empty elimination
- FP8 NP splits=1 (bs=64/kv<=1024): saves ~2µs, MAYBE_FINAL_OUT=True (no stage2)
"""
from task import input_t, output_t
import torch
import aiter as _aiter
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, 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 _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 persistent DIRECT (bs<=4/kv<=1024) ----
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)
rpm_sz = m['reduce_partial_map'].size(0)
inter = _gi(('bfp_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(
q, kv_bf16.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,
)
_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
# ---- FP8 NP splits=1 DIRECT (bs=64/kv<=1024) ----
# MAYBE_FINAL_OUT=True (v_dim=512<=512, mgc=0): stage1 writes directly to o
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),
})
inter = _gi(('fn_i', tq, nh), lambda: {
'attn_lse': torch.empty((tq, 1, nh, 1), dtype=torch.float32, device=dev),
})
logits = o.view(tq, 1, nh, vd) # View of output — stage1 writes here directly
_aiter.mla_decode_stage1_asm_fwd(
q8, kd.unsqueeze(1), qo_indptr, kv_indptr,
t['ki'], t['kl'], t['si'],
None, None, None,
1, 1, 1, sms,
logits, inter['attn_lse'], o,
t['qs'], ks,
)
# NO stage2 reduce needed — MAYBE_FINAL_OUT=True
return o
# ---- a8w8 persistent DIRECT (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)
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)
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 · 155 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 641982.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """MLA v86 — Hybrid: mla_decode_fwd for NP paths, direct stage1+reduce for persistent.+ """MLA v108 — Full bypass: ALL paths use direct stage1_asm (+ reduce where needed).- 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.+ v90 bypassed mla_decode_fwd only for a8w8 persistent. v108 also bypasses for:+ - BF16 persistent (bs<=4/kv<=1024): saves ~3µs from torch.empty elimination+ - FP8 NP splits=1 (bs=64/kv<=1024): saves ~2µs, MAYBE_FINAL_OUT=True (no stage2)"""from task import input_t, output_timport torchimport aiter as _aiter- from aiter.mla import mla_decode_fwdfrom aiter import dtypes as aiter_dtypesfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1⋯ 7 unchanged linesif v is None: v = fn(); _tc[key] = vreturn v- def _gm(bs, tot, nh, ns, ps, qi, ki, kl, dev):- key = (bs, tot, nh, ns, ps)+ 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- 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)+ 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 = workget_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)+ 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⋯ 22 unchanged lineskd, ks = kv_data["fp8"]tot = kd.shape[0]- # ---- BF16 non-persistent via mla_decode_fwd (bs<=4, kv<=1024) ----+ # ---- BF16 persistent DIRECT (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),+ 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),})- 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)+ m = _gm(bs, tot, nh, ns, ps, qo_indptr, t['kip'], t['kl'], dev, torch.bfloat16, torch.bfloat16)+ rpm_sz = m['reduce_partial_map'].size(0)+ inter = _gi(('bfp_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(+ q, kv_bf16.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,+ )+ _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- # ---- FP8 NP splits=1 via mla_decode_fwd (bs=64, kv<=1024) ----+ # ---- FP8 NP splits=1 DIRECT (bs=64/kv<=1024) ----+ # MAYBE_FINAL_OUT=True (v_dim=512<=512, mgc=0): stage1 writes directly to oif kvl <= 1024 and bs == 64:q8 = q.to(torch.float8_e4m3fn)t = _gt(('fn', bs, tot), lambda: {⋯ 2 unchanged lines'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)+ inter = _gi(('fn_i', tq, nh), lambda: {+ 'attn_lse': torch.empty((tq, 1, nh, 1), dtype=torch.float32, device=dev),+ })+ logits = o.view(tq, 1, nh, vd) # View of output — stage1 writes here directly+ _aiter.mla_decode_stage1_asm_fwd(+ q8, kd.unsqueeze(1), qo_indptr, kv_indptr,+ t['ki'], t['kl'], t['si'],+ None, None, None,+ 1, 1, 1, sms,+ logits, inter['attn_lse'], o,+ t['qs'], ks,+ )+ # NO stage2 reduce needed — MAYBE_FINAL_OUT=Truereturn o- # ---- a8w8 persistent: DIRECT stage1_asm + reduce_v1 (bypass mla_decode_fwd) ----+ # ---- a8w8 persistent DIRECT (all other shapes) ----q8 = q.to(torch.float8_e4m3fn)if kvl <= 1024:ps, ns = 2, (16 if bs <= 32 else 8)⋯ 7 unchanged lines'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)+ m = _gm(bs, tot, nh, ns, ps, qo_indptr, t['kip'], t['kl'], dev, aiter_dtypes.fp8, aiter_dtypes.fp8)- # 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),⋯ 2 unchanged lines_aiter.mla_decode_stage1_asm_fwd(q8, kd.view(np, ps, 1, 576), qo_indptr, t['kip'],- t['ki'], t['kl'],- None,+ 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,
scrolls · 140 diff lines total
Best evidence level for this revision: reported
JSON