submission 619946
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 171 lines, June 9 Researcher Reciprocity License v1.0.
v50.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-619946?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:507b71e0ca2f971486284ffac08e13042cb9dc290c4b33c495b2b0f388501d3e
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15
Kernel source
v50.py171 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v50 — v47 + aggressive tensor caching to eliminate GPU kernel launch overhead.
All shape-dependent tensors (kv_indices, kv_last_page_lens, kv_indptr_pages,
num_kv_splits_indptr) are cached in module-level dicts keyed by shape params.
Eliminates 3-5 GPU kernel launches (torch.arange, torch.ones, etc.) per call.
"""
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 = {}
_tensor_cache = {}
def _get_cached_tensors(key, create_fn):
cached = _tensor_cache.get(key)
if cached is None:
cached = create_fn()
_tensor_cache[key] = cached
return cached
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):
batch_size = config['batch_size']
total_kv = kv_bf16.shape[0]
kv_buffer = kv_bf16.unsqueeze(1)
tensors = _get_cached_tensors(
('bf16', batch_size, total_kv),
lambda: {
'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=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],
max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
)
def _run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config):
batch_size = config['batch_size']
total_kv = kv_fp8_data.shape[0]
q_fp8 = q.to(torch.float8_e4m3fn)
kv_buffer = kv_fp8_data.unsqueeze(1)
tensors = _get_cached_tensors(
('fp8np', batch_size, total_kv),
lambda: {
'q_scale': torch.ones(1, dtype=torch.float32, device=q.device),
'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),
'num_kv_splits_indptr': torch.arange(batch_size + 1, dtype=torch.int32, device=q.device),
}
)
mla_decode_fwd(
q=q_fp8, kv_buffer=kv_buffer, o=output,
qo_indptr=qo_indptr, kv_indptr=kv_indptr,
kv_indices=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],
max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
num_kv_splits=1, num_kv_splits_indptr=tensors['num_kv_splits_indptr'],
q_scale=tensors['q_scale'], kv_scale=kv_fp8_scale,
)
def _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):
batch_size = config['batch_size']
num_heads = config['num_heads']
total_kv = kv_fp8_data.shape[0]
num_pages = total_kv // page_size
kv_buffer = kv_fp8_data.view(num_pages, page_size, 1, 576)
tensors = _get_cached_tensors(
('a16w8', batch_size, total_kv, page_size),
lambda: {
'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_splits, page_size,
torch.bfloat16, aiter_dtypes.fp8,
qo_indptr, tensors['kv_indptr_pages'], tensors['kv_last_page_lens'], q.device,
)
mla_decode_fwd(
q, kv_buffer, output,
qo_indptr, tensors['kv_indptr_pages'], tensors['kv_indices'], tensors['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_splits,
q_scale=None, 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 kv_seq_len <= 1024 and batch_size <= 4:
_run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
elif kv_seq_len <= 1024 and batch_size == 64:
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
_run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config)
elif kv_seq_len <= 1024:
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
splits = 16 if batch_size <= 32 else 8
_run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=2, num_splits=splits)
else:
kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
if batch_size <= 4:
splits = 16
elif batch_size <= 32:
splits = 8
else:
splits = 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 · 171 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 618890.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """MLA v47 — v46 + FP8 NP for bs=64/kv=1024 + splits=32 for bs=4/kv=8192.+ """MLA v50 — v47 + aggressive tensor caching to eliminate GPU kernel launch overhead.- Changes from v46:- bs=64, kv=1024 → FP8 non-persist splits=1 (was BF16, 50→40µs expected)- bs=4, kv=8192 → splits=32 (was 16, testing higher CU utilization)+ All shape-dependent tensors (kv_indices, kv_last_page_lens, kv_indptr_pages,+ num_kv_splits_indptr) are cached in module-level dicts keyed by shape params.+ Eliminates 3-5 GPU kernel launches (torch.arange, torch.ones, etc.) per call."""from task import input_t, output_timport torch⋯ 2 unchanged linesfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1_meta_cache = {}+ _tensor_cache = {}+ def _get_cached_tensors(key, create_fn):+ cached = _tensor_cache.get(key)+ if cached is None:+ cached = create_fn()+ _tensor_cache[key] = cached+ return cached++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))⋯ 28 unchanged linesdef _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):- """BF16 non-persistent — exact correctness."""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)+ tensors = _get_cached_tensors(+ ('bf16', batch_size, total_kv),+ lambda: {+ '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,+ kv_indices=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],)def _run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config):- """FP8+FP8 non-persistent splits=1 — single kernel, no reduce."""batch_size = config['batch_size']total_kv = kv_fp8_data.shape[0]q_fp8 = q.to(torch.float8_e4m3fn)- q_scale = torch.ones(1, dtype=torch.float32, device=q.device)-kv_buffer = kv_fp8_data.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)- num_kv_splits = 1- num_kv_splits_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=q.device)+ tensors = _get_cached_tensors(+ ('fp8np', batch_size, total_kv),+ lambda: {+ 'q_scale': torch.ones(1, dtype=torch.float32, device=q.device),+ '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),+ 'num_kv_splits_indptr': torch.arange(batch_size + 1, dtype=torch.int32, device=q.device),+ }+ )mla_decode_fwd(q=q_fp8, 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,+ kv_indices=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],- num_kv_splits=num_kv_splits, num_kv_splits_indptr=num_kv_splits_indptr,- q_scale=q_scale, kv_scale=kv_fp8_scale,+ num_kv_splits=1, num_kv_splits_indptr=tensors['num_kv_splits_indptr'],+ q_scale=tensors['q_scale'], kv_scale=kv_fp8_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."""batch_size = config['batch_size']num_heads = config['num_heads']total_kv = kv_fp8_data.shape[0]num_pages = total_kv // page_sizekv_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)+ tensors = _get_cached_tensors(+ ('a16w8', batch_size, total_kv, page_size),+ lambda: {+ '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_splits, page_size,torch.bfloat16, aiter_dtypes.fp8,- qo_indptr, kv_indptr_pages, kv_last_page_lens, q.device,+ qo_indptr, tensors['kv_indptr_pages'], tensors['kv_last_page_lens'], q.device,)mla_decode_fwd(q, kv_buffer, output,- qo_indptr, kv_indptr_pages, kv_indices, kv_last_page_lens,+ qo_indptr, tensors['kv_indptr_pages'], tensors['kv_indices'], tensors['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_splits,q_scale=None, kv_scale=kv_fp8_scale,⋯ 12 unchanged linesoutput = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)if kv_seq_len <= 1024 and batch_size <= 4:- # BF16 non-persistent: fastest for tiny shapes_run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)elif kv_seq_len <= 1024 and batch_size == 64:- # FP8 non-persistent splits=1: faster than BF16 (40µs vs 50µs)kv_fp8_data, kv_fp8_scale = kv_data["fp8"]_run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config)elif kv_seq_len <= 1024:- # a16w8 persistent ps=2 for bs=32 and bs=256kv_fp8_data, kv_fp8_scale = kv_data["fp8"]splits = 16 if batch_size <= 32 else 8_run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=2, num_splits=splits)else:- # a16w8 persistent ps=8 for kv=8192 with per-bs optimal splitskv_fp8_data, kv_fp8_scale = kv_data["fp8"]if batch_size <= 4:- splits = 16 # splits=32 tested neutral, keep 16+ splits = 16elif batch_size <= 32:splits = 8else:
scrolls · 156 diff lines total
Best evidence level for this revision: reported
JSON