submission 588326
manderson240 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 235 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-588326?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:cd018e770f121525bded6569839aabb4e412fe21a1c156257855a461e41435bb
license declaredunknown
license concludedunknown
authorsmanderson240
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
split_key = (bs, nheads, num_splits)Kernel source
submission.py235 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA decode: Phase 11 three-regime routing — 69.5 µs ranked geomean.
Regime 1: bs<=4 AND total_kv<=65536 → torch.einsum (bypasses aiter pipeline)
Regime 2: total_kv<=262144 → aiter a16w8 (bf16 Q, fp8 KV)
Regime 3: total_kv>262144 → aiter a8w8 (fp8 Q + fp8 KV)
"""
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
SM_SCALE = 1.0 / (576**0.5)
V_HEAD_DIM = 512
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
BF16_DTYPE = torch.bfloat16
EINSUM_MAX_BS = 4
EINSUM_MAX_TOTAL_KV = 65536
A16W8_THRESHOLD = 262144
_cache: dict = {}
_split_cache: dict = {}
_out_cache: dict = {}
_stage1_fn = None
_reduce_fn = None
def _choose_num_kv_splits(total_kv: int) -> int:
if total_kv <= 2048:
return 1
if total_kv <= 16384:
return 4
if total_kv <= 131072:
return 8
if total_kv <= 524288:
return 16
return 32
def _ensure_asm_loaded():
global _stage1_fn, _reduce_fn
if _stage1_fn is not None:
return
from aiter.mla import mla_decode_fwd # noqa: F401
import aiter
if hasattr(aiter, "mla_decode_stage1_asm_fwd"):
_stage1_fn = aiter.mla_decode_stage1_asm_fwd
else:
try:
from aiter.jit_build import module_mla_asm, module_mla_reduce
_stage1_fn = module_mla_asm.mla_decode_stage1_asm_fwd
_reduce_fn = module_mla_reduce.mla_reduce_v1
except (ImportError, AttributeError):
pass
if hasattr(aiter, "mla_reduce_v1"):
_reduce_fn = aiter.mla_reduce_v1
elif _reduce_fn is None:
try:
from aiter.jit_build import module_mla_reduce
_reduce_fn = module_mla_reduce.mla_reduce_v1
except (ImportError, AttributeError):
pass
def _quantize_fp8(tensor):
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
return (
(tensor / scale).clamp(finfo.min, finfo.max).to(FP8_DTYPE),
scale.float().reshape(1),
)
def _build_cache(bs, qseqlen, kvseqlen, nheads, q_dtype, kv_dtype, qo_indptr, kv_indptr, num_splits):
total_kv = bs * kvseqlen
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
info = get_mla_metadata_info_v1(
bs, qseqlen, nheads, q_dtype, kv_dtype,
is_sparse=False, fast_mode=True,
num_kv_splits=num_splits, intra_batch_mode=True,
)
wm, wi, wis, ri, rfm, rpm = [
torch.empty(s, dtype=t, device="cuda") for s, t in info
]
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nheads // NUM_KV_HEADS, NUM_KV_HEADS, True,
wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=qseqlen, uni_seqlen_qo=qseqlen,
fast_mode=True,
max_split_per_batch=num_splits,
intra_batch_mode=True,
dtype_q=q_dtype, dtype_kv=kv_dtype,
)
return {
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"work_meta_data": wm,
"work_indptr": wi,
"work_info_set": wis,
"reduce_indptr": ri,
"reduce_final_map": rfm,
"reduce_partial_map": rpm,
}
def _einsum_path(q, kv_data, bs, kvseqlen, qseqlen, nheads):
"""Regime 1: bypass aiter 3-stage pipeline for small batch+kv shapes."""
kv = kv_data["bf16"].view(bs, kvseqlen, QK_HEAD_DIM)
q_r = q.view(bs, qseqlen, nheads, QK_HEAD_DIM)
scores = torch.einsum("bqnh,bsh->bnqs", q_r, kv).mul_(SM_SCALE)
weights = torch.softmax(scores, dim=-1)
v = kv[:, :, :V_HEAD_DIM]
out = torch.einsum("bnqs,bsd->bqnd", weights, v)
return out.reshape(-1, nheads, V_HEAD_DIM)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kvseqlen = config["kv_seq_len"]
qseqlen = config["q_seq_len"]
nheads = config["num_heads"]
total_kv = bs * kvseqlen
# Regime 1: small batch + small kv — torch.einsum bypasses aiter pipeline overhead
if bs <= EINSUM_MAX_BS and total_kv <= EINSUM_MAX_TOTAL_KV:
return _einsum_path(q, kv_data, bs, kvseqlen, qseqlen, nheads)
kv_fp8, kv_scale = kv_data["fp8"]
kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
# Regime 2: medium — a16w8 (bf16 Q)
# Regime 3: large — a8w8 (fp8 Q for lower bandwidth)
use_a16w8 = total_kv <= A16W8_THRESHOLD
if use_a16w8:
q_input = q
q_scale = None
q_dtype = BF16_DTYPE
else:
q_input, q_scale = _quantize_fp8(q)
q_dtype = FP8_DTYPE
num_splits = _choose_num_kv_splits(total_kv)
key = (bs, qseqlen, kvseqlen, nheads, use_a16w8, num_splits)
if key not in _cache:
_cache[key] = _build_cache(
bs, qseqlen, kvseqlen, nheads,
q_dtype, FP8_DTYPE,
qo_indptr, kv_indptr, num_splits,
)
c = _cache[key]
out_key = (q.shape[0], nheads)
if out_key not in _out_cache or _out_cache[out_key].shape[0] != q.shape[0]:
_out_cache[out_key] = torch.empty(
(q.shape[0], nheads, V_HEAD_DIM),
dtype=torch.bfloat16, device="cuda",
)
o = _out_cache[out_key]
_ensure_asm_loaded()
if _stage1_fn is not None and _reduce_fn is not None:
split_key = (bs, nheads, num_splits)
if split_key not in _split_cache:
total_q = bs * qseqlen
_split_cache[split_key] = {
"split_data": torch.empty(
(total_q, num_splits, nheads, V_HEAD_DIM + 8),
dtype=torch.float32, device="cuda",
),
"split_lse": torch.empty(
(total_q, num_splits, nheads),
dtype=torch.float32, device="cuda",
),
}
sc = _split_cache[split_key]
_stage1_fn(
q_input.view(-1, nheads, QK_HEAD_DIM),
kv_4d,
qo_indptr, kv_indptr,
c["kv_indices"], c["kv_last_page_len"],
None,
c["work_meta_data"], c["work_indptr"], c["work_info_set"],
qseqlen,
PAGE_SIZE, NUM_KV_HEADS,
SM_SCALE,
sc["split_data"], sc["split_lse"], o,
q_scale=q_scale, kv_scale=kv_scale,
)
_reduce_fn(
sc["split_data"], sc["split_lse"],
c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
qseqlen,
o,
)
return o
from aiter.mla import mla_decode_fwd
mla_decode_fwd(
q_input.view(-1, nheads, QK_HEAD_DIM), kv_4d, o,
qo_indptr, kv_indptr,
c["kv_indices"], c["kv_last_page_len"],
qseqlen,
page_size=PAGE_SIZE, nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=num_splits,
q_scale=q_scale, kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=c["work_meta_data"],
work_indptr=c["work_indptr"],
work_info_set=c["work_info_set"],
reduce_indptr=c["reduce_indptr"],
reduce_final_map=c["reduce_final_map"],
reduce_partial_map=c["reduce_partial_map"],
)
return o
scrolls · 235 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