submission 664856
Infatoshi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 61 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-664856?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:94e6799b98676993c125b8d352f6870cd2ef50d2569c2d751b71f5eda75175ba
license declaredunknown
license concludedunknown
authorsInfatoshi
imported2026-08-26
Kernel source
submission.py61 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""Split on kv_len: a16w8 for kv<=1024, a8w8+dynquant for kv>=8192."""
import torch
from task import input_t, output_t
NUM_HEADS = 16; NUM_KV_HEADS = 1; QK_HEAD_DIM = 576; V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
from aiter.ops.quant import dynamic_per_tensor_quant
FP8_DTYPE = aiter_dtypes.fp8
_cache = {}
def _get_config(bs, kvl):
if kvl <= 1024:
ns = 8 if bs <= 32 else 4
return (ns, False, 2, True) # a16w8
else: # kv >= 8192
if bs <= 4:
return (32, False, 2, True) # a16w8 + high ns (40us)
elif bs <= 32:
return (8, True, 1, False) # a8w8 + dynquant (80us)
elif bs <= 64:
return (16, True, 1, False) # a8w8 (130us)
else:
return (32, True, 1, False) # a8w8 (310us)
def _get_or_build(bs, kvl, qd, kvd, qo, kvi, ns, dev, ps, fm):
key = (bs, kvl, ns, qd, ps, fm)
if key in _cache: return _cache[key]
tkv = bs * kvl
kl = (kvi[1:] - kvi[:-1]).to(torch.int32)
ki = torch.arange(tkv, dtype=torch.int32, device=dev)
info = get_mla_metadata_info_v1(bs, 1, NUM_HEADS, qd, kvd, is_sparse=False, fast_mode=fm, num_kv_splits=ns, intra_batch_mode=True)
w = [torch.empty(s, dtype=t, device=dev) for s, t in info]
wm, wi, ws, ri, rf, rp = w
get_mla_metadata_v1(qo, kvi, kl, NUM_HEADS//NUM_KV_HEADS, NUM_KV_HEADS, True, wm, ws, wi, ri, rf, rp, page_size=ps, kv_granularity=max(ps,16), max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=fm, max_split_per_batch=ns, intra_batch_mode=True, dtype_q=qd, dtype_kv=kvd)
e = {"meta": {"work_meta_data":wm,"work_indptr":wi,"work_info_set":ws,"reduce_indptr":ri,"reduce_final_map":rf,"reduce_partial_map":rp}, "kl":kl, "ki":ki, "out":torch.empty((bs,NUM_HEADS,V_HEAD_DIM),dtype=torch.bfloat16,device=dev)}
_cache[key] = e
return e
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = int(config["batch_size"]); kvl = int(config["kv_seq_len"])
ns, use_a8w8, ps, fm = _get_config(bs, kvl)
kv_fp8, kv_scale = kv_data["fp8"]
kv_4d = kv_fp8.view(kv_fp8.shape[0], 1, NUM_KV_HEADS, kv_fp8.shape[-1])
if use_a8w8:
bkey = ("dq", q.numel())
if bkey not in _cache:
_cache[bkey] = (torch.empty_like(q, dtype=FP8_DTYPE), torch.empty(1, dtype=torch.float32, device=q.device))
qi, qs = _cache[bkey]
dynamic_per_tensor_quant(qi, q, qs)
qv = qi.view(-1, NUM_HEADS, QK_HEAD_DIM)
else:
qv = q.view(-1, NUM_HEADS, QK_HEAD_DIM); qs = None
c = _get_or_build(bs, kvl, qv.dtype, kv_fp8.dtype, qo_indptr, kv_indptr, ns, q.device, ps, fm)
mla_decode_fwd(qv, kv_4d, c["out"], qo_indptr, kv_indptr, c["ki"], c["kl"], 1, page_size=ps, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE, logit_cap=0.0, num_kv_splits=ns, q_scale=qs, kv_scale=kv_scale, intra_batch_mode=True, **c["meta"])
return c["out"]
scrolls · 61 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