submission 691466
nkh5845 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 139 lines, June 9 Researcher Reciprocity License v1.0.
mixed-mla_89170b4_leaderboard.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-691466?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:db459c96e4499ef81330d33eff33927310d25cdab629d7179b3ed1b7d167a6f3
license declaredunknown
license concludedunknown
authorsnkh5845
imported2026-08-26
Kernel source
mixed-mla_89170b4_leaderboard.py139 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import torch
from task import input_t, output_t
from utils import make_match_reference
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
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_FP8_MAX = torch.finfo(FP8_DTYPE).max
_FP8_MIN = torch.finfo(FP8_DTYPE).min
# ---------------------------------------------------------------------------
# All caches keyed by shape tuples — never by data_ptr
# ---------------------------------------------------------------------------
_meta_cache = {}
_meta_filled = {}
_kv_idx_cache = {}
_out_cache = {}
_klp_cache = {}
# ---------------------------------------------------------------------------
# Metadata: allocate once, fill once per shape key
# No .item() calls — all shape info from config dict (pure Python, no GPU sync)
# ---------------------------------------------------------------------------
def _get_meta(bs, kv_seq, mql, nq, nkv, qdt, kvdt, qoi, kvi, nks):
ck = (bs, kv_seq, mql, nq, nkv, str(qdt), str(kvdt), nks)
if ck not in _meta_cache:
info = get_mla_metadata_info_v1(
bs, mql, nq, qdt, kvdt,
is_sparse=False, fast_mode=True,
num_kv_splits=nks, intra_batch_mode=True,
)
_meta_cache[ck] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = _meta_cache[ck]
if ck not in _meta_filled:
klp_key = (bs, kv_seq)
if klp_key not in _klp_cache:
_klp_cache[klp_key] = torch.full((bs,), kv_seq, dtype=torch.int32, device="cuda")
klp = _klp_cache[klp_key]
get_mla_metadata_v1(
qoi, kvi, klp, nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE, kv_granularity=64,
max_seqlen_qo=mql, uni_seqlen_qo=mql,
fast_mode=True,
max_split_per_batch=nks,
intra_batch_mode=True, dtype_q=qdt, dtype_kv=kvdt,
)
_meta_filled[ck] = True
return {"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
"reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm}
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
mql = config["q_seq_len"]
kv_seq = config["kv_seq_len"]
total_q = q.shape[0]
tot_kv = bs * kv_seq # NO .item() — pure Python, no GPU sync
# Output buffer cache
okey = (total_q, nq, dv)
if okey not in _out_cache:
_out_cache[okey] = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
o = _out_cache[okey]
# KV indices cache
if tot_kv not in _kv_idx_cache:
_kv_idx_cache[tot_kv] = torch.arange(tot_kv, dtype=torch.int32, device="cuda")
kv_idx = _kv_idx_cache[tot_kv]
# Per-shape config: (num_kv_splits, use_bf16)
# bf16 skips ~30us quantize_fp8 overhead — wins when not bandwidth-bound
if bs <= 4:
nks, use_bf16 = (16, True) if kv_seq <= 1024 else (32, True)
elif bs <= 32:
nks, use_bf16 = (8, True) if kv_seq <= 1024 else (16, False)
elif bs <= 64:
nks, use_bf16 = (8 if kv_seq <= 1024 else 4, False)
else:
nks, use_bf16 = (4, False)
if use_bf16:
kv_buf = kv_data["bf16"]
q_in, q_scale, kv_scale = q, None, None
else:
kv_fp8, kv_scale = kv_data["fp8"]
kv_buf = kv_fp8
# Fused quantize: fewer intermediate tensors
amax = q.abs().amax().clamp(min=1e-12)
scale = amax / _FP8_MAX
q_in = (q / scale).clamp(min=_FP8_MIN, max=_FP8_MAX).to(FP8_DTYPE)
q_scale = scale.to(torch.float32).reshape(1)
kv4d = kv_buf.view(tot_kv, PAGE_SIZE, nkv, kv_buf.shape[-1])
meta = _get_meta(bs, kv_seq, mql, nq, nkv, q_in.dtype, kv_buf.dtype,
qo_indptr, kv_indptr, nks)
# KV last page length — cached, no .item()
klp_key = (bs, kv_seq)
if klp_key not in _klp_cache:
_klp_cache[klp_key] = torch.full((bs,), kv_seq, dtype=torch.int32, device="cuda")
klp = _klp_cache[klp_key]
mla_decode_fwd(
q_in.view(-1, nq, dq), kv4d, o,
qo_indptr, kv_indptr, kv_idx, klp, mql,
page_size=PAGE_SIZE, nhead_kv=nkv,
sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=nks,
q_scale=q_scale, kv_scale=kv_scale,
intra_batch_mode=True, **meta,
)
return o
check_implementation = make_match_reference(custom_kernel, rtol=1e-01, atol=1e-01)
scrolls · 139 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