Skip to content
KernelIndex
Search⌘K

submission 716430

somethingobscurefordevstuff · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-716430?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
35.9µs
#80 of 766
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3af2846a6f2f8add6b226670bea510754973649832eef374a1ee88a986c273e1
license declaredunknown
license concludedunknown
authorssomethingobscurefordevstuff
imported2026-08-15

Kernel source

submission.py131 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

# v382 plus mla_decode_fwd only on the medium-short bf16 region (48<=bs<128, kv<=2048).

import aiter
import torch
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
from task import input_t, output_t

NH = 16; NKV = 1; DQ = 576; DV = 512; SM = DQ**-0.5
FP8 = torch.float8_e4m3fn; BF16 = torch.bfloat16
_s1 = aiter.mla_decode_stage1_asm_fwd; _rd = aiter.mla_reduce_v1
_c = {}


def _cfg(bs, kvl):
    large_kv = kvl > 2048
    med_bs = 16 < bs < 128
    medium_short = 48 <= bs < 128 and kvl <= 2048

    use_fp8 = (large_kv and med_bs) or (bs >= 128)
    kv_dt = FP8 if use_fp8 else BF16

    if use_fp8 and large_kv and med_bs:
        fm = False
        ibm = bs >= 48
    else:
        fm = True
        ibm = True

    if bs <= 8 and kvl <= 1024:
        ps = 1
    elif kvl <= 1024:
        ps = 2
    else:
        ps = 8

    if bs >= 128:
        ns = 1
    elif bs >= 48:
        ns = 4
    elif bs >= 16:
        ns = 8
    else:
        ns = 16

    return kv_dt, BF16, fm, ibm, ps, ns, medium_short


def _setup(bs, kvl, dev):
    k = (bs, kvl)
    if k in _c:
        return _c[k]

    if bs <= 8 and kvl <= 1024:
        total = bs * kvl
        qo = torch.arange(bs + 1, dtype=torch.int32, device=dev)
        kvi = torch.arange(bs + 1, dtype=torch.int32, device=dev) * kvl
        klp = torch.full((bs,), 1, dtype=torch.int32, device=dev)
        ki = torch.arange(total, dtype=torch.int32, device=dev)
        o = torch.empty((bs, NH, DV), dtype=BF16, device=dev)
        _c[k] = ("nonpersist", o, qo, kvi, klp, ki)
        return _c[k]

    kv_dt, q_dt, fm, ibm, ps, ns, medium_short = _cfg(bs, kvl)
    pp = kvl // ps
    qo = torch.arange(bs + 1, dtype=torch.int32, device=dev)
    kvi = torch.arange(bs + 1, dtype=torch.int32, device=dev) * pp
    klp = torch.full((bs,), ps, dtype=torch.int32, device=dev)
    ki = torch.arange(bs * pp, dtype=torch.int32, device=dev)
    info = get_mla_metadata_info_v1(
        bs, 1, NH, q_dt, kv_dt,
        is_sparse=False, fast_mode=fm, num_kv_splits=ns, intra_batch_mode=ibm,
    )
    wm, wi, ws, ri, rf, rp = [torch.empty(s, dtype=t, device=dev) for s, t in info]
    kvg = 8 if (bs >= 128 or (bs >= 64 and kvl > 2048)) else max(ps, 16)
    get_mla_metadata_v1(
        qo, kvi, klp, NH, NKV, False, wm, ws, wi, ri, rf, rp,
        page_size=ps, kv_granularity=kvg, max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=fm, max_split_per_batch=ns, intra_batch_mode=ibm, dtype_q=q_dt, dtype_kv=kv_dt,
    )
    o = torch.empty((bs, NH, DV), dtype=BF16, device=dev)
    po = torch.empty((bs, ns, NH, DV), dtype=torch.float32, device=dev)
    pl = torch.empty((bs, ns, NH, 1), dtype=torch.float32, device=dev)
    use_bf16 = (kv_dt == BF16)
    use_fwd_bf16 = medium_short and use_bf16
    _c[k] = ("persist", o, po, pl, qo, kvi, klp, ki, wm, wi, ws, ri, rf, rp, ps, ns, ibm, use_bf16, use_fwd_bf16)
    return _c[k]

def custom_kernel(data: input_t) -> output_t:
    q, kv_data, _, _, _ = data
    bs = q.shape[0]
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_bf16 = kv_data["bf16"]
    kvl = kv_fp8.shape[0] // bs
    entry = _setup(bs, kvl, q.device)
    if entry[0] == "nonpersist":
        _, o, qo, kvi, klp, ki = entry
        kv_buf = kv_bf16.view(bs * kvl, 1, NKV, DQ)
        mla_decode_fwd(
            q=q, kv_buffer=kv_buf, o=o,
            qo_indptr=qo, kv_indptr=kvi, kv_indices=ki, kv_last_page_lens=klp,
            max_seqlen_q=1, page_size=1, nhead_kv=NKV, sm_scale=SM,
        )
        return o

    _, o, po, pl, qo, kvi, klp, ki, wm, wi, ws, ri, rf, rp, ps, ns, ibm, use_bf16, use_fwd_bf16 = entry
    if use_fwd_bf16:
        kv4 = kv_bf16.view(kv_bf16.shape[0] // ps, ps, NKV, DQ)
        mla_decode_fwd(
            q, kv4, o, qo, kvi, ki, klp, 1,
            page_size=ps, nhead_kv=NKV, sm_scale=SM,
            num_kv_splits=ns, q_scale=None, kv_scale=None,
            intra_batch_mode=ibm,
            work_meta_data=wm, work_indptr=wi, work_info_set=ws,
            reduce_indptr=ri, reduce_final_map=rf, reduce_partial_map=rp,
        )
        return o

    if use_bf16:
        kv4 = kv_bf16.view(kv_bf16.shape[0] // ps, ps, NKV, DQ)
        _s1(q, kv4, qo, kvi, ki, klp, None, wm, wi, ws, 1, ps, NKV, SM, po, pl, o, None, None)
    else:
        kv4 = kv_fp8.view(kv_fp8.shape[0] // ps, ps, NKV, DQ)
        _s1(q, kv4, qo, kvi, ki, klp, None, wm, wi, ws, 1, ps, NKV, SM, po, pl, o, None, kv_scale)
    if ns > 1:
        _rd(po, pl, ri, rf, rp, 1, o, None)
    return o
scrolls · 131 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