submission 716406
janice.jiayao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 154 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-716406?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:728eeb60d43b97b9d7c8fee625e47523fb97ff17db5354ae8d1db17a30c68c94
license declaredunknown
license concludedunknown
authorsjanice.jiayao
imported2026-08-26
Kernel source
submission.py154 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import torch
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
FP8_DTYPE = aiter_dtypes.fp8
_finfo = torch.finfo(FP8_DTYPE)
FP8_MAX = _finfo.max
FP8_MIN = _finfo.min
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
_cache = {}
def get_config(batch_size: int, kv_seq_len: int):
"""
Dispatch table fitted from MI355X probe data.
Best observed steady-state configs:
bs=4, kv=1024 -> a16w8 split=32 fast_mode=True intra_batch=False
bs=4, kv=8192 -> a16w8 split=256 fast_mode=False intra_batch=True
bs=32, kv=1024 -> a16w8 split=256 fast_mode=True intra_batch=False
bs=32, kv=8192 -> a8w8 split=256 fast_mode=False intra_batch=True
bs=64, kv=1024 -> a16w8 split=32 fast_mode=False intra_batch=True
bs=64, kv=8192 -> a8w8 split=32 fast_mode=False intra_batch=True
bs=256, kv=1024 -> a8w8 split=32 fast_mode=False intra_batch=True
bs=256, kv=8192 -> a8w8 split=256 fast_mode=False intra_batch=True
"""
if kv_seq_len <= 1024:
if batch_size <= 4:
return False, 32, True, False
if batch_size <= 32:
return False, 256, True, False
if batch_size <= 64:
return False, 32, False, True
return True, 32, False, True
if batch_size <= 4:
return False, 256, False, True
if batch_size <= 32:
return True, 256, False, True
if batch_size <= 64:
return True, 32, False, True
return True, 256, False, True
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
total_kv_len = batch_size * kv_seq_len
use_a8w8, num_kv_splits, fast_mode, intra_batch_mode = get_config(batch_size, kv_seq_len)
cache_key = (batch_size, kv_seq_len, use_a8w8, num_kv_splits, fast_mode, intra_batch_mode)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
if cache_key not in _cache:
q_dtype = FP8_DTYPE if use_a8w8 else torch.bfloat16
kv_dtype = FP8_DTYPE
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(
batch_size, 1, 16, q_dtype, kv_dtype,
is_sparse=False, fast_mode=fast_mode,
num_kv_splits=num_kv_splits, intra_batch_mode=intra_batch_mode,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
16, 1, True, work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=1, kv_granularity=16,
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=fast_mode, topk=-1, max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
logits = torch.empty((reduce_partial_map.size(0), 1, 16, 512), dtype=torch.float32, device="cuda")
attn_lse = torch.empty((reduce_partial_map.size(0), 1, 16, 1), dtype=torch.float32, device="cuda")
o = torch.empty((batch_size, 16, 512), dtype=torch.bfloat16, device="cuda")
entry = {
"use_a8w8": use_a8w8,
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"work_metadata": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
"logits": logits,
"attn_lse": attn_lse,
"o": o,
}
if use_a8w8:
entry["q_fp8"] = torch.empty((batch_size, 16, 576), dtype=FP8_DTYPE, device="cuda")
entry["q_scratch"] = torch.empty((batch_size, 16, 576), dtype=torch.bfloat16, device="cuda")
entry["q_amax"] = torch.empty((), dtype=q.dtype, device="cuda")
entry["q_scale"] = torch.empty((1,), dtype=torch.float32, device="cuda")
_cache[cache_key] = entry
entry = _cache[cache_key]
if entry["use_a8w8"]:
q_scratch = entry["q_scratch"]
q_amax = entry["q_amax"]
q_scale = entry["q_scale"]
q_fp8 = entry["q_fp8"]
torch.abs(q, out=q_scratch)
torch.amax(q_scratch, out=q_amax)
q_amax.clamp_(min=1e-12)
q_scale.copy_(q_amax)
q_scale.div_(FP8_MAX)
torch.div(q, q_scale, out=q_scratch)
q_scratch.clamp_(min=FP8_MIN, max=FP8_MAX)
q_fp8.copy_(q_scratch)
aiter.mla_decode_stage1_asm_fwd(
q_fp8.view(-1, 16, 576), kv_buffer_fp8.view(-1, 1, 1, 576),
qo_indptr, kv_indptr, entry["kv_indices"], entry["kv_last_page_len"], None,
entry["work_metadata"], entry["work_indptr"], entry["work_info_set"],
1, PAGE_SIZE, 1, SM_SCALE, entry["logits"], entry["attn_lse"], entry["o"], q_scale, kv_scale,
)
else:
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, 16, 576), kv_buffer_fp8.view(-1, 1, 1, 576),
qo_indptr, kv_indptr, entry["kv_indices"], entry["kv_last_page_len"], None,
entry["work_metadata"], entry["work_indptr"], entry["work_info_set"],
1, PAGE_SIZE, 1, SM_SCALE, entry["logits"], entry["attn_lse"], entry["o"], None, kv_scale,
)
aiter.mla_reduce_v1(
entry["logits"], entry["attn_lse"],
entry["reduce_indptr"], entry["reduce_final_map"], entry["reduce_partial_map"],
1, entry["o"], None,
)
return entry["o"]
scrolls · 154 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