submission 674864
yanchaomei · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 140 lines, June 9 Researcher Reciprocity License v1.0.
submission_mla_bf16.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-674864?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:38566e524a90171c76a9a2e6be241337ca45be283d9eeaae1129f1735a90704a
license declaredunknown
license concludedunknown
authorsyanchaomei
imported2026-08-26
Kernel source
submission_mla_bf16.py140 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode — BF16 path (no FP8 quant overhead).
Uses mla_dec_stage1_bf16_a16w16_subQ16_mqa16.co kernel.
Trade: 2x more KV bandwidth vs 5µs saved on Q quant.
For small batches, quant overhead dominates → BF16 wins.
"""
import torch
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import (
get_mla_metadata_info_v1, get_mla_metadata_v1,
mla_decode_stage1_asm_fwd, mla_reduce_v1,
per_tensor_quant_hip,
)
from aiter.mla import mla_decode_fwd
FP8 = aiter_dtypes.fp8
BF16 = torch.bfloat16
NH, NKV, QKD, VD = 16, 1, 576, 512
SM = 1.0 / (QKD ** 0.5)
CU = 256
_c = {}
def _splits_fp8(bs, kl):
oh = 84.1
best_s, best = 1, float("-inf")
for s in range(1, 17):
w = (bs * s + CU - 1) // CU
score = bs * s / (w * CU) * kl / (kl + oh * s)
if score > best:
best, best_s = score, s
return min(best_s, max(1, (kl + 127) // 128))
def _splits_bf16(bs, kl):
oh = 42.0 # BF16 has lower overhead per split (bigger tiles)
best_s, best = 1, float("-inf")
for s in range(1, 17):
w = (bs * s + CU - 1) // CU
score = bs * s / (w * CU) * kl / (kl + oh * s)
if score > best:
best, best_s = score, s
return min(best_s, max(1, (kl + 127) // 128))
def _init_fp8(bs, qs, kl, qoi, kvi):
ns = _splits_fp8(bs, kl)
ki = torch.arange(bs*kl, dtype=torch.int32, device="cuda")
klp = (kvi[1:]-kvi[:-1]).to(torch.int32)
info = get_mla_metadata_info_v1(bs, qs, NH, FP8, FP8,
is_sparse=False, fast_mode=False, num_kv_splits=ns, 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(qoi,kvi,klp, NH,NKV,True, wm,wis,wi,ri,rfm,rpm,
page_size=1, kv_granularity=16, max_seqlen_qo=qs, uni_seqlen_qo=qs,
fast_mode=False, max_split_per_batch=ns, intra_batch_mode=True,
dtype_q=FP8, dtype_kv=FP8)
np_ = rpm.numel()
return {"ns":ns, "wm":wm,"wi":wi,"wis":wis,"ri":ri,"rfm":rfm,"rpm":rpm,
"ki":ki,"klp":klp,
"sd":torch.empty((np_*qs,1,NH,VD),dtype=torch.float32,device="cuda"),
"sl":torch.empty((np_*qs,1,NH,1),dtype=torch.float32,device="cuda")}
def _init_bf16(bs, qs, kl, qoi, kvi):
ns = _splits_bf16(bs, kl)
ki = torch.arange(bs*kl, dtype=torch.int32, device="cuda")
klp = (kvi[1:]-kvi[:-1]).to(torch.int32)
# BF16 dtype for metadata
info = get_mla_metadata_info_v1(bs, qs, NH, BF16, BF16,
is_sparse=False, fast_mode=False, num_kv_splits=ns, 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(qoi,kvi,klp, NH,NKV,True, wm,wis,wi,ri,rfm,rpm,
page_size=1, kv_granularity=16, max_seqlen_qo=qs, uni_seqlen_qo=qs,
fast_mode=False, max_split_per_batch=ns, intra_batch_mode=True,
dtype_q=BF16, dtype_kv=BF16)
np_ = rpm.numel()
return {"ns":ns, "wm":wm,"wi":wi,"wis":wis,"ri":ri,"rfm":rfm,"rpm":rpm,
"ki":ki,"klp":klp,
"sd":torch.empty((np_*qs,1,NH,VD),dtype=torch.float32,device="cuda"),
"sl":torch.empty((np_*qs,1,NH,1),dtype=torch.float32,device="cuda")}
# Use BF16 for small batches (quant overhead dominates), FP8 for large
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs, kl = config["batch_size"], config["kv_seq_len"]
qs, qt = config.get("q_seq_len",1), q.shape[0]
# BF16 wins for kv=1024 with bs=32/64 (saves Q quant overhead)
use_bf16 = (kl <= 1024 and 16 <= bs <= 64)
if use_bf16:
key = ("bf16", bs, kl)
if key not in _c:
_c[key] = _init_bf16(bs, qs, kl, qo_indptr, kv_indptr)
c = _c[key]
kv_bf16 = kv_data["bf16"]
kv4 = kv_bf16.view(kv_bf16.shape[0], 1, NKV, kv_bf16.shape[-1])
q_bf16 = q.view(-1, NH, QKD)
out = torch.empty((qt, NH, VD), dtype=torch.bfloat16, device="cuda")
mla_decode_stage1_asm_fwd(
q_bf16, kv4, qo_indptr, kv_indptr,
c["ki"], c["klp"], None,
c["wm"], c["wi"], c["wis"],
qs, 1, NKV, SM,
c["sd"], c["sl"], out,
None, None, # No scales for BF16
)
mla_reduce_v1(c["sd"], c["sl"], c["ri"], c["rfm"], c["rpm"], qs, out, None)
return out
else:
key = ("fp8", bs, kl)
if key not in _c:
_c[key] = _init_fp8(bs, qs, kl, qo_indptr, kv_indptr)
c = _c[key]
kf, ks = kv_data["fp8"]
qf, qsc = per_tensor_quant_hip(q.view(-1,NH,QKD), quant_dtype=FP8)
qsc = qsc.reshape(1)
kv4 = kf.view(kf.shape[0],1,NKV,kf.shape[-1])
out = torch.empty((qt, NH, VD), dtype=torch.bfloat16, device="cuda")
mla_decode_stage1_asm_fwd(
qf.view(-1,NH,QKD), kv4, qo_indptr, kv_indptr,
c["ki"], c["klp"], None,
c["wm"], c["wi"], c["wis"],
qs, 1, NKV, SM,
c["sd"], c["sl"], out, qsc, ks,
)
mla_reduce_v1(c["sd"], c["sl"], c["ri"], c["rfm"], c["rpm"], qs, out, None)
return out
scrolls · 140 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