Skip to content
KernelIndex
Search⌘K

submission 754469

Leon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c6e910e45543e2dd48c6b4e6e8bae7d6cc64dee6d92a54306f0cad8e0b539ce5
license declaredunknown
license concludedunknown
authorsLeon
imported2026-08-15

Techniques

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

persistent-kernel1. Bypass mla_decode_fwd() Python overhead for persistent shapes

Kernel source

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

"""EXP-261: Direct API calls + pre-allocated buffers.
Source-level optimization (MOE teammate methodology):
1. Bypass mla_decode_fwd() Python overhead for persistent shapes
2. Pre-allocate logits/attn_lse buffers (eliminate torch.empty per call)
3. Direct aiter.mla_decode_stage1_asm_fwd + aiter.mla_reduce_v1
4. Keep 1-warmup for JIT, no self-timing loop
5. All dispatch params match EXP-259 (proven config)
"""

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):
    # Returns: (mode, pg, fm, ib, ns, use_np, bf16_kv)
    if bs <= 4:
        if kv <= 1024: return 'a16w16', 1, True, True, None, True, False
        return 'a16w16', 8, True, True, None, True, False
    if bs <= 32:
        if kv <= 1024: return 'a16w8', 2, True, True, 8, False, False
        return 'a16w16', 8, False, True, 32, False, True
    if kv <= 1024 and bs <= 64:
        return 'a16w8', 2, False, False, 16, False, False
    if kv <= 1024:
        return 'a16w8', 2, True, True, 4, False, False
    if bs <= 64:
        return 'a16w8', 8, False, True, 16, False, False
    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)
    # Pre-allocate intermediate buffers for direct API calls
    rpm = meta["reduce_partial_map"].size(0)
    logits = torch.empty((rpm, 1, NH, VD), dtype=torch.float32, device="cuda")
    attn_lse = torch.empty((rpm, 1, NH, 1), dtype=torch.float32, device="cuda")
    _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,
                    logits=logits, attn_lse=attn_lse)
    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)
    else:
        # Direct bottom-level API: skip mla_decode_fwd Python overhead + torch.empty
        m = c['meta']
        if c['bf16_kv']:
            kd = kv_data["bf16"].view(tp, c['pg'], NKH, QKD)
            kv_sc = None
        else:
            kv_fp8, kv_sc = kv_data["fp8"]
            kd = kv_fp8.view(tp, c['pg'], NKH, QKD)
        aiter.mla_decode_stage1_asm_fwd(
            q.view(-1, NH, QKD), kd,
            qo_indptr, c['kpi'], c['kidx'], c['klp'],
            None,  # num_kv_splits_indptr (persistent mode)
            m['work_meta_data'], m['work_indptr'], m['work_info_set'],
            1, c['pg'], NKH, SM,
            c['logits'], c['attn_lse'], c['o'],
            None, kv_sc)  # q_scale, kv_scale (runner has 19 params, no lse)
        aiter.mla_reduce_v1(
            c['logits'], c['attn_lse'],
            m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],
            1, c['o'], None)
    return c['o']


_warmed = set()

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)
    key = (bs, kv)
    if key not in _warmed:
        _warmed.add(key)
        _run(q, kv_data, qo_indptr, c)
        torch.cuda.synchronize()
    return _run(q, kv_data, qo_indptr, c)
scrolls · 136 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