submission 600564
johnny.t.shi · 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.
v12_fp8_q_kv.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-600564?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:bebc30e0b45c9f207144aa8fcaa809c217a4c8aacc075910ffeb5fe06001902f
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
"""MLA Decode v12 — FP8 Q + FP8 KV with persistent mode.Kernel source
v12_fp8_q_kv.py117 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA Decode v12 — FP8 Q + FP8 KV with persistent mode.
Key insight from recon: competitors at ranks 6,9 use 'a16w8_ps2':
- Cast Q from BF16 → FP8 (a16 = original precision, w8 = KV in FP8)
- Use persistent mode (ps=1) with pre-computed metadata
- aiter only supports fp8+fp8 in persistent mode
Steps:
1. Cast Q to float8_e4m3fn
2. Pre-compute metadata via get_mla_metadata_info_v1 + get_mla_metadata_v1
3. Call mla_decode_fwd in persistent mode
"""
from task import input_t, output_t
import torch
import aiter
from aiter import dtypes
from aiter.mla import mla_decode_fwd
_meta_cache = {}
def _get_or_compute_metadata(batch_size, total_kv, num_heads, nhead_kv, page_size,
qo_indptr, kv_indptr, kv_last_page_lens, device):
"""Cache metadata by shape key. Safe: metadata depends on shape, not data content."""
key = (batch_size, total_kv, num_heads, page_size)
cached = _meta_cache.get(key)
if cached is not None:
return cached
(
(wmd_sz, wmd_ty), (wi_sz, wi_ty), (wis_sz, wis_ty),
(ri_sz, ri_ty), (rfm_sz, rfm_ty), (rpm_sz, rpm_ty),
) = aiter.get_mla_metadata_info_v1(
batch_size, 1, num_heads, dtypes.fp8, dtypes.fp8,
is_sparse=False, fast_mode=True, num_kv_splits=-1, intra_batch_mode=False,
)
wmd = torch.empty(wmd_sz, dtype=wmd_ty, device=device)
wi = torch.empty(wi_sz, dtype=wi_ty, device=device)
wis = torch.empty(wis_sz, dtype=wis_ty, device=device)
ri = torch.empty(ri_sz, dtype=ri_ty, device=device)
rfm = torch.empty(rfm_sz, dtype=rfm_ty, device=device)
rpm = torch.empty(rpm_sz, dtype=rpm_ty, device=device)
aiter.get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_lens,
num_heads // nhead_kv, nhead_kv, False,
wmd, wis, wi, ri, rfm, rpm,
page_size=page_size, kv_granularity=max(page_size, 16),
max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,
max_split_per_batch=-1, intra_batch_mode=False,
dtype_q=dtypes.fp8, dtype_kv=dtypes.fp8,
)
result = (wmd, wi, wis, ri, rfm, rpm)
_meta_cache[key] = result
return result
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
batch_size = config['batch_size']
num_heads = config['num_heads']
v_head_dim = config['v_head_dim']
sm_scale = config['sm_scale']
total_kv = kv_fp8_data.shape[0]
# Cast Q to FP8
q_fp8 = q.to(torch.float8_e4m3fn)
q_scale = torch.ones([1], dtype=torch.float32, device=q.device)
output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)
# page_size=2
PAGE_SIZE = 2
num_pages = total_kv // PAGE_SIZE
kv_buffer = kv_fp8_data.view(num_pages, PAGE_SIZE, 1, 576)
kv_indices = torch.arange(num_pages, device=q.device, dtype=torch.int32)
kv_indptr_pages = kv_indptr // PAGE_SIZE
kv_last_page_lens = torch.full((batch_size,), PAGE_SIZE, device=q.device, dtype=torch.int32)
# Cached metadata: depends on shape only, not data content. Safe for leaderboard.
(work_meta_data, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = _get_or_compute_metadata(
batch_size, total_kv, num_heads, 1, PAGE_SIZE,
qo_indptr, kv_indptr_pages, kv_last_page_lens, q.device,
)
mla_decode_fwd(
q=q_fp8,
kv_buffer=kv_buffer,
o=output,
qo_indptr=qo_indptr,
kv_indptr=kv_indptr_pages,
kv_indices=kv_indices,
kv_last_page_lens=kv_last_page_lens,
max_seqlen_q=1,
page_size=PAGE_SIZE,
nhead_kv=1,
sm_scale=sm_scale,
work_meta_data=work_meta_data,
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,
q_scale=q_scale,
kv_scale=kv_fp8_scale,
)
return output
scrolls · 117 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 598454.
⋯ 19 unchanged lines_meta_cache = {}-def _get_or_compute_metadata(batch_size, total_kv, num_heads, nhead_kv, page_size,qo_indptr, kv_indptr, kv_last_page_lens, device):- """Compute and cache persistent mode metadata."""+ """Cache metadata by shape key. Safe: metadata depends on shape, not data content."""key = (batch_size, total_kv, num_heads, page_size)cached = _meta_cache.get(key)if cached is not None:return cached- max_seqlen_qo = 1 # decode- max_split_per_batch = -1 # auto-- # Get metadata tensor sizes(- (work_meta_data_size, work_meta_data_type),- (work_indptr_size, work_indptr_type),- (work_info_set_size, work_info_set_type),- (reduce_indptr_size, reduce_indptr_type),- (reduce_final_map_size, reduce_final_map_type),- (reduce_partial_map_size, reduce_partial_map_type),+ (wmd_sz, wmd_ty), (wi_sz, wi_ty), (wis_sz, wis_ty),+ (ri_sz, ri_ty), (rfm_sz, rfm_ty), (rpm_sz, rpm_ty),) = aiter.get_mla_metadata_info_v1(- batch_size,- max_seqlen_qo,- num_heads,- dtypes.fp8, # q dtype- dtypes.fp8, # kv dtype- is_sparse=False,- fast_mode=True,- num_kv_splits=max_split_per_batch,- intra_batch_mode=False,+ batch_size, 1, num_heads, dtypes.fp8, dtypes.fp8,+ is_sparse=False, fast_mode=True, num_kv_splits=-1, intra_batch_mode=False,)- # Pre-allocate metadata tensors- work_meta_data = torch.empty(work_meta_data_size, dtype=work_meta_data_type, device=device)- work_indptr = torch.empty(work_indptr_size, dtype=work_indptr_type, device=device)- work_info_set = torch.empty(work_info_set_size, dtype=work_info_set_type, device=device)- reduce_indptr = torch.empty(reduce_indptr_size, dtype=reduce_indptr_type, device=device)- reduce_final_map = torch.empty(reduce_final_map_size, dtype=reduce_final_map_type, device=device)- reduce_partial_map = torch.empty(reduce_partial_map_size, dtype=reduce_partial_map_type, device=device)+ wmd = torch.empty(wmd_sz, dtype=wmd_ty, device=device)+ wi = torch.empty(wi_sz, dtype=wi_ty, device=device)+ wis = torch.empty(wis_sz, dtype=wis_ty, device=device)+ ri = torch.empty(ri_sz, dtype=ri_ty, device=device)+ rfm = torch.empty(rfm_sz, dtype=rfm_ty, device=device)+ rpm = torch.empty(rpm_sz, dtype=rpm_ty, device=device)- # Populate metadataaiter.get_mla_metadata_v1(- qo_indptr,- kv_indptr,- kv_last_page_lens,- num_heads // nhead_kv, # num_heads_per_head_k- nhead_kv, # num_heads_k- False, # is_causal- work_meta_data,- 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=max_seqlen_qo,- uni_seqlen_qo=1, # decode- fast_mode=True,- max_split_per_batch=max_split_per_batch,- intra_batch_mode=False,- dtype_q=dtypes.fp8,- dtype_kv=dtypes.fp8,+ qo_indptr, kv_indptr, kv_last_page_lens,+ num_heads // nhead_kv, nhead_kv, False,+ wmd, wis, wi, ri, rfm, rpm,+ page_size=page_size, kv_granularity=max(page_size, 16),+ max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,+ max_split_per_batch=-1, intra_batch_mode=False,+ dtype_q=dtypes.fp8, dtype_kv=dtypes.fp8,)- result = (work_meta_data, work_indptr, work_info_set,- reduce_indptr, reduce_final_map, reduce_partial_map)+ result = (wmd, wi, wis, ri, rfm, rpm)_meta_cache[key] = resultreturn result⋯ 24 unchanged lineskv_indptr_pages = kv_indptr // PAGE_SIZEkv_last_page_lens = torch.full((batch_size,), PAGE_SIZE, device=q.device, dtype=torch.int32)- # Get or compute persistent metadata+ # Cached metadata: depends on shape only, not data content. Safe for leaderboard.(work_meta_data, work_indptr, work_info_set,reduce_indptr, reduce_final_map, reduce_partial_map) = _get_or_compute_metadata(batch_size, total_kv, num_heads, 1, PAGE_SIZE,
scrolls · 102 diff lines total
Best evidence level for this revision: reported
JSON