submission 723450
olezhka_007 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 93 lines, June 9 Researcher Reciprocity License v1.0.
probe_optimal_ps_apr6.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-723450?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:c8654019653c8c0820170360ee5528c93be142d93b90c7cbb629aff71d1661e0
license declaredunknown
license concludedunknown
authorsolezhka_007
imported2026-08-15
Kernel source
probe_optimal_ps_apr6.py93 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""[PROBE] Optimal page_size: ps=4 for bs<=64 kv=8192, ps=8 for bs=256 kv=8192.
Benchmark data:
ps=4 all 8k: (4,8192) 24.1µs (-16%), (32,8192) 36.9µs (-29%), (64,8192) 50.9µs (-37%), (256,8192) 118µs (-33%)
ps=8 (256,8192): 82.1µs (-53%)
Combined: ps=4 for small batches + ps=8 for bs=256 → estimated geomean ~34.8µs (from 43.6µs = 20% improvement)
"""
import torch
from task import input_t, output_t
from aiter import (mla_decode_stage1_asm_fwd, mla_reduce_v1,
get_mla_metadata_info_v1, get_mla_metadata_v1)
from aiter import dtypes as aiter_dtypes
FP8 = aiter_dtypes.fp8; FM = True; IBM = True
SPLITS = {
(4, 1024): 16,
(4, 8192): 16,
(32, 1024): 1,
(32, 8192): 4,
(64, 1024): 1,
(64, 8192): 4,
(256, 1024): 1,
(256, 8192): 1,
}
_C = {}
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]; nq = 16; nkv = 1
dv = config["v_head_dim"]; sm = config["sm_scale"]
kvl = config["kv_seq_len"]; total_q = q.shape[0]
# Optimal page_size per shape
if bs == 256 and kvl == 8192:
ps = 8 # ps=8 gave 82.1µs (from 176µs)
elif kvl == 8192:
ps = 4 # ps=4 gave -16% to -37% for smaller batches
elif kvl == 1024 and bs >= 32:
ps = 2
else:
ps = 1
ns = SPLITS.get((bs, kvl), 16)
k = (bs, kvl)
if k not in _C:
dev = q.device
qs = torch.ones((1,), dtype=torch.float32, device=dev)
o = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device=dev)
sd = torch.empty((total_q, ns, nq, dv), dtype=torch.float32, device=dev)
sl = torch.empty((total_q, ns, nq, 1), dtype=torch.float32, device=dev)
ppb = kvl // ps; tp = ppb * bs
kl = torch.full((bs,), ps, dtype=torch.int32, device=dev)
ki = torch.arange(0, bs + 1, dtype=torch.int32, device=dev) * ppb
kx = torch.arange(tp, dtype=torch.int32, device=dev)
info = get_mla_metadata_info_v1(bs, 1, nq, FP8, FP8,
is_sparse=False, fast_mode=FM, num_kv_splits=ns, intra_batch_mode=IBM)
w = [torch.empty(s, dtype=t, device=dev) for s, t in info]
wm, wi, wis, ri, rfm, rpm = w
get_mla_metadata_v1(qo_indptr, ki, kl, nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm,
page_size=ps, kv_granularity=max(ps, 16),
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=FM, max_split_per_batch=ns,
intra_batch_mode=IBM, dtype_q=FP8, dtype_kv=FP8)
_C[k] = {
'qs': qs, 'o': o, 'sd': sd, 'sl': sl,
'kl': kl, 'ki': ki, 'kx': kx, 'tp': tp, 'ps': ps,
'wm': wm, 'wi': wi, 'wis': wis,
'ri': ri, 'rfm': rfm, 'rpm': rpm, 'ns': ns,
}
c = _C[k]
qf = q.to(FP8)
kvf, kvs = kv_data["fp8"]
mla_decode_stage1_asm_fwd(
qf, kvf.view(c['tp'], c['ps'], nkv, kvf.shape[-1]),
qo_indptr, c['ki'], c['kx'], c['kl'],
None, c['wm'], c['wi'], c['wis'],
1, c['ps'], nkv, sm, c['sd'], c['sl'], c['o'], c['qs'], kvs)
if c['ns'] > 1:
mla_reduce_v1(c['sd'], c['sl'], c['ri'], c['rfm'], c['rpm'], 1, c['o'])
return c['o']
scrolls · 93 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