Skip to content
KernelIndex
Search⌘K

submission 682562

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mla_v93.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-682562?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
54.2µs
#180 of 766
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8ae5968021b92eb890c83c73ef72bb604f980747a510a145fccd5402ae70dac9
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15

Kernel source

mla_v93.py123 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""v93: Per-shape fast_mode — cherry-pick best of v67 and v90.

fast_mode=True helps: (4,8192) -3us, (64,1024) -3.7us
fast_mode=False helps: (32,8192) -1us, (64,8192) -2.5us
Neutral: (256,*), (4,1024), (32,1024)

Use True for small batch+large kv and medium batch+small kv.
Use False for large batch or large kv+large batch.
"""

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

FP8 = aiter_dtypes.fp8
_NH = 16
_NKV = 1
_QK = 576
_VD = 512
_SC = 1.0 / (_QK ** 0.5)
_PS = 2
_NSPLIT = 32

# Per-shape: (fast_mode,) — True where it helps, False where it hurts
_FM = {
    (4, 8192): True,
    (64, 1024): True,
    # Everything else uses False (v67 default, better for large shapes)
}

_cache = {}


def _get_pg1(bs, sl, dev):
    k = ("pg1", bs, sl)
    if k not in _cache:
        _cache[k] = (
            torch.arange(bs * sl, dtype=torch.int32, device=dev),
            torch.full((bs,), sl, dtype=torch.int32, device=dev),
        )
    return _cache[k]


def _get_a16w8_persist(bs, sl, fm, qo_indptr, dev):
    k = ("a16w8p", bs, sl, fm)
    if k not in _cache:
        npages = (bs * sl) // _PS
        kv_idx = torch.arange(npages, dtype=torch.int32, device=dev)
        ki = torch.arange(bs + 1, dtype=torch.int32, device=dev) * (sl // _PS)
        lp = torch.full((bs,), _PS, dtype=torch.int32, device=dev)

        info = get_mla_metadata_info_v1(
            bs, 1, _NH, torch.bfloat16, FP8,
            is_sparse=False, fast_mode=fm,
            num_kv_splits=_NSPLIT, intra_batch_mode=True,
        )
        work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
        wm, wi, wis, ri, rfm, rpm = work

        get_mla_metadata_v1(
            qo_indptr, ki, lp,
            _NH // _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=_NSPLIT,
            intra_batch_mode=True,
            dtype_q=torch.bfloat16, dtype_kv=FP8,
        )

        meta = {
            "work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
            "reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm,
        }
        _cache[k] = (meta, kv_idx, ki, lp)
    return _cache[k]


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = int(config["batch_size"])
    sl = int(config["kv_seq_len"])
    nt = q.shape[0]
    dev = q.device

    q_r = q.view(nt, _NH, _QK)
    out = torch.empty((nt, _NH, _VD), dtype=torch.bfloat16, device=dev)

    if bs <= 32 and sl <= 1024:
        # bf16 non-persistent pg1
        kv_raw = kv_data["bf16"]
        kv_4d = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])
        pg, lp = _get_pg1(bs, sl, dev)
        mla_decode_fwd(
            q_r, kv_4d, out, qo_indptr, kv_indptr,
            pg, lp, 1,
            page_size=1, nhead_kv=_NKV, sm_scale=_SC,
            intra_batch_mode=False,
        )
    else:
        # a16w8 persist pg2 with per-shape fast_mode
        fm = _FM.get((bs, sl), False)
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_4d = kv_fp8.view(-1, _PS, _NKV, kv_fp8.shape[-1])
        meta, kv_idx, ki, lp = _get_a16w8_persist(bs, sl, fm, qo_indptr, dev)
        mla_decode_fwd(
            q_r, kv_4d, out, qo_indptr, ki,
            kv_idx, lp, 1,
            page_size=_PS, nhead_kv=_NKV, sm_scale=_SC,
            logit_cap=0.0, num_kv_splits=_NSPLIT,
            kv_scale=kv_scale,
            intra_batch_mode=True, **meta,
        )

    return out
scrolls · 123 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 673461.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """v67: a16w8 persistent pg2 — bf16 Q + fp8 KV, page_size=2, persistent mode.
- Non-persistent a16w8 pg2 has no ASM kernel (ps:0 error).
- Persistent mode dispatches through metadata scheduler which supports pg2.
- bf16 non-persistent pg1 for tiny shapes (proven by v41)."""
+ """v93: Per-shape fast_mode — cherry-pick best of v67 and v90.
+ fast_mode=True helps: (4,8192) -3us, (64,1024) -3.7us
+ fast_mode=False helps: (32,8192) -1us, (64,8192) -2.5us
+ Neutral: (256,*), (4,1024), (32,1024)
+
+ Use True for small batch+large kv and medium batch+small kv.
+ Use False for large batch or large kv+large batch.
+ """
+
import torch
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
⋯ 9 unchanged lines
_PS = 2
_NSPLIT = 32
+ # Per-shape: (fast_mode,) — True where it helps, False where it hurts
+ _FM = {
+ (4, 8192): True,
+ (64, 1024): True,
+ # Everything else uses False (v67 default, better for large shapes)
+ }
+
_cache = {}
⋯ 7 unchanged lines
return _cache[k]
- def _get_a16w8_persist(bs, sl, nt, qo_indptr, dev):
- k = ("a16w8p", bs, sl)
+ def _get_a16w8_persist(bs, sl, fm, qo_indptr, dev):
+ k = ("a16w8p", bs, sl, fm)
if k not in _cache:
npages = (bs * sl) // _PS
kv_idx = torch.arange(npages, dtype=torch.int32, device=dev)
⋯ 2 unchanged lines
info = get_mla_metadata_info_v1(
bs, 1, _NH, torch.bfloat16, FP8,
- is_sparse=False, fast_mode=False,
+ is_sparse=False, fast_mode=fm,
num_kv_splits=_NSPLIT, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
⋯ 6 unchanged lines
page_size=_PS,
kv_granularity=max(_PS, 16),
max_seqlen_qo=1, uni_seqlen_qo=1,
- fast_mode=False,
+ fast_mode=fm,
max_split_per_batch=_NSPLIT,
intra_batch_mode=True,
dtype_q=torch.bfloat16, dtype_kv=FP8,
⋯ 17 unchanged lines
q_r = q.view(nt, _NH, _QK)
out = torch.empty((nt, _NH, _VD), dtype=torch.bfloat16, device=dev)
- if bs <= 4 and sl <= 1024:
- # Tiny: bf16 non-persistent pg1 (proven by v41)
+ if bs <= 32 and sl <= 1024:
+ # bf16 non-persistent pg1
kv_raw = kv_data["bf16"]
kv_4d = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])
pg, lp = _get_pg1(bs, sl, dev)
⋯ 4 unchanged lines
intra_batch_mode=False,
)
else:
- # All other: a16w8 persistent pg2
+ # a16w8 persist pg2 with per-shape fast_mode
+ fm = _FM.get((bs, sl), False)
kv_fp8, kv_scale = kv_data["fp8"]
kv_4d = kv_fp8.view(-1, _PS, _NKV, kv_fp8.shape[-1])
- meta, kv_idx, ki, lp = _get_a16w8_persist(bs, sl, nt, qo_indptr, dev)
+ meta, kv_idx, ki, lp = _get_a16w8_persist(bs, sl, fm, qo_indptr, dev)
mla_decode_fwd(
q_r, kv_4d, out, qo_indptr, ki,
kv_idx, lp, 1,
page_size=_PS, nhead_kv=_NKV, sm_scale=_SC,
logit_cap=0.0, num_kv_splits=_NSPLIT,
kv_scale=kv_scale,
- intra_batch_mode=True,
- **meta,
+ intra_batch_mode=True, **meta,
)
return out
scrolls · 97 diff lines total

Best evidence level for this revision: reported

JSON