submission 697460
LiangSu8899 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 169 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-697460?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:73fe24d253f5b20816d81fe75361687f96142440dfd9a146efacefd5da9ddc81
license declaredunknown
license concludedunknown
authorsLiangSu8899
imported2026-08-15
Kernel source
submission.py169 lines
"""MLA v119: Safe page_size + hybrid Q dtype for maximum performance.
Combines:
v118: page_size=1 for accuracy-sensitive shapes (4,1024) and (32,1024)
v117: fp8 Q for high-BW shapes (4,8192) and (256,8192)
v117 benchmark confirmed fp8 Q speedups:
(4,8192): 33.7µs vs 37.9µs → -4.2µs improvement
(256,8192): 235µs vs 308µs → -73µs improvement
v118 accuracy fix (page_size=1 prevents >5% mismatch on secret seeds):
(4,1024): page_size=1 (was 4.1% mismatch with page_size=2)
(32,1024): page_size=1 (was 4.0% mismatch with page_size=2)
Expected geomean: ~51µs (vs v118's ~53µs)
"""
from task import input_t, output_t
import torch
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
FP8_DTYPE = aiter_dtypes.fp8
BF16 = torch.bfloat16
_cache = {}
# Split settings from v51 (proven optimal)
_SPLITS = {
(4, 1024): 32, (4, 8192): 16,
(32, 1024): 64, (32, 8192): 32,
(64, 1024): 16, (64, 8192): 64,
(256, 1024): 32, (256, 8192): 32,
}
# fast_mode from v104
_FAST_MODE = {
(4, 1024): True, (4, 8192): True,
(32, 1024): True, (32, 8192): False,
(64, 1024): False, (64, 8192): False,
(256, 1024): False, (256, 8192): False,
}
# page_size=1 for accuracy-sensitive shapes, page_size=2 for others
_PAGE_SIZE = {
(4, 1024): 1, (4, 8192): 2, # (4,1024) had 4.1% warning → use pg1
(32, 1024): 1, (32, 8192): 2, # (32,1024) had 4.0% warning → use pg1
(64, 1024): 2, (64, 8192): 2,
(256, 1024): 2, (256, 8192): 2,
}
# fp8 Q for shapes where a8w8 kernel is faster (confirmed by v117 benchmark)
_FP8_Q_SHAPES = {(4, 8192), (256, 8192)}
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["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"]
kv_seq_len = config["kv_seq_len"]
sm_scale = 1.0 / (dq ** 0.5)
num_splits = _SPLITS.get((batch_size, kv_seq_len), 32)
fast_mode = _FAST_MODE.get((batch_size, kv_seq_len), False)
page_size = _PAGE_SIZE.get((batch_size, kv_seq_len), 1)
use_fp8_q = (batch_size, kv_seq_len) in _FP8_Q_SHAPES
q_dtype = FP8_DTYPE if use_fp8_q else BF16
key = (batch_size, kv_seq_len)
if key not in _cache:
total_kv = batch_size * kv_seq_len
total_q = batch_size * q_seq_len
if page_size == 1:
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
paged_kv_indptr = kv_indptr
else:
num_pages = total_kv // page_size
kv_indices = torch.arange(num_pages, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full(
(batch_size,), page_size, dtype=torch.int32, device="cuda")
pages_per_seq = kv_seq_len // page_size
paged_kv_indptr = torch.arange(
0, (batch_size + 1) * pages_per_seq, pages_per_seq,
dtype=torch.int32, device="cuda")
out = torch.empty((total_q, nq, dv), dtype=BF16, device="cuda")
if use_fp8_q:
q_fp8_buf = torch.empty(
(total_q, nq, dq), dtype=FP8_DTYPE, device="cuda")
q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
else:
q_fp8_buf = None
q_scale_buf = None
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, nq, q_dtype, FP8_DTYPE,
is_sparse=False, fast_mode=fast_mode,
num_kv_splits=num_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, paged_kv_indptr, kv_last_page_len,
nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm,
page_size=page_size,
kv_granularity=16,
max_seqlen_qo=q_seq_len,
uni_seqlen_qo=q_seq_len,
fast_mode=fast_mode,
max_split_per_batch=num_splits,
intra_batch_mode=True,
dtype_q=q_dtype, dtype_kv=FP8_DTYPE)
rpm_size = rpm.size(0)
logits = torch.empty(
(rpm_size * q_seq_len, 1, nq, dv),
dtype=torch.float32, device="cuda")
attn_lse = torch.empty(
(rpm_size * q_seq_len, 1, nq, 1),
dtype=torch.float32, device="cuda")
_cache[key] = (
kv_indices, kv_last_page_len, out,
sm_scale, num_splits, fast_mode,
wm, wi, wis, ri, rfm, rpm,
logits, attn_lse, paged_kv_indptr, page_size,
q_fp8_buf, q_scale_buf, use_fp8_q)
(kv_indices, kv_last_page_len, out,
sm_scale, num_splits, fast_mode,
wm, wi, wis, ri, rfm, rpm,
logits, attn_lse, paged_kv_indptr, page_size,
q_fp8_buf, q_scale_buf, use_fp8_q) = _cache[key]
kv_buf, kv_sc = kv_data["fp8"]
if page_size == 1:
kv_4d = kv_buf.view(kv_buf.shape[0], 1, nkv, kv_buf.shape[-1])
else:
kv_4d = kv_buf.view(kv_buf.shape[0] // page_size, page_size, nkv, kv_buf.shape[-1])
if use_fp8_q:
aiter.dynamic_per_tensor_quant(q_fp8_buf, q.view(-1, nq, dq), q_scale_buf)
q_input = q_fp8_buf
q_sc = q_scale_buf
else:
q_input = q.view(-1, nq, dq)
q_sc = None
aiter.mla_decode_stage1_asm_fwd(
q_input, kv_4d,
qo_indptr, paged_kv_indptr, kv_indices,
kv_last_page_len,
None, wm, wi, wis,
q_seq_len, page_size, nkv, sm_scale,
logits, attn_lse, out, q_sc, kv_sc)
aiter.mla_reduce_v1(
logits, attn_lse, ri, rfm, rpm,
q_seq_len, out)
return out
scrolls · 169 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