Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
34.6µs
#63 of 766
2026-04-04

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