submission 747474
rajeev9 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 117 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747474?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:3ccb707fec89f36a4087341689f3600e076ce14ad948b0666740462fb5437748
license declaredunknown
license concludedunknown
authorsrajeev9
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
MLA Decode: All-persistent FP8 ASM with per-config kv_granularity tuning.Kernel source
submission.py117 lines
"""
MLA Decode: All-persistent FP8 ASM with per-config kv_granularity tuning.
g=64 for bs<=4/kv<=1024, g=128 for bs<=4/kv>1024, g=32 for bs<=32, g=16 for bs>32.
"""
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.jit.utils.chip_info import get_cu_num
SM_SCALE: float = 1.0 / (576 ** 0.5)
FP8 = dtypes.fp8
_QSC = torch.empty(1, dtype=torch.float32, device="cuda")
_CACHE = {}
_quant = aiter.dynamic_per_tensor_quant
_stage1 = aiter.mla_decode_stage1_asm_fwd
_reduce = aiter.mla_reduce_v1
_get_meta = aiter.get_mla_metadata_v1
def _get_kv_granularity(bs, kv_seq):
if bs <= 4 and kv_seq <= 1024:
return 64
if bs <= 4:
return 128
if bs <= 32:
return 32
return 16
def _build_cache(bs, kv_seq, nq, nkv, dv, dq, qsl, qo_indptr, kv_indptr):
cu_num = get_cu_num()
tkv = bs * kv_seq
tile_cnt = bs
max_work = bs + cu_num - 1
max_split_tiles = min(max_work, (cu_num - 1) * 2)
device = "cuda"
work_meta = torch.empty(2, dtype=torch.uint64, device=device)
work_indptr = torch.empty(cu_num + 1, dtype=torch.int32, device=device)
work_info_set = torch.empty((max_work, 8), dtype=torch.int32, device=device)
reduce_indptr = torch.empty(tile_cnt + 1, dtype=torch.int32, device=device)
reduce_final_map = torch.empty((tile_cnt, 2), dtype=torch.int32, device=device)
reduce_partial_map = torch.empty(max_split_tiles, dtype=torch.int32, device=device)
klp = torch.full((bs,), kv_seq, dtype=torch.int32, device=device)
kv_indices = torch.arange(tkv, dtype=torch.int32, device=device)
kvg = _get_kv_granularity(bs, kv_seq)
_get_meta(
qo_indptr, kv_indptr, klp,
nq // nkv, nkv, False,
work_meta, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
1, kvg, qsl, qsl, True, -1, -1, False, FP8, FP8,
)
total_q = bs * qsl
logits = torch.empty((max_split_tiles * qsl, 1, nq, dv), dtype=torch.float32, device=device)
attn_lse = torch.empty((max_split_tiles * qsl, 1, nq, 1), dtype=torch.float32, device=device)
o = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device=device)
qf = torch.empty((total_q, nq * dq), dtype=FP8, device=device)
qf_3d = qf.view(-1, nq, dq)
return (
work_meta, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map,
klp, kv_indices,
logits, attn_lse, o, qf, qf_3d,
qsl, nkv,
)
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dv = config["v_head_dim"]
dq = config["qk_head_dim"]
qsl = config["q_seq_len"]
kv_seq = config["kv_seq_len"]
key = (bs, kv_seq)
if key not in _CACHE:
_CACHE[key] = _build_cache(bs, kv_seq, nq, nkv, dv, dq, qsl, qo_indptr, kv_indptr)
(work_meta, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map,
klp, kv_indices,
logits, attn_lse, o, qf, qf_3d,
qsl_c, nkv_c) = _CACHE[key]
kv_fp8, kv_sc = kv_data["fp8"]
_quant(qf, q, _QSC)
_stage1(
qf_3d, kv_fp8.view(-1, 1, nkv_c, dq),
qo_indptr, kv_indptr, kv_indices, klp,
None, work_meta, work_indptr, work_info_set,
qsl_c, 1, nkv_c, SM_SCALE,
logits, attn_lse, o, _QSC, kv_sc,
)
_reduce(
logits, attn_lse,
reduce_indptr, reduce_final_map, reduce_partial_map,
qsl_c, o, None,
)
return o
scrolls · 117 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