submission 648412
John Hahn · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 169 lines, June 9 Researcher Reciprocity License v1.0.
v36_pg2_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-648412?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:409b82e7441aa2e1dfdc131d602926b828b045edecfb70b52d5786870b181f82
license declaredunknown
license concludedunknown
authorsJohn Hahn
imported2026-08-15
Kernel source
v36_pg2_hybrid.py169 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
v36_pg2_hybrid: ps=2 for kv=1024 with per-case kv_granularity.
- kv=1024: ps=2, kv_granularity=2 (huge speedup, especially bs=256: 67→45µs)
- kv=8192: ps=8, kv_granularity=16 (keep original, kvgran=8 regressed bs=4)
"""
import torch
from task import input_t, output_t
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
QK_DIM = 576
V_DIM = 512
SM_SCALE = 1.0 / (QK_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
# (ps, splits, fast_mode, kv_granularity)
_TUNE = {
(4, 1024): (2, 8, True, 2),
(4, 8192): (8, 16, True, 16),
(32, 1024): (2, 8, False, 2),
(32, 8192): (8, 8, False, 16),
(64, 1024): (2, 8, False, 2),
(64, 8192): (8, 8, False, 16),
(256, 1024): (2, 8, False, 2),
(256, 8192): (8, 8, False, 16),
}
_cache = {}
_q_scale = None
def _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, dev):
global _q_scale
if cfg_key in _cache:
return _cache[cfg_key]
if _q_scale is None:
_q_scale = torch.ones(1, dtype=torch.float32, device=dev)
ps, num_splits, fast_mode, kvgran = _TUNE.get((bs, kvsl), (1, 8, bs <= 4, 16))
nkv = 1
effective_bs = bs * qsl
if qsl > 1:
eff_kv_indptr = torch.zeros(effective_bs + 1, dtype=torch.int32, device=dev)
for i in range(bs):
kv_len = kv_indptr[i + 1].item() - kv_indptr[i].item()
for j in range(qsl):
idx = i * qsl + j
eff_kv_indptr[idx + 1] = eff_kv_indptr[idx] + kv_len
else:
eff_kv_indptr = kv_indptr
eff_qo_indptr = torch.arange(effective_bs + 1, dtype=torch.int32, device=dev)
total_eff_kv = int(eff_kv_indptr[-1].item())
if ps > 1:
pages_per_seq = (kvsl + ps - 1) // ps
kv_last_page_len = torch.full((effective_bs,),
kvsl % ps if kvsl % ps != 0 else ps,
dtype=torch.int32, device=dev)
kv_indices = torch.arange(effective_bs * pages_per_seq, dtype=torch.int32, device=dev)
paged_kv_indptr = torch.arange(effective_bs + 1, dtype=torch.int32, device=dev) * pages_per_seq
else:
kv_last_page_len = torch.full((effective_bs,), ps, dtype=torch.int32, device=dev)
kv_indices = torch.arange(total_eff_kv, dtype=torch.int32, device=dev)
paged_kv_indptr = eff_kv_indptr
info = get_mla_metadata_info_v1(
effective_bs, 1, nh, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=fast_mode,
num_kv_splits=num_splits, intra_batch_mode=(not fast_mode),
)
work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
eff_qo_indptr, paged_kv_indptr, kv_last_page_len,
nh // nkv, nkv, False,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=ps,
kv_granularity=kvgran,
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=fast_mode,
max_split_per_batch=num_splits,
intra_batch_mode=(not fast_mode),
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
total_q = bs * qsl
o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)
q_fp8_buf = torch.empty((total_q, nh, QK_DIM), dtype=FP8_DTYPE, device=dev)
entry = {
'ps': ps, 'num_splits': num_splits, 'fast_mode': fast_mode,
'eff_qo_indptr': eff_qo_indptr,
'paged_kv_indptr': paged_kv_indptr,
'kv_indices': kv_indices,
'kv_last_page_len': kv_last_page_len,
'meta': {
'work_meta_data': work_metadata,
'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,
},
'o': o,
'q_fp8_buf': q_fp8_buf,
'qsl': qsl, 'bs': bs, 'nh': nh, 'nkv': nkv,
'effective_bs': effective_bs,
'kvsl': kvsl,
}
_cache[cfg_key] = entry
return entry
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
nh = config["num_heads"]
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
cfg_key = (bs, qsl, kvsl, nh)
c = _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, q.device)
q_fp8 = c['q_fp8_buf']
q_fp8.copy_(q.view(bs * qsl, nh, QK_DIM))
kv_fp8_raw, kv_fp8_scale = kv_data["fp8"]
kv_scale = kv_fp8_scale.view(1) if kv_fp8_scale.numel() == 1 else kv_fp8_scale
ps = c['ps']
total_kv = bs * kvsl
if ps > 1:
pages_per_seq = (kvsl + ps - 1) // ps
effective_bs = c['effective_bs']
kv_buf = kv_fp8_raw.view(bs, kvsl, 1, QK_DIM)
if qsl > 1:
kv_buf = kv_buf.repeat_interleave(qsl, dim=0)
effective_bs = bs * qsl
kv_4d = kv_buf.reshape(effective_bs * pages_per_seq, ps, 1, QK_DIM)
else:
kv_4d = kv_fp8_raw.view(total_kv, 1, 1, QK_DIM)
o = c['o']
mla_decode_fwd(
q_fp8, kv_4d, o,
c['eff_qo_indptr'], c['paged_kv_indptr'],
c['kv_indices'], c['kv_last_page_len'],
1, page_size=ps, nhead_kv=1,
sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=c['num_splits'],
q_scale=_q_scale, kv_scale=kv_scale,
intra_batch_mode=(not c['fast_mode']),
**c['meta'],
)
return o
scrolls · 169 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 645796.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- v25_best_combo2: Cherry-pick best splits per case from v21(splits=16) + v22(splits=8).- splits=8 wins most cases; splits=16 wins bs=4/kv=8192.+ v36_pg2_hybrid: ps=2 for kv=1024 with per-case kv_granularity.+ - kv=1024: ps=2, kv_granularity=2 (huge speedup, especially bs=256: 67→45µs)+ - kv=8192: ps=8, kv_granularity=16 (keep original, kvgran=8 regressed bs=4)"""import torchfrom task import input_t, output_t⋯ 6 unchanged linesSM_SCALE = 1.0 / (QK_DIM ** 0.5)FP8_DTYPE = aiter_dtypes.fp8+ # (ps, splits, fast_mode, kv_granularity)_TUNE = {- (4, 1024): (1, 8, True),- (4, 8192): (8, 16, True),- (32, 1024): (1, 8, False),- (32, 8192): (8, 8, False),- (64, 1024): (1, 8, False),- (64, 8192): (8, 8, False),- (256, 1024): (1, 8, False),- (256, 8192): (8, 8, False),+ (4, 1024): (2, 8, True, 2),+ (4, 8192): (8, 16, True, 16),+ (32, 1024): (2, 8, False, 2),+ (32, 8192): (8, 8, False, 16),+ (64, 1024): (2, 8, False, 2),+ (64, 8192): (8, 8, False, 16),+ (256, 1024): (2, 8, False, 2),+ (256, 8192): (8, 8, False, 16),}_cache = {}⋯ 8 unchanged linesif _q_scale is None:_q_scale = torch.ones(1, dtype=torch.float32, device=dev)- ps, num_splits, fast_mode = _TUNE.get((bs, kvsl), (1, 8, bs <= 4))+ ps, num_splits, fast_mode, kvgran = _TUNE.get((bs, kvsl), (1, 8, bs <= 4, 16))nkv = 1effective_bs = bs * qsl⋯ 33 unchanged linesget_mla_metadata_v1(eff_qo_indptr, paged_kv_indptr, kv_last_page_len,- nh // nkv, nkv, True,+ nh // nkv, nkv, False,work_metadata, work_info_set, work_indptr,reduce_indptr, reduce_final_map, reduce_partial_map,page_size=ps,- kv_granularity=max(ps, 16),+ kv_granularity=kvgran,max_seqlen_qo=1,uni_seqlen_qo=1,fast_mode=fast_mode,
scrolls · 60 diff lines total
Best evidence level for this revision: reported
JSON