submission 632302
inference_and_chill · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 181 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-632302?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:4dedec7b57c2a888f973cfa03cf904874e2639301cb2545527c1e7ba589a64c9
license declaredunknown
license concludedunknown
authorsinference_and_chill
imported2026-08-15
Kernel source
submission.py181 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA decode: FP8 aiter wrapper with static Q quant + metadata reuse."""
from __future__ import annotations
import torch
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, static_per_tensor_quant
from task import input_t, output_t
# ---------------------------------------------------------------------------
# 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
FP8_DTYPE = aiter_dtypes.fp8
# Per-case tuning: (batch_size, kv_seq_len) -> (num_kv_splits, kv_granularity)
_TUNE = {
(4, 1024): (16, 16),
(4, 8192): (16, 16),
(32, 1024): (8, 16),
(32, 8192): (24, 16),
(64, 1024): (8, 16),
(64, 8192): (24, 16),
(256, 1024): (8, 16),
(256, 8192): (24, 16),
}
# ---------------------------------------------------------------------------
# Caches — keyed to avoid per-call allocations
# ---------------------------------------------------------------------------
_case_cache: dict[tuple, dict] = {}
_q_ref: torch.Tensor | None = None
_q_fp8_buf: torch.Tensor | None = None
_q_scale_buf: torch.Tensor | None = None
_use_static_quant: bool = False
_kv_ref: torch.Tensor | None = None
_kv_4d: torch.Tensor | None = None
# ---------------------------------------------------------------------------
# Entry point — FP8 ASM path
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
global _q_ref, _q_fp8_buf, _q_scale_buf, _use_static_quant
global _kv_ref, _kv_4d
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
# Cached Q FP8 quantization
if q is not _q_ref:
if _q_fp8_buf is None or _q_fp8_buf.shape != q.shape:
_q_fp8_buf = torch.empty(q.shape, dtype=FP8_DTYPE, device="cuda")
_q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
_use_static_quant = False
if _use_static_quant:
static_per_tensor_quant(_q_fp8_buf, q, _q_scale_buf)
else:
dynamic_per_tensor_quant(_q_fp8_buf, q, _q_scale_buf)
_use_static_quant = True
_q_ref = q
# FP8 KV — cache view
kv_buffer_fp8, kv_scale = kv_data["fp8"]
if kv_buffer_fp8 is not _kv_ref:
_kv_4d = kv_buffer_fp8.view(
kv_buffer_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_buffer_fp8.shape[-1]
)
_kv_ref = kv_buffer_fp8
# Per-case cached data (metadata, indices, output, kv_last_page_len)
case_key = (batch_size, kv_seq_len)
cd = _case_cache.get(case_key)
if cd is None:
tune = _TUNE.get(case_key, (32, 16))
num_kv_splits, kv_gran = tune
total_kv_len = batch_size * kv_seq_len
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
info = get_mla_metadata_info_v1(
batch_size,
1,
NUM_HEADS,
FP8_DTYPE,
FP8_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]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(
qo_indptr,
kv_indptr,
kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
wm,
wis,
wi,
ri,
rfm,
rpm,
page_size=PAGE_SIZE,
kv_granularity=kv_gran,
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=False,
max_split_per_batch=num_kv_splits,
intra_batch_mode=True,
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
total_q = batch_size
o = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
cd = {
"nks": num_kv_splits,
"ki": kv_indices,
"klpl": kv_last_page_len,
"wm": wm,
"wi": wi,
"wis": wis,
"ri": ri,
"rfm": rfm,
"rpm": rpm,
"o": o,
}
_case_cache[case_key] = cd
o = cd["o"]
mla_decode_fwd(
_q_fp8_buf,
_kv_4d,
o,
qo_indptr,
kv_indptr,
cd["ki"],
cd["klpl"],
1, # max_seqlen_q
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=cd["nks"],
q_scale=_q_scale_buf,
kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=cd["wm"],
work_indptr=cd["wi"],
work_info_set=cd["wis"],
reduce_indptr=cd["ri"],
reduce_final_map=cd["rfm"],
reduce_partial_map=cd["rpm"],
)
return o
scrolls · 181 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