submission 616673
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 125 lines, June 9 Researcher Reciprocity License v1.0.
v31.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-616673?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:e3bb08f19bc92bf53d432029f80aaf804b529aa3446aef70c9e431466c88692f
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 v29 — Optimal hybrid: BF16 non-persist + a16w8 (BF16 Q + FP8 KV) persistent.Kernel source
v31.py125 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v29 — Optimal hybrid: BF16 non-persist + a16w8 (BF16 Q + FP8 KV) persistent.
Dispatch by best-of per shape:
bs≤4, kv≤1024 → BF16 non-persistent (26µs, fastest for tiny shapes)
everything else → a16w8 persistent (45-98µs, FP8 KV bandwidth + no Q quant overhead)
ps=1 for kv≤1024, ps=8 for kv≥8192
"""
from task import input_t, output_t
import torch
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_meta_cache = {}
def _get_or_make_metadata(batch_size, total_kv, num_heads, nhead_kv, num_splits, page_size,
q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_lens, device):
key = (batch_size, total_kv, num_heads, num_splits, page_size, str(q_dtype), str(kv_dtype))
cached = _meta_cache.get(key)
if cached is not None:
return cached
info = get_mla_metadata_info_v1(
batch_size, 1, num_heads, q_dtype, kv_dtype,
is_sparse=False, fast_mode=True,
num_kv_splits=num_splits, intra_batch_mode=False,
)
work = [torch.empty(s, dtype=t, device=device) for s, t in info]
(wmd, wi, wis, ri, rfm, rpm) = work
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=num_splits, intra_batch_mode=False,
dtype_q=q_dtype, dtype_kv=kv_dtype,
)
result = dict(
work_meta_data=wmd, work_indptr=wi, work_info_set=wis,
reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
)
_meta_cache[key] = result
return result
def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):
"""BF16 non-persistent — fastest for tiny shapes."""
batch_size = config['batch_size']
total_kv = kv_bf16.shape[0]
kv_buffer = kv_bf16.unsqueeze(1)
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)
mla_decode_fwd(
q=q, kv_buffer=kv_buffer, o=output,
qo_indptr=qo_indptr, kv_indptr=kv_indptr,
kv_indices=kv_indices, kv_last_page_lens=kv_last_page_lens,
max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
)
def _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):
"""a16w8: BF16 Q + FP8 KV persistent — no Q quant, FP8 bandwidth."""
batch_size = config['batch_size']
num_heads = config['num_heads']
total_kv = kv_fp8_data.shape[0]
NUM_KV_SPLITS = num_splits
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)
meta = _get_or_make_metadata(
batch_size, total_kv, num_heads, 1, NUM_KV_SPLITS, page_size,
torch.bfloat16, aiter_dtypes.fp8,
qo_indptr, kv_indptr_pages, kv_last_page_lens, q.device,
)
mla_decode_fwd(
q, # BF16 Q — no quantization!
kv_buffer, # FP8 KV
output,
qo_indptr, kv_indptr_pages, kv_indices, kv_last_page_lens,
1, page_size=page_size, nhead_kv=1, sm_scale=config['sm_scale'],
logit_cap=0.0, num_kv_splits=NUM_KV_SPLITS,
q_scale=None, # BF16 Q doesn't need scale
kv_scale=kv_fp8_scale,
intra_batch_mode=False, **meta,
)
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']
v_head_dim = config['v_head_dim']
kv_seq_len = config['kv_seq_len']
output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)
if batch_size <= 4 and kv_seq_len <= 1024:
# BF16 non-persistent: fastest for tiny shapes (26µs)
_run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
elif kv_seq_len <= 1024:
# a16w8 persistent ps=1: adaptive splits (more for small batch, fewer for large)
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
splits = 8 if batch_size <= 32 else 4
_run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=1, num_splits=splits)
else:
# a16w8 persistent ps=8: adaptive splits
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
splits = 8 if batch_size <= 32 else 4
_run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)
return output
scrolls · 125 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 600564.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """MLA Decode v12 — FP8 Q + FP8 KV with persistent mode.+ """MLA v29 — Optimal hybrid: BF16 non-persist + a16w8 (BF16 Q + FP8 KV) persistent.- 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+ Dispatch by best-of per shape:+ bs≤4, kv≤1024 → BF16 non-persistent (26µs, fastest for tiny shapes)+ everything else → a16w8 persistent (45-98µs, FP8 KV bandwidth + no Q quant overhead)+ ps=1 for kv≤1024, ps=8 for kv≥8192"""from task import input_t, output_timport torch- import aiter- from aiter import dtypesfrom aiter.mla import mla_decode_fwd+ from aiter import dtypes as aiter_dtypes+ from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1_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)++ def _get_or_make_metadata(batch_size, total_kv, num_heads, nhead_kv, num_splits, page_size,+ q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_lens, device):+ key = (batch_size, total_kv, num_heads, num_splits, page_size, str(q_dtype), str(kv_dtype))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,+ info = get_mla_metadata_info_v1(+ batch_size, 1, num_heads, q_dtype, kv_dtype,+ is_sparse=False, fast_mode=True,+ num_kv_splits=num_splits, intra_batch_mode=False,)+ work = [torch.empty(s, dtype=t, device=device) for s, t in info]+ (wmd, wi, wis, ri, rfm, rpm) = work- 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(+ 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,+ max_split_per_batch=num_splits, intra_batch_mode=False,+ dtype_q=q_dtype, dtype_kv=kv_dtype,)- result = (wmd, wi, wis, ri, rfm, rpm)+ result = dict(+ work_meta_data=wmd, work_indptr=wi, work_info_set=wis,+ reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,+ )_meta_cache[key] = resultreturn 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"]+ def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):+ """BF16 non-persistent — fastest for tiny shapes."""+ batch_size = config['batch_size']+ total_kv = kv_bf16.shape[0]+ kv_buffer = kv_bf16.unsqueeze(1)+ 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)+ mla_decode_fwd(+ q=q, kv_buffer=kv_buffer, o=output,+ qo_indptr=qo_indptr, kv_indptr=kv_indptr,+ kv_indices=kv_indices, kv_last_page_lens=kv_last_page_lens,+ max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],+ )+++ def _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):+ """a16w8: BF16 Q + FP8 KV persistent — no Q quant, FP8 bandwidth."""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]+ NUM_KV_SPLITS = num_splits- # 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)-+ 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)+ 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,+ meta = _get_or_make_metadata(+ batch_size, total_kv, num_heads, 1, NUM_KV_SPLITS, page_size,+ torch.bfloat16, aiter_dtypes.fp8,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,+ q, # BF16 Q — no quantization!+ kv_buffer, # FP8 KV+ output,+ qo_indptr, kv_indptr_pages, kv_indices, kv_last_page_lens,+ 1, page_size=page_size, nhead_kv=1, sm_scale=config['sm_scale'],+ logit_cap=0.0, num_kv_splits=NUM_KV_SPLITS,+ q_scale=None, # BF16 Q doesn't need scalekv_scale=kv_fp8_scale,+ intra_batch_mode=False, **meta,)++ 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']+ v_head_dim = config['v_head_dim']+ kv_seq_len = config['kv_seq_len']++ output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)++ if batch_size <= 4 and kv_seq_len <= 1024:+ # BF16 non-persistent: fastest for tiny shapes (26µs)+ _run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)+ elif kv_seq_len <= 1024:+ # a16w8 persistent ps=1: adaptive splits (more for small batch, fewer for large)+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]+ splits = 8 if batch_size <= 32 else 4+ _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=1, num_splits=splits)+ else:+ # a16w8 persistent ps=8: adaptive splits+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]+ splits = 8 if batch_size <= 32 else 4+ _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)+return output
scrolls · 198 diff lines total
Best evidence level for this revision: reported
JSON