submission 746532
kkosey · 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_e99_noibm_ps8.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-746532?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:809672bf4e72c30ec3266c2092da1cda3e2c08b18bf911e6b2823b3df5408744
license declaredunknown
license concludedunknown
authorskkosey
imported2026-08-15
Kernel source
submission_e99_noibm_ps8.py169 lines
#!POPCORN gpu MI355X
#!POPCORN leaderboard amd-mixed-mla
"""E93: E88 optimal dispatch + hybrid PAGE_SIZE (PS=4 for kv>1024, PS=2 for kv≤1024).
Combines:
- E88: a16w8 for all except bs≥256 kv>1024 (a8w8), IBM=True for kv>1024 bs≥32
- E89: PAGE_SIZE=4 gives -16% to -32% speedup for kv=8192 shapes
Expected: best of both worlds.
"""
import torch
import aiter
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
FP8_DTYPE = aiter_dtypes.fp8
NUM_KV_HEADS = 1
SM_SCALE = 1.0 / (576 ** 0.5)
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
_cache = {}
_stage1_op = torch.ops.aiter.mla_decode_stage1_asm_fwd
_reduce_op = torch.ops.aiter.mla_reduce_v1
_dpt_quant_op = torch.ops.aiter.dynamic_per_tensor_quant
_MIN_PROGRAMS = 120
_TARGET_GRID = 512
def _choose_strategy(batch_size, q_seq_len, kv_seq_len):
total_q = batch_size * q_seq_len
# bs≤4: always a16w8 NS=1
if batch_size <= 4:
return True, 1
# E98: ALL a16w8 (no a8w8 branch)
ns = max(1, (_TARGET_GRID + total_q - 1) // total_q)
return True, min(16, ns)
def _init_shape(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr):
nkv = NUM_KV_HEADS
use_a16w8, num_kv_splits = _choose_strategy(batch_size, q_seq_len, kv_seq_len)
q_dtype = torch.bfloat16 if use_a16w8 else FP8_DTYPE
kv_dtype = FP8_DTYPE
# Hybrid PAGE_SIZE: PS=4 for kv>1024, PS=2 for kv≤1024
page_size = 8 if kv_seq_len > 1024 else 2
# IBM: kv>1024 AND bs>=32 (E88 optimal)
ibm = False
kv_indptr_pages = (kv_indptr // page_size).to(torch.int32)
kv_last_page_len = (kv_indptr_pages[1:] - kv_indptr_pages[:-1]).to(torch.int32)
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, num_heads, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=num_kv_splits, intra_batch_mode=ibm,
)
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_pages, kv_last_page_len,
num_heads // nkv, nkv, True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=page_size,
kv_granularity=max(page_size, 16),
max_seqlen_qo=q_seq_len,
uni_seqlen_qo=q_seq_len,
fast_mode=False,
max_split_per_batch=num_kv_splits,
intra_batch_mode=ibm,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
total_q = batch_size * q_seq_len
total_kv = batch_size * kv_seq_len
total_pages = total_kv // page_size
partial_size = reduce_partial_map.size(0) * q_seq_len
result = {
"use_a16w8": use_a16w8,
"page_size": page_size,
"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,
"kv_last_page_len": kv_last_page_len,
"kv_indptr_pages": kv_indptr_pages,
"kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
"logits": torch.empty((partial_size, 1, num_heads, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((partial_size, 1, num_heads, 1), dtype=torch.float32, device="cuda"),
"output": torch.empty((total_q, num_heads, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
}
if not use_a16w8:
result["q_fp8"] = torch.empty((total_q, num_heads, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
result["q_scale"] = torch.empty(1, dtype=torch.float32, device="cuda")
return result
def _get_cache(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr):
key = (batch_size, q_seq_len, num_heads, kv_seq_len)
if key not in _cache:
_cache[key] = _init_shape(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr)
return _cache[key]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
num_heads = config["num_heads"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
c = _get_cache(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr)
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
total_kv = batch_size * kv_seq_len
page_size = c["page_size"]
kv_buffer_4d = kv_buffer_fp8.view(
total_kv // page_size, page_size, NUM_KV_HEADS, kv_buffer_fp8.shape[-1]
)
o = c["output"]
kv_indptr_pages = c["kv_indptr_pages"]
if c["use_a16w8"]:
_stage1_op(
q, kv_buffer_4d, qo_indptr, kv_indptr_pages,
c["kv_indices"], c["kv_last_page_len"], None,
c["work_metadata"], c["work_indptr"], c["work_info_set"],
q_seq_len, page_size, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], o,
None, kv_scale_fp8,
)
else:
q_fp8 = c["q_fp8"]
q_scale = c["q_scale"]
_dpt_quant_op(q_fp8, q, q_scale)
_stage1_op(
q_fp8, kv_buffer_4d, qo_indptr, kv_indptr_pages,
c["kv_indices"], c["kv_last_page_len"], None,
c["work_metadata"], c["work_indptr"], c["work_info_set"],
q_seq_len, page_size, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], o,
q_scale, kv_scale_fp8,
)
_reduce_op(
c["logits"], c["attn_lse"],
c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
q_seq_len, o, None,
)
return o
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