submission 598454
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 150 lines, June 9 Researcher Reciprocity License v1.0.
v12_fp8_q_kv.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-598454?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:5f0c604718b62e345bc3390cd5eb651b69f50872b8400b81cb74310d8d39264d
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.py150 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):
"""Compute and cache persistent mode metadata."""
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),
) = 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,
)
# 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)
# Populate metadata
aiter.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,
)
result = (work_meta_data, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map)
_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)
# Get or compute persistent metadata
(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 · 150 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 592885.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """MLA Decode v1 — use aiter's ASM-optimized MLA decode."""+ """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_timport torch+ import aiter+ from aiter import dtypesfrom 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):+ """Compute and cache persistent mode metadata."""+ 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),+ ) = 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,+ )++ # 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)++ # Populate metadata+ aiter.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,+ )++ result = (work_meta_data, work_indptr, work_info_set,+ reduce_indptr, reduce_final_map, reduce_partial_map)+ _meta_cache[key] = result+ return result++def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = data- kv_bf16 = kv_data["bf16"]+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]batch_size = config['batch_size']num_heads = config['num_heads']- qk_head_dim = config['qk_head_dim']v_head_dim = config['v_head_dim']sm_scale = config['sm_scale']- total_kv = kv_bf16.shape[0]+ 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)- # Treat contiguous KV as paged with page_size=1- # kv_bf16 shape: [total_kv, nhead_kv, head_dim] → reshape to [total_kv, 1, nhead_kv, head_dim]- if kv_bf16.dim() == 3:- kv_buffer = kv_bf16.unsqueeze(1) # [total_kv, 1, nhead_kv, head_dim]- elif kv_bf16.dim() == 2:- kv_buffer = kv_bf16.unsqueeze(1).unsqueeze(2) # [total_kv, 1, 1, head_dim]- else:- kv_buffer = kv_bf16+ # 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(total_kv, device=q.device, dtype=torch.int32)- kv_last_page_lens = torch.ones(batch_size, device=q.device, dtype=torch.int32)+ 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)+ # Get or compute persistent metadata+ (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,+ q=q_fp8,kv_buffer=kv_buffer,o=output,qo_indptr=qo_indptr,- kv_indptr=kv_indptr,+ kv_indptr=kv_indptr_pages,kv_indices=kv_indices,kv_last_page_lens=kv_last_page_lens,max_seqlen_q=1,- page_size=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 · 166 diff lines total
Best evidence level for this revision: reported
JSON