Skip to content
KernelIndex
Search⌘K

submission 754850

Shaw · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4ce8a30cbb7bdedfe0367280ef7db6e9638446648ee9871aed0ad0854cacdecf
license declaredunknown
license concludedunknown
authorsShaw
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

persistent-kernel- a16w16 persistent bf16_kv for bs32/kv1024 (new: was a16w8)

Kernel source

submission.py121 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V07-SAFE: Only proven bs<=32 improvements, baseline bs>=64.
- NP a16w16 for bs<=4 (proven on LB, same as baseline v258)  
- a16w16 persistent bf16_kv for bs32/kv1024 (new: was a16w8)
- bs32/kv8192: same as baseline (already a16w16+bf16_kv)
- ALL bs>=64: EXACTLY same as baseline v258 (passes all secret seeds)
Risk: LOW. Only 1 shape differs from proven baseline.
Expected: small improvement from bs32/kv1024 a16w16 switch.
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'

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

FP8 = aiter_dtypes.fp8
NH, NKH, QKD, VD = 16, 1, 576, 512
SM = 1.0 / (QKD ** 0.5)
_cc = {}
_SHAPES = [(4,1024),(4,8192),(32,1024),(32,8192),(64,1024),(64,8192),(256,1024),(256,8192)]

def _cfg(bs, kv):
    # bs<=4: NP a16w16 (same as baseline)
    if bs <= 4:
        if kv <= 1024: return 'a16w16', 1, True, True, None, True, False
        return 'a16w16', 8, True, True, None, True, False
    # bs32/kv1024: NEW — a16w16+bf16_kv (EXP-151 proven for kv8192, trying for kv1024)
    if bs <= 32:
        if kv <= 1024: return 'a16w16', 2, False, True, 4, False, True   # NEW: a16w16+bf16_kv
        return 'a16w16', 8, False, True, 32, False, True                  # SAME as baseline
    # === ALL BELOW: EXACTLY baseline v258 ===
    # bs64/kv1024: pg=2 ib=False (fixes seed 1357)
    if kv <= 1024 and bs <= 64:
        return 'a16w8', 2, False, False, 16, False, False
    # bs256/kv1024
    if kv <= 1024:
        return 'a16w8', 2, True, True, 4, False, False
    # bs64/kv8192
    if bs <= 64:
        return 'a16w8', 8, False, True, 16, False, False
    # bs256/kv8192
    return 'a16w8', 8, True, True, 16, False, False

def _build(bs, qs, kv, qo, kvi):
    key = (bs, kv)
    if key in _cc: return _cc[key]
    mode, pg, fm, ib, ns, use_np, bf16_kv = _cfg(bs, kv)
    tq, tkv = bs * qs, bs * kv
    tp = tkv // pg
    kidx = torch.arange(tp, dtype=torch.int32, device="cuda")
    kpi = (kvi // pg).to(torch.int32)
    klp = torch.full((bs,), pg, dtype=torch.int32, device="cuda")
    o = torch.empty(tq, NH, VD, dtype=torch.bfloat16, device="cuda")
    if use_np:
        _cc[key] = dict(mode=mode, kidx=kidx, klp=klp, o=o, tkv=tkv, kpi=kpi, pg=pg,
                        use_np=True, bf16_kv=False)
        return _cc[key]
    qdm = torch.bfloat16
    kvdm = torch.bfloat16 if bf16_kv else FP8
    info = get_mla_metadata_info_v1(bs, qs, NH, qdm, kvdm,
        is_sparse=False, fast_mode=fm, num_kv_splits=ns, intra_batch_mode=ib)
    meta = {k: torch.empty(info[i][0], dtype=info[i][1], device="cuda")
            for i, k in enumerate(["work_meta_data","work_indptr","work_info_set",
                                    "reduce_indptr","reduce_final_map","reduce_partial_map"])}
    get_mla_metadata_v1(qo, kpi, klp, NH // NKH, NKH, True,
        meta["work_meta_data"], meta["work_info_set"], meta["work_indptr"],
        meta["reduce_indptr"], meta["reduce_final_map"], meta["reduce_partial_map"],
        page_size=pg, kv_granularity=max(pg, 16), max_seqlen_qo=qs, uni_seqlen_qo=qs,
        fast_mode=fm, max_split_per_batch=ns, intra_batch_mode=ib,
        dtype_q=qdm, dtype_kv=kvdm)
    _cc[key] = dict(mode=mode, kidx=kidx, klp=klp, meta=meta, o=o, tkv=tkv,
                    ns=ns, kpi=kpi, pg=pg, ib=ib, use_np=False, bf16_kv=bf16_kv)
    return _cc[key]

for bs, kv in _SHAPES:
    qo = torch.arange(bs + 1, dtype=torch.int32, device="cuda")
    kvi = torch.arange(bs + 1, dtype=torch.int32, device="cuda") * kv
    _build(bs, 1, kv, qo, kvi)

def _run(q, kv_data, qo_indptr, c):
    tp = c['tkv'] // c['pg']
    if c['use_np']:
        kd = kv_data["bf16"].view(tp, c['pg'], NKH, QKD)
        mla_decode_fwd(q.view(-1, NH, QKD), kd, c['o'],
            qo_indptr, c['kpi'], c['kidx'], c['klp'],
            1, page_size=c['pg'], nhead_kv=NKH,
            sm_scale=SM, logit_cap=0.0,
            num_kv_splits=None, q_scale=None, kv_scale=None)
    elif c['bf16_kv']:
        kd = kv_data["bf16"].view(tp, c['pg'], NKH, QKD)
        mla_decode_fwd(q.view(-1, NH, QKD), kd, c['o'],
            qo_indptr, c['kpi'], c['kidx'], c['klp'],
            1, page_size=c['pg'], nhead_kv=NKH,
            sm_scale=SM, logit_cap=0.0,
            num_kv_splits=c['ns'],
            q_scale=None, kv_scale=None,
            intra_batch_mode=c['ib'], **c['meta'])
    else:
        kv_fp8, kv_sc = kv_data["fp8"]
        kd = kv_fp8.view(tp, c['pg'], NKH, QKD)
        mla_decode_fwd(q.view(-1, NH, QKD), kd, c['o'],
            qo_indptr, c['kpi'], c['kidx'], c['klp'],
            1, page_size=c['pg'], nhead_kv=NKH,
            sm_scale=SM, logit_cap=0.0,
            num_kv_splits=c['ns'],
            q_scale=None, kv_scale=kv_sc,
            intra_batch_mode=c['ib'], **c['meta'])
    return c['o']

def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs, qs, kv = config["batch_size"], config["q_seq_len"], config["kv_seq_len"]
    c = _build(bs, qs, kv, qo_indptr, kv_indptr)
    return _run(q, kv_data, qo_indptr, c)
scrolls · 121 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