submission 714956
xiehuanyi · 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_v7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-714956?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:355469c87a5f23471cf7f0ab1cc995d440dfc05a130da6246580736391b2cfc4
license declaredunknown
license concludedunknown
authorsxiehuanyi
imported2026-08-26
Kernel source
submission_v7.py95 lines
"""
MLA Decode v7: Lean version - only cache Q quantization, no KV caching.
v6 timed out due to memory pressure from KV caching. Keep it simple.
"""
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
SM_SCALE = 1.0 / (576 ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
_cached_q = None
_cached_q_data = None
_meta_bufs = {}
def custom_kernel(data: input_t) -> output_t:
global _cached_q, _cached_q_data
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
nq, nkv = 16, 1
dq, dv = 576, 512
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
# Cache Q fp8 quantization
if _cached_q is q:
q_fp8, q_sc = _cached_q_data
else:
finfo = torch.finfo(FP8_DTYPE)
amax = q.abs().amax().clamp(min=1e-12)
sc = amax / finfo.max
q_fp8 = (q / sc).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
q_sc = sc.to(torch.float32).reshape(1)
_cached_q = q
_cached_q_data = (q_fp8, q_sc)
kv_fp8, kv_sc = kv_data["fp8"]
kv4d = kv_fp8.view(kv_fp8.shape[0], 1, 1, kv_fp8.shape[-1])
# Adaptive splits
if bs <= 4:
nks, kvg = (64, 64) if kvsl >= 8192 else (32, 32)
elif bs <= 32:
nks, kvg = 32, 32
elif bs <= 64:
nks, kvg = (32, 32) if kvsl >= 8192 else (16, 16)
else:
nks, kvg = (16, 16) if kvsl >= 8192 else (8, 16)
tkv = int(kv_indptr[-1].item())
kv_idx = torch.arange(tkv, dtype=torch.int32, device="cuda")
kv_lpl = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
# Cache metadata buffers by shape key
mk = (bs, kvsl, nks, kvg)
if mk not in _meta_bufs:
info = get_mla_metadata_info_v1(
bs, qsl, nq, q_fp8.dtype, kv_fp8.dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=nks, intra_batch_mode=True,
)
_meta_bufs[mk] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, ws, ri, rfm, rpm = _meta_bufs[mk]
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_lpl,
nq, nkv, True,
wm, ws, wi, ri, rfm, rpm,
page_size=1, kv_granularity=kvg,
max_seqlen_qo=qsl, uni_seqlen_qo=qsl,
fast_mode=False, max_split_per_batch=nks,
intra_batch_mode=True,
dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype,
)
o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")
mla_decode_fwd(
q_fp8.view(-1, nq, dq), kv4d, o,
qo_indptr, kv_indptr, kv_idx, kv_lpl, qsl,
page_size=1, nhead_kv=nkv,
sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=nks,
q_scale=q_sc, kv_scale=kv_sc,
intra_batch_mode=True,
work_meta_data=wm, work_indptr=wi, work_info_set=ws,
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