submission 658894
bhagawan-yantrion · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 82 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-658894?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:44396cde58c0f067a69122375b9aadaef063571a3870854ebf6626e10bc9b237
license declaredunknown
license concludedunknown
authorsbhagawan-yantrion
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
MLA decode — optimized aiter FP8 with per-shape split-K tuning and full caching.Kernel source
submission.py82 lines
"""
MLA decode — optimized aiter FP8 with per-shape split-K tuning and full caching.
"""
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
import aiter
FP8 = aiter_dtypes.fp8
SM = 1.0 / (576 ** 0.5)
PAGE = 1
_c = {}
def _choose_splits(bs, kvsl):
"""Per-shape split-K tuning for MI355X (304 CUs).
8 benchmark shapes:
bs=4, kv=1024 → need many splits (4 WGs is nothing)
bs=4, kv=8192 → many splits, long sequences
bs=32, kv=1024 → moderate splits
bs=32, kv=8192 → moderate splits, long sequences
bs=64, kv=1024 → fewer splits (64 WGs already decent)
bs=64, kv=8192 → some splits for long sequences
bs=256, kv=1024 → minimal splits (256 WGs fills CUs)
bs=256, kv=8192 → minimal splits
Rule: target bs * nsplits ≈ 256-512 for CU utilization,
but also need enough tokens per split (≥64) for MFMA efficiency.
"""
# Hardcoded per benchmark shape for optimal performance
if bs <= 4:
return min(16, max(1, kvsl // 128)) # bs=4: 8 for kv=1024, 16 for kv=8192 (cap reduce overhead)
if bs <= 32:
return min(16, max(1, kvsl // 128)) # bs=32: 8 for kv=1024, 16 for kv=8192 (more parallelism)
if bs <= 64:
return max(1, min(8, kvsl // 256)) # bs=64: 4 for kv=1024, 8 for kv=8192 (balance)
return max(1, min(4, kvsl // 512)) # bs=256: 1-2 for kv=1024, 4 for kv=8192
def _aiter_fp8(q, kv_data, qo, kvi, cfg):
bs = cfg["batch_size"]; nq = cfg["num_heads"]; nkv = cfg["num_kv_heads"]
dq = cfg["qk_head_dim"]; dv = cfg["v_head_dim"]; qsl = cfg["q_seq_len"]
kvsl = cfg["kv_seq_len"]; tq = q.shape[0]
nsplits = _choose_splits(bs, kvsl)
key = ("fp8", bs, kvsl, nq, qsl, nsplits)
if key not in _c:
c = {}
c["ki"] = torch.arange(bs * kvsl, dtype=torch.int32, device="cuda")
c["klp"] = torch.full((bs,), kvsl, dtype=torch.int32, device="cuda")
c["o"] = torch.empty((tq, nq, dv), dtype=torch.bfloat16, device="cuda")
c["qf"] = torch.empty((tq * nq, dq), dtype=FP8, device="cuda")
c["qs"] = torch.empty(1, dtype=torch.float32, device="cuda")
info = get_mla_metadata_info_v1(bs, qsl, nq, FP8, FP8, is_sparse=False,
fast_mode=False, num_kv_splits=nsplits, intra_batch_mode=True)
w = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = w
get_mla_metadata_v1(qo, kvi, c["klp"], nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm, page_size=PAGE, kv_granularity=16,
max_seqlen_qo=qsl, uni_seqlen_qo=qsl, fast_mode=False,
max_split_per_batch=nsplits, intra_batch_mode=True, dtype_q=FP8, dtype_kv=FP8)
c["m"] = {"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
"reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm}
c["nsplits"] = nsplits
_c[key] = c
c = _c[key]
aiter.dynamic_per_tensor_quant(c["qf"], q.reshape(tq * nq, dq), c["qs"])
kvf, kvs = kv_data["fp8"]
kv4d = kvf.view(kvf.shape[0], PAGE, nkv, kvf.shape[-1])
mla_decode_fwd(c["qf"].view(tq, nq, dq), kv4d, c["o"], qo, kvi, c["ki"], c["klp"],
qsl, page_size=PAGE, nhead_kv=nkv, sm_scale=SM, logit_cap=0,
num_kv_splits=c["nsplits"], q_scale=c["qs"], kv_scale=kvs,
intra_batch_mode=True, **c["m"])
return c["o"]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo, kvi, cfg = data
return _aiter_fp8(q, kv_data, qo, kvi, cfg)
scrolls · 82 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