submission 662826
yuzhou2 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 163 lines, June 9 Researcher Reciprocity License v1.0.
submission_current.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-662826?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:af8989ef9efaa6c61e732023fdf2af9123a9455e339d1bcc1108a408e2daf59c
license declaredunknown
license concludedunknown
authorsyuzhou2
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
None, # num_kv_splits_indptr (not used in persistent mode)Kernel source
submission_current.py163 lines
"""MLA decode v82 — inference_mode + micro-optimizations.
Based on v77 (bypass mla_decode_fwd dispatch, pre-allocate intermediates).
Additional optimizations:
1. @torch.inference_mode() on custom_kernel — disables autograd tracking
for all tensor ops, saving ~0.5-1us overhead per call.
2. Pre-compute q_fp8 3D view in _ensure_cached() — saves 1 view() call
per invocation since views on the same storage are cheap to cache.
3. Unpack frequently-used cache values to local variables — dict lookups
are ~50ns each in CPython; locals are a single LOAD_FAST opcode.
"""
import math, torch
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
SM_SCALE = 1.0 / math.sqrt(576)
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
H = 16
NKV = 1
DK = 576
DV = 512
_cache = {}
def _ensure_cached(bs, kvl, device):
key = (bs, kvl)
if key in _cache:
return _cache[key]
total_kv = bs * kvl
qo_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device)
kv_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device) * kvl
kv_last_page_len = torch.full((bs,), kvl, dtype=torch.int32, device=device)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
info = get_mla_metadata_info_v1(
bs, 1, H, 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=device) 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, kv_last_page_len,
H // 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=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,
)
out = torch.empty((bs, H, DV), dtype=torch.bfloat16, device=device)
# Pre-allocate Q FP8 buffer
q_fp8 = torch.empty((bs * H, DK), dtype=FP8_DTYPE, device=device)
# Pre-compute 3D view (saves 1 view() call per invocation)
q_fp8_3d = q_fp8.view(bs, H, DK)
# q_scale = 1.0 (direct BF16->FP8 cast, no dynamic scaling)
q_scale = torch.ones(1, dtype=torch.float32, device=device)
# Pre-allocate intermediate buffers for stage1 ASM kernel
# These are normally allocated every call inside mla_decode_fwd()
rp_size = reduce_partial_map.size(0)
logits = torch.empty((rp_size, 1, H, DV), dtype=torch.float32, device=device)
attn_lse = torch.empty((rp_size, 1, H, 1), dtype=torch.float32, device=device)
_cache[key] = {
"qo_indptr": qo_indptr,
"kv_indptr": kv_indptr,
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"work_meta_data": 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,
"output": out,
"q_fp8": q_fp8,
"q_fp8_3d": q_fp8_3d,
"q_scale": q_scale,
"logits": logits,
"attn_lse": attn_lse,
}
return _cache[key]
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr_in, kv_indptr_in, cfg = data
bs = cfg["batch_size"]
kvl = cfg["kv_seq_len"]
c = _ensure_cached(bs, kvl, q.device)
# Unpack frequently-used cache values to locals (avoid repeated dict lookups)
q_fp8 = c["q_fp8"]
q_fp8_3d = c["q_fp8_3d"]
logits = c["logits"]
attn_lse = c["attn_lse"]
output = c["output"]
q_scale = c["q_scale"]
# Direct BF16->FP8 cast (1 HIP launch vs 3 for dynamic_per_tensor_quant)
q_2d = q.view(-1, DK)
q_fp8.copy_(q_2d)
# FP8 KV
kv_fp8, kv_scale = kv_data["fp8"]
tkv = bs * kvl
kv_4d = kv_fp8[:tkv].view(tkv, PAGE_SIZE, NKV, DK)
# Direct ASM dispatch — bypass mla_decode_fwd() Python wrapper
# Saves: 2 tensor allocations (logits, attn_lse) + Python dispatch overhead
aiter.mla_decode_stage1_asm_fwd(
q_fp8_3d,
kv_4d,
c["qo_indptr"],
c["kv_indptr"],
c["kv_indices"],
c["kv_last_page_len"],
None, # num_kv_splits_indptr (not used in persistent mode)
c["work_meta_data"],
c["work_indptr"],
c["work_info_set"],
1, # max_seqlen_q
PAGE_SIZE,
NKV,
SM_SCALE,
logits,
attn_lse,
output,
q_scale,
kv_scale,
)
aiter.mla_reduce_v1(
logits,
attn_lse,
c["reduce_indptr"],
c["reduce_final_map"],
c["reduce_partial_map"],
1, # max_seqlen_q
output,
None, # final_lse
)
return output
scrolls · 163 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