submission 673461
Maxwell Cipher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 111 lines, June 9 Researcher Reciprocity License v1.0.
mla_v67.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-673461?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:cfe14e9550c88fbe55f8348657e95b21232136510b845208cb6422d636a6114a
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
"""v67: a16w8 persistent pg2 — bf16 Q + fp8 KV, page_size=2, persistent mode.Kernel source
mla_v67.py111 lines
#!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)."""
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
_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, nt, qo_indptr, dev):
k = ("a16w8p", bs, sl)
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=False,
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=False,
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 <= 4 and sl <= 1024:
# Tiny: bf16 non-persistent pg1 (proven by v41)
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:
# All other: a16w8 persistent pg2
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)
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 · 111 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 663004.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """v41: Full bf16 non-persistent MLA decode.- Skip fp8 quantization entirely — single kernel launch with bf16 Q and KV.- Overhead savings (no quant, no metadata, no reduce) outweigh 2x bandwidth cost."""+ """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)."""import torchfrom task import input_t, output_tfrom 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_DIM = 576- _V_DIM = 512- _SM_SC = 1.0 / (_QK_DIM ** 0.5)+ _QK = 576+ _VD = 512+ _SC = 1.0 / (_QK ** 0.5)+ _PS = 2+ _NSPLIT = 32- _shape_bufs = {}+ _cache = {}- def _get_shape_buffers(bs, seq_len, dev):- k = (bs, seq_len)- if k not in _shape_bufs:- n = bs * seq_len- page_ids = torch.arange(n, dtype=torch.int32, device=dev)- seq_lens = torch.full((bs,), seq_len, dtype=torch.int32, device=dev)- _shape_bufs[k] = (page_ids, seq_lens)- return _shape_bufs[k]+ 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, nt, qo_indptr, dev):+ k = ("a16w8p", bs, sl)+ 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=False,+ 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=False,+ 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"])- seq_len = int(config["kv_seq_len"])- n_tokens = q.shape[0]+ sl = int(config["kv_seq_len"])+ nt = q.shape[0]+ dev = q.device- q_reshaped = q.view(n_tokens, _NH, _QK_DIM)- kv_raw = kv_data["bf16"]- kv_paged = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])+ q_r = q.view(nt, _NH, _QK)+ out = torch.empty((nt, _NH, _VD), dtype=torch.bfloat16, device=dev)- page_ids, seq_lens = _get_shape_buffers(bs, seq_len, q.device)- out = torch.empty((n_tokens, _NH, _V_DIM), dtype=torch.bfloat16, device=q.device)+ if bs <= 4 and sl <= 1024:+ # Tiny: bf16 non-persistent pg1 (proven by v41)+ 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:+ # All other: a16w8 persistent pg2+ 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)+ 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,+ )- mla_decode_fwd(- q_reshaped, kv_paged, out,- qo_indptr, kv_indptr,- page_ids, seq_lens, 1,- page_size=1, nhead_kv=_NKV, sm_scale=_SM_SC,- intra_batch_mode=False,- )return out
scrolls · 140 diff lines total
Best evidence level for this revision: reported
JSON