Skip to content
KernelIndex
Search⌘K

submission 747564

Chivier · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d329016f8140217bc487434b75596c4ea6b114035649a8addcd033cfb781caaa
license declaredunknown
license concludedunknown
authorsChivier
imported2026-08-26

Kernel source

submission.py95 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA V28 — Hybrid: Triton bf16 (kv≤1024) + V18 FP8 ASM (kv>1024).
V26 showed Triton bf16 wins on kv=1024: 21-44µs vs ASM 27-49µs.
V18 FP8 ASM proven on kv=8192: 36-312µs (ranked 76µs).
V28 uses V18's exact ASM path (proven metadata caching) for large kv."""
import torch
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1, dtypes as _d
from aiter.ops.quant import dynamic_per_tensor_quant as _dq
from aiter.ops.triton.attention.mla_decode_rope import decode_attention_fwd_grouped_rope

_F = _d.fp8
_S = 1.0 / (576 ** 0.5)

_SP = {
    (4, 1024): 8, (4, 8192): 16,
    (32, 1024): 8, (32, 8192): 8,
    (64, 1024): 4, (64, 8192): 4,
    (256, 1024): 1, (256, 8192): 1,
}
_ct = {}  # Triton bf16 cache
_ca = {}  # ASM FP8 cache (V18-style)

def custom_kernel(data):
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    f8, ks = kv_data["fp8"]
    T = f8.shape[0]
    kl = T // bs
    ns = _SP.get((bs, kl), 8)

    if kl <= 1024:
        # ── Triton bf16 path (20-44µs for kv=1024) ──
        kv_bf16 = kv_data["bf16"]
        k = (bs, T, ns)
        if k not in _ct:
            ki = torch.arange(T, dtype=torch.int32, device="cuda")
            o = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
            al = torch.empty((bs, 16, ns, 513), dtype=torch.float32, device="cuda")
            dummy = torch.empty(0, device="cuda")
            _ct[k] = (ki, o, al, dummy)
        ki, o, al, dummy = _ct[k]

        decode_attention_fwd_grouped_rope(
            q.view(bs, 16, 576), kv_bf16, kv_bf16[:, :, :512], o,
            kv_indptr, ki, dummy, 512, 64, dummy, dummy,
            al, ns, _S,
            logit_cap=0.0, use_rope=False, is_neox_style=False,
        )
        return o
    else:
        # ── V18 FP8 ASM path (proven 36-312µs for kv=8192) ──
        k = (bs, T, ns)
        if k not in _ca:
            ki = torch.arange(T, dtype=torch.int32, device="cuda")
            kl_t = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
            o = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
            qf = torch.empty(q.shape, dtype=_F, device="cuda")
            qs = torch.empty(1, dtype=torch.float32, device="cuda")

            info = get_mla_metadata_info_v1(
                bs, 1, 16, _F, _F,
                is_sparse=False, fast_mode=False,
                num_kv_splits=ns, intra_batch_mode=True)
            work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
            wm, wi, wis, ri, rfm, rpm = work
            get_mla_metadata_v1(
                qo_indptr, kv_indptr, kl_t,
                16, 1, True, wm, wis, wi, ri, rfm, rpm,
                page_size=1, kv_granularity=16,
                max_seqlen_qo=1, uni_seqlen_qo=1,
                fast_mode=False, max_split_per_batch=ns,
                intra_batch_mode=True, dtype_q=_F, dtype_kv=_F)
            _ca[k] = (ki, kl_t, o, qf, qs, wm, wi, wis, ri, rfm, rpm)

        ki, kl_t, o, qf, qs, wm, wi, wis, ri, rfm, rpm = _ca[k]

        _dq(qf, q, qs)
        kv4 = f8.view(T, 1, 1, 576)

        mla_decode_fwd(
            qf.view(-1, 16, 576), kv4, o,
            qo_indptr, kv_indptr, ki, kl_t,
            1, page_size=1, nhead_kv=1,
            sm_scale=_S, logit_cap=0.0,
            num_kv_splits=ns,
            q_scale=qs, kv_scale=ks,
            intra_batch_mode=True,
            work_meta_data=wm, work_indptr=wi, work_info_set=wis,
            reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
        )
        return o
scrolls · 95 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