submission 747564
Chivier · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 95 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747564?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:d329016f8140217bc487434b75596c4ea6b114035649a8addcd033cfb781caaa
license declaredunknown
license concludedunknown
authorsChivier
imported2026-08-26
Kernel source
submission.py95 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA V28 — Hybrid: Triton bf16 (kv≤1024) + V18 FP8 ASM (kv>1024).
V26 showed Triton bf16 wins on kv=1024: 21-44µs vs ASM 27-49µs.
V18 FP8 ASM proven on kv=8192: 36-312µs (ranked 76µs).
V28 uses V18's exact ASM path (proven metadata caching) for large kv."""
import torch
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1, dtypes as _d
from aiter.ops.quant import dynamic_per_tensor_quant as _dq
from aiter.ops.triton.attention.mla_decode_rope import decode_attention_fwd_grouped_rope
_F = _d.fp8
_S = 1.0 / (576 ** 0.5)
_SP = {
(4, 1024): 8, (4, 8192): 16,
(32, 1024): 8, (32, 8192): 8,
(64, 1024): 4, (64, 8192): 4,
(256, 1024): 1, (256, 8192): 1,
}
_ct = {} # Triton bf16 cache
_ca = {} # ASM FP8 cache (V18-style)
def custom_kernel(data):
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
f8, ks = kv_data["fp8"]
T = f8.shape[0]
kl = T // bs
ns = _SP.get((bs, kl), 8)
if kl <= 1024:
# ── Triton bf16 path (20-44µs for kv=1024) ──
kv_bf16 = kv_data["bf16"]
k = (bs, T, ns)
if k not in _ct:
ki = torch.arange(T, dtype=torch.int32, device="cuda")
o = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
al = torch.empty((bs, 16, ns, 513), dtype=torch.float32, device="cuda")
dummy = torch.empty(0, device="cuda")
_ct[k] = (ki, o, al, dummy)
ki, o, al, dummy = _ct[k]
decode_attention_fwd_grouped_rope(
q.view(bs, 16, 576), kv_bf16, kv_bf16[:, :, :512], o,
kv_indptr, ki, dummy, 512, 64, dummy, dummy,
al, ns, _S,
logit_cap=0.0, use_rope=False, is_neox_style=False,
)
return o
else:
# ── V18 FP8 ASM path (proven 36-312µs for kv=8192) ──
k = (bs, T, ns)
if k not in _ca:
ki = torch.arange(T, dtype=torch.int32, device="cuda")
kl_t = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
o = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
qf = torch.empty(q.shape, dtype=_F, device="cuda")
qs = torch.empty(1, dtype=torch.float32, device="cuda")
info = get_mla_metadata_info_v1(
bs, 1, 16, _F, _F,
is_sparse=False, fast_mode=False,
num_kv_splits=ns, intra_batch_mode=True)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kl_t,
16, 1, True, wm, wis, wi, ri, rfm, rpm,
page_size=1, kv_granularity=16,
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=False, max_split_per_batch=ns,
intra_batch_mode=True, dtype_q=_F, dtype_kv=_F)
_ca[k] = (ki, kl_t, o, qf, qs, wm, wi, wis, ri, rfm, rpm)
ki, kl_t, o, qf, qs, wm, wi, wis, ri, rfm, rpm = _ca[k]
_dq(qf, q, qs)
kv4 = f8.view(T, 1, 1, 576)
mla_decode_fwd(
qf.view(-1, 16, 576), kv4, o,
qo_indptr, kv_indptr, ki, kl_t,
1, page_size=1, nhead_kv=1,
sm_scale=_S, logit_cap=0.0,
num_kv_splits=ns,
q_scale=qs, kv_scale=ks,
intra_batch_mode=True,
work_meta_data=wm, work_indptr=wi, work_info_set=wis,
reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
)
return o
scrolls · 95 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON