submission 687741
Jianian-Xu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 245 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-687741?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:da6f4eaa0eea0df12725635c61d219193fadc53989dfe978fe45ec38df130f16
license declaredunknown
license concludedunknown
authorsJianian-Xu
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
For kv<=1024: non-persistent mode with num_kv_splits=1, skips reduce entirely (2 launches).split-k
_SINGLE_SPLIT_KV_THRESHOLD = 0 # 0 = never use single-splitKernel source
submission.py245 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
"""
MLA decode — bypasses mla_decode_fwd, calls stage1+reduce directly.
For kv<=1024: non-persistent mode with num_kv_splits=1, skips reduce entirely (2 launches).
For kv>1024: persistent mode with splits + reduce (3 launches).
"""
import torch
from task import input_t, output_t
from utils import make_match_reference
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter import dynamic_per_tensor_quant as _dpq
from aiter import static_per_tensor_quant as _spq
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576
V_HEAD_DIM = KV_LORA_RANK # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 16
FP8_DTYPE = aiter_dtypes.fp8
# Threshold: kv_per_seq <= this uses non-persistent single-split (skip reduce)
# DISABLED: non-persistent with 1 split has worse CU utilization than persistent
_SINGLE_SPLIT_KV_THRESHOLD = 0 # 0 = never use single-split
# ---------------------------------------------------------------------------
# Caches
# ---------------------------------------------------------------------------
_metadata_cache: dict = {}
_kv_indices_cache: dict = {}
_kv_last_page_len_cache: dict = {}
_kv_indptr_cache: dict = {}
_qo_indptr_cache: dict = {}
_splits_indptr_cache: dict = {}
_fast_path_cache: dict = {}
def _get_cached_kv_indices(n):
if n not in _kv_indices_cache:
_kv_indices_cache[n] = torch.arange(n, dtype=torch.int32, device="cuda")
return _kv_indices_cache[n]
def _get_cached_kv_last_page_len(total_kv, bs):
k = (total_kv, bs)
if k not in _kv_last_page_len_cache:
_kv_last_page_len_cache[k] = torch.full((bs,), total_kv // bs, dtype=torch.int32, device="cuda")
return _kv_last_page_len_cache[k]
def _get_cached_indptr(cache, bs, step):
k = (bs, step)
if k not in cache:
cache[k] = torch.arange(0, (bs + 1) * step, step, dtype=torch.int32, device="cuda")
return cache[k]
def _get_cached_metadata(bs, max_q_len, nq, nkv, q_dtype, kv_dtype,
qo_indptr, kv_indptr, kv_last_page_len,
total_kv_len, num_kv_splits):
ck = (bs, max_q_len, nq, nkv, q_dtype, kv_dtype, num_kv_splits, total_kv_len)
if ck not in _metadata_cache:
info = get_mla_metadata_info_v1(
bs, max_q_len, nq, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=num_kv_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
d = {
"work_meta_data": work[0], "work_indptr": work[1], "work_info_set": work[2],
"reduce_indptr": work[3], "reduce_final_map": work[4], "reduce_partial_map": work[5],
}
_metadata_cache[ck] = (d, [None])
cached, pop_store = _metadata_cache[ck]
pk = (qo_indptr.data_ptr(), kv_indptr.data_ptr(), kv_last_page_len.data_ptr())
if pop_store[0] != pk:
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nq // nkv, nkv, True,
cached["work_meta_data"], cached["work_info_set"], cached["work_indptr"],
cached["reduce_indptr"], cached["reduce_final_map"], cached["reduce_partial_map"],
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=max_q_len, uni_seqlen_qo=max_q_len,
fast_mode=False, max_split_per_batch=num_kv_splits,
intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
pop_store[0] = pk
return cached
# ---------------------------------------------------------------------------
# Main kernel
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_buffer_fp8, kv_scale = kv_data["fp8"]
kv_input = kv_buffer_fp8
fp_key = (config["batch_size"], kv_input.shape[0])
fp = _fast_path_cache.get(fp_key)
if fp is not None:
# --- Fast path: always run quant (Q/KV data change in leaderboard) ---
_spq(fp["q_fp8"], q, fp["q_scale"])
kv_4d = kv_input.view(fp["kv_view"])
if fp["single_split"]:
aiter.mla_decode_stage1_asm_fwd(
fp["q_fp8_3d"], kv_4d,
fp["qo_indptr"], fp["kv_indptr"], fp["kv_indices"], fp["kv_last_page_len"],
fp["splits_indptr"], None, None, None,
1, PAGE_SIZE, fp["nkv"], SM_SCALE,
fp["logits_view"], fp["attn_lse"], fp["o"],
q_scale=fp["q_scale"], kv_scale=kv_scale,
)
else:
aiter.mla_decode_stage1_asm_fwd(
fp["q_fp8_3d"], kv_4d,
fp["qo_indptr"], fp["kv_indptr"], fp["kv_indices"], fp["kv_last_page_len"],
None, fp["work_meta_data"], fp["work_indptr"], fp["work_info_set"],
1, PAGE_SIZE, fp["nkv"], SM_SCALE,
fp["logits"], fp["attn_lse"], fp["o"],
q_scale=fp["q_scale"], kv_scale=kv_scale,
)
aiter.mla_reduce_v1(
fp["logits"], fp["attn_lse"],
fp["reduce_indptr"], fp["reduce_final_map"], fp["reduce_partial_map"],
1, fp["o"], None,
)
return fp["o"]
# --- Slow path: first call per config ---
batch_size, total_kv_len = fp_key
kv_per_seq = total_kv_len // batch_size
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
q_seq_len = config["q_seq_len"]
total_q = q.shape[0]
# Decide: single-split non-persistent or multi-split persistent
single_split = (kv_per_seq <= _SINGLE_SPLIT_KV_THRESHOLD)
# Common buffers
kv_indptr_c = _get_cached_indptr(_kv_indptr_cache, batch_size, kv_per_seq)
qo_indptr_c = _get_cached_indptr(_qo_indptr_cache, batch_size, q_seq_len)
kv_indices = _get_cached_kv_indices(total_kv_len)
kv_last_page_len = _get_cached_kv_last_page_len(total_kv_len, batch_size)
q_fp8_buf = torch.empty(q.shape, dtype=FP8_DTYPE, device="cuda")
q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
o_buf = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
q_fp8_3d = q_fp8_buf.view(-1, nq, dq)
kv_4d = kv_input.view(total_kv_len, PAGE_SIZE, nkv, kv_input.shape[-1])
_dpq(q_fp8_buf, q, q_scale_buf)
fp_entry = {
"q_fp8": q_fp8_buf, "q_fp8_3d": q_fp8_3d, "q_scale": q_scale_buf,
"kv_view": (total_kv_len, PAGE_SIZE, nkv, kv_input.shape[-1]),
"kv_4d": kv_4d,
"o": o_buf,
"qo_indptr": qo_indptr_c, "kv_indptr": kv_indptr_c,
"kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
"nkv": nkv, "single_split": single_split,
"last_q_ptr": None,
}
if single_split:
# Non-persistent, num_kv_splits=1: stage1 writes directly to o
splits_indptr = _get_cached_indptr(_splits_indptr_cache, batch_size, 1)
# logits = o.view(total_q, 1, nq, dv) — view of output, no alloc
logits_view = o_buf.view(total_q, 1, nq, dv)
attn_lse = torch.empty((total_q, 1, nq, 1), dtype=torch.float32, device="cuda")
aiter.mla_decode_stage1_asm_fwd(
q_fp8_3d, kv_4d,
qo_indptr_c, kv_indptr_c, kv_indices, kv_last_page_len,
splits_indptr, None, None, None, # non-persistent
q_seq_len, PAGE_SIZE, nkv, SM_SCALE,
logits_view, attn_lse, o_buf,
q_scale=q_scale_buf, kv_scale=kv_scale,
)
# NO reduce needed
fp_entry.update({
"splits_indptr": splits_indptr,
"logits_view": logits_view,
"attn_lse": attn_lse,
})
else:
# Persistent mode
num_splits = 64 if batch_size >= 256 else NUM_KV_SPLITS
meta = _get_cached_metadata(
batch_size, q_seq_len, nq, nkv,
FP8_DTYPE, kv_input.dtype,
qo_indptr_c, kv_indptr_c, kv_last_page_len,
total_kv_len, num_splits,
)
rpm_size = meta["reduce_partial_map"].size(0)
logits_buf = torch.empty((rpm_size * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda")
attn_lse_buf = torch.empty((rpm_size * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda")
aiter.mla_decode_stage1_asm_fwd(
q_fp8_3d, kv_4d,
qo_indptr_c, kv_indptr_c, kv_indices, kv_last_page_len,
None, meta["work_meta_data"], meta["work_indptr"], meta["work_info_set"],
q_seq_len, PAGE_SIZE, nkv, SM_SCALE,
logits_buf, attn_lse_buf, o_buf,
q_scale=q_scale_buf, kv_scale=kv_scale,
)
aiter.mla_reduce_v1(
logits_buf, attn_lse_buf,
meta["reduce_indptr"], meta["reduce_final_map"], meta["reduce_partial_map"],
q_seq_len, o_buf, None,
)
fp_entry.update({
"logits": logits_buf, "attn_lse": attn_lse_buf,
"work_meta_data": meta["work_meta_data"],
"work_indptr": meta["work_indptr"],
"work_info_set": meta["work_info_set"],
"reduce_indptr": meta["reduce_indptr"],
"reduce_final_map": meta["reduce_final_map"],
"reduce_partial_map": meta["reduce_partial_map"],
})
_fast_path_cache[fp_key] = fp_entry
return o_buf
scrolls · 245 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