submission 682562
Maxwell Cipher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 123 lines, June 9 Researcher Reciprocity License v1.0.
mla_v93.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-682562?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:8ae5968021b92eb890c83c73ef72bb604f980747a510a145fccd5402ae70dac9
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Kernel source
mla_v93.py123 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""v93: Per-shape fast_mode — cherry-pick best of v67 and v90.
fast_mode=True helps: (4,8192) -3us, (64,1024) -3.7us
fast_mode=False helps: (32,8192) -1us, (64,8192) -2.5us
Neutral: (256,*), (4,1024), (32,1024)
Use True for small batch+large kv and medium batch+small kv.
Use False for large batch or large kv+large batch.
"""
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
FP8 = aiter_dtypes.fp8
_NH = 16
_NKV = 1
_QK = 576
_VD = 512
_SC = 1.0 / (_QK ** 0.5)
_PS = 2
_NSPLIT = 32
# Per-shape: (fast_mode,) — True where it helps, False where it hurts
_FM = {
(4, 8192): True,
(64, 1024): True,
# Everything else uses False (v67 default, better for large shapes)
}
_cache = {}
def _get_pg1(bs, sl, dev):
k = ("pg1", bs, sl)
if k not in _cache:
_cache[k] = (
torch.arange(bs * sl, dtype=torch.int32, device=dev),
torch.full((bs,), sl, dtype=torch.int32, device=dev),
)
return _cache[k]
def _get_a16w8_persist(bs, sl, fm, qo_indptr, dev):
k = ("a16w8p", bs, sl, fm)
if k not in _cache:
npages = (bs * sl) // _PS
kv_idx = torch.arange(npages, dtype=torch.int32, device=dev)
ki = torch.arange(bs + 1, dtype=torch.int32, device=dev) * (sl // _PS)
lp = torch.full((bs,), _PS, dtype=torch.int32, device=dev)
info = get_mla_metadata_info_v1(
bs, 1, _NH, torch.bfloat16, FP8,
is_sparse=False, fast_mode=fm,
num_kv_splits=_NSPLIT, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(
qo_indptr, ki, lp,
_NH // _NKV, _NKV, True,
wm, wis, wi, ri, rfm, rpm,
page_size=_PS,
kv_granularity=max(_PS, 16),
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=fm,
max_split_per_batch=_NSPLIT,
intra_batch_mode=True,
dtype_q=torch.bfloat16, dtype_kv=FP8,
)
meta = {
"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
"reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm,
}
_cache[k] = (meta, kv_idx, ki, lp)
return _cache[k]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = int(config["batch_size"])
sl = int(config["kv_seq_len"])
nt = q.shape[0]
dev = q.device
q_r = q.view(nt, _NH, _QK)
out = torch.empty((nt, _NH, _VD), dtype=torch.bfloat16, device=dev)
if bs <= 32 and sl <= 1024:
# bf16 non-persistent pg1
kv_raw = kv_data["bf16"]
kv_4d = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])
pg, lp = _get_pg1(bs, sl, dev)
mla_decode_fwd(
q_r, kv_4d, out, qo_indptr, kv_indptr,
pg, lp, 1,
page_size=1, nhead_kv=_NKV, sm_scale=_SC,
intra_batch_mode=False,
)
else:
# a16w8 persist pg2 with per-shape fast_mode
fm = _FM.get((bs, sl), False)
kv_fp8, kv_scale = kv_data["fp8"]
kv_4d = kv_fp8.view(-1, _PS, _NKV, kv_fp8.shape[-1])
meta, kv_idx, ki, lp = _get_a16w8_persist(bs, sl, fm, qo_indptr, dev)
mla_decode_fwd(
q_r, kv_4d, out, qo_indptr, ki,
kv_idx, lp, 1,
page_size=_PS, nhead_kv=_NKV, sm_scale=_SC,
logit_cap=0.0, num_kv_splits=_NSPLIT,
kv_scale=kv_scale,
intra_batch_mode=True, **meta,
)
return out
scrolls · 123 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 673461.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """v67: a16w8 persistent pg2 — bf16 Q + fp8 KV, page_size=2, persistent mode.- Non-persistent a16w8 pg2 has no ASM kernel (ps:0 error).- Persistent mode dispatches through metadata scheduler which supports pg2.- bf16 non-persistent pg1 for tiny shapes (proven by v41)."""+ """v93: Per-shape fast_mode — cherry-pick best of v67 and v90.+ fast_mode=True helps: (4,8192) -3us, (64,1024) -3.7us+ fast_mode=False helps: (32,8192) -1us, (64,8192) -2.5us+ Neutral: (256,*), (4,1024), (32,1024)++ Use True for small batch+large kv and medium batch+small kv.+ Use False for large batch or large kv+large batch.+ """+import torchfrom task import input_t, output_tfrom aiter.mla import mla_decode_fwd⋯ 9 unchanged lines_PS = 2_NSPLIT = 32+ # Per-shape: (fast_mode,) — True where it helps, False where it hurts+ _FM = {+ (4, 8192): True,+ (64, 1024): True,+ # Everything else uses False (v67 default, better for large shapes)+ }+_cache = {}⋯ 7 unchanged linesreturn _cache[k]- def _get_a16w8_persist(bs, sl, nt, qo_indptr, dev):- k = ("a16w8p", bs, sl)+ def _get_a16w8_persist(bs, sl, fm, qo_indptr, dev):+ k = ("a16w8p", bs, sl, fm)if k not in _cache:npages = (bs * sl) // _PSkv_idx = torch.arange(npages, dtype=torch.int32, device=dev)⋯ 2 unchanged linesinfo = get_mla_metadata_info_v1(bs, 1, _NH, torch.bfloat16, FP8,- is_sparse=False, fast_mode=False,+ is_sparse=False, fast_mode=fm,num_kv_splits=_NSPLIT, intra_batch_mode=True,)work = [torch.empty(s, dtype=t, device=dev) for s, t in info]⋯ 6 unchanged linespage_size=_PS,kv_granularity=max(_PS, 16),max_seqlen_qo=1, uni_seqlen_qo=1,- fast_mode=False,+ fast_mode=fm,max_split_per_batch=_NSPLIT,intra_batch_mode=True,dtype_q=torch.bfloat16, dtype_kv=FP8,⋯ 17 unchanged linesq_r = q.view(nt, _NH, _QK)out = torch.empty((nt, _NH, _VD), dtype=torch.bfloat16, device=dev)- if bs <= 4 and sl <= 1024:- # Tiny: bf16 non-persistent pg1 (proven by v41)+ if bs <= 32 and sl <= 1024:+ # bf16 non-persistent pg1kv_raw = kv_data["bf16"]kv_4d = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])pg, lp = _get_pg1(bs, sl, dev)⋯ 4 unchanged linesintra_batch_mode=False,)else:- # All other: a16w8 persistent pg2+ # a16w8 persist pg2 with per-shape fast_mode+ fm = _FM.get((bs, sl), False)kv_fp8, kv_scale = kv_data["fp8"]kv_4d = kv_fp8.view(-1, _PS, _NKV, kv_fp8.shape[-1])- meta, kv_idx, ki, lp = _get_a16w8_persist(bs, sl, nt, qo_indptr, dev)+ meta, kv_idx, ki, lp = _get_a16w8_persist(bs, sl, fm, qo_indptr, dev)mla_decode_fwd(q_r, kv_4d, out, qo_indptr, ki,kv_idx, lp, 1,page_size=_PS, nhead_kv=_NKV, sm_scale=_SC,logit_cap=0.0, num_kv_splits=_NSPLIT,kv_scale=kv_scale,- intra_batch_mode=True,- **meta,+ intra_batch_mode=True, **meta,)return out
scrolls · 97 diff lines total
Best evidence level for this revision: reported
JSON