Skip to content
KernelIndex
Search⌘K

submission 645796

John Hahn · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 167 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-645796?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
AMD Instinct MI355X
36.0µs
#81 of 766
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0d4529dfcba35146a8076839bd1592df1b1de541047d236530cb0357e112f475
license declaredunknown
license concludedunknown
authorsJohn Hahn
imported2026-08-15

Kernel source

submission.py167 lines
#!POPCORN leaderboard amd-mixed-mla
#!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.
"""
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

_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),
}

_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 = _TUNE.get((bs, kvsl), (1, 8, bs <= 4))
    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, True,
        work_metadata, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=ps,
        kv_granularity=max(ps, 16),
        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 · 167 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 645036.

⋯ 1 unchanged lines
#!POPCORN gpu MI355X
"""
- v21_best_combo: Cherry-pick best config per shape from v12 + v19.
- kv=1024: keep v12 splits (32 for bs=4, 16 for rest) — fast_mode for bs=4.
- kv=8192: use v19 splits=16 everywhere (was 32 for bs=4,32 in v12).
+ 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.
"""
import torch
from task import input_t, output_t
⋯ 7 unchanged lines
FP8_DTYPE = aiter_dtypes.fp8
_TUNE = {
- (4, 1024): (1, 32, True),
+ (4, 1024): (1, 8, True),
(4, 8192): (8, 16, True),
- (32, 1024): (1, 16, False),
- (32, 8192): (8, 16, False),
- (64, 1024): (1, 16, False),
- (64, 8192): (8, 16, False),
- (256, 1024): (1, 16, False),
- (256, 8192): (8, 16, False),
+ (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),
}
_cache = {}
⋯ 8 unchanged lines
if _q_scale is None:
_q_scale = torch.ones(1, dtype=torch.float32, device=dev)
- ps, num_splits, fast_mode = _TUNE.get((bs, kvsl), (1, 32, bs <= 4))
+ ps, num_splits, fast_mode = _TUNE.get((bs, kvsl), (1, 8, bs <= 4))
nkv = 1
effective_bs = bs * qsl
⋯ 101 unchanged lines
if qsl > 1:
kv_buf = kv_buf.repeat_interleave(qsl, dim=0)
effective_bs = bs * qsl
- if kvsl % ps == 0:
- kv_4d = kv_buf.reshape(effective_bs * pages_per_seq, ps, 1, QK_DIM)
- else:
- pad_len = pages_per_seq * ps - kvsl
- kv_buf = torch.nn.functional.pad(kv_buf, (0, 0, 0, 0, 0, pad_len))
- kv_4d = kv_buf.reshape(effective_bs * pages_per_seq, ps, 1, QK_DIM)
+ kv_4d = kv_buf.reshape(effective_bs * pages_per_seq, ps, 1, QK_DIM)
else:
- if qsl > 1:
- kv_buf = kv_fp8_raw.view(bs, kvsl, 1, QK_DIM)
- kv_4d = kv_buf.repeat(qsl, 1, 1, 1).reshape(bs * qsl * kvsl, 1, 1, QK_DIM)
- else:
- kv_4d = kv_fp8_raw.view(total_kv, 1, 1, QK_DIM)
+ kv_4d = kv_fp8_raw.view(total_kv, 1, 1, QK_DIM)
o = c['o']
mla_decode_fwd(
scrolls · 64 diff lines total

Best evidence level for this revision: reported

JSON