Skip to content
KernelIndex
Search⌘K

submission 673461

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mla_v67.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-673461?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
55.6µs
#191 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

persistent-kernel"""v67: a16w8 persistent pg2 — bf16 Q + fp8 KV, page_size=2, persistent mode.

Kernel source

mla_v67.py111 lines
#!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)."""

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

_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, nt, qo_indptr, dev):
    k = ("a16w8p", bs, sl)
    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=False,
            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=False,
            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 <= 4 and sl <= 1024:
        # Tiny: bf16 non-persistent pg1 (proven by v41)
        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:
        # All other: a16w8 persistent pg2
        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)
        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 · 111 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 663004.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """v41: Full bf16 non-persistent MLA decode.
- Skip fp8 quantization entirely — single kernel launch with bf16 Q and KV.
- Overhead savings (no quant, no metadata, no reduce) outweigh 2x bandwidth cost."""
+ """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)."""
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_DIM = 576
- _V_DIM = 512
- _SM_SC = 1.0 / (_QK_DIM ** 0.5)
+ _QK = 576
+ _VD = 512
+ _SC = 1.0 / (_QK ** 0.5)
+ _PS = 2
+ _NSPLIT = 32
- _shape_bufs = {}
+ _cache = {}
- def _get_shape_buffers(bs, seq_len, dev):
- k = (bs, seq_len)
- if k not in _shape_bufs:
- n = bs * seq_len
- page_ids = torch.arange(n, dtype=torch.int32, device=dev)
- seq_lens = torch.full((bs,), seq_len, dtype=torch.int32, device=dev)
- _shape_bufs[k] = (page_ids, seq_lens)
- return _shape_bufs[k]
+ 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, nt, qo_indptr, dev):
+ k = ("a16w8p", bs, sl)
+ 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=False,
+ 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=False,
+ 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"])
- seq_len = int(config["kv_seq_len"])
- n_tokens = q.shape[0]
+ sl = int(config["kv_seq_len"])
+ nt = q.shape[0]
+ dev = q.device
- q_reshaped = q.view(n_tokens, _NH, _QK_DIM)
- kv_raw = kv_data["bf16"]
- kv_paged = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])
+ q_r = q.view(nt, _NH, _QK)
+ out = torch.empty((nt, _NH, _VD), dtype=torch.bfloat16, device=dev)
- page_ids, seq_lens = _get_shape_buffers(bs, seq_len, q.device)
- out = torch.empty((n_tokens, _NH, _V_DIM), dtype=torch.bfloat16, device=q.device)
+ if bs <= 4 and sl <= 1024:
+ # Tiny: bf16 non-persistent pg1 (proven by v41)
+ 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:
+ # All other: a16w8 persistent pg2
+ 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)
+ 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,
+ )
- mla_decode_fwd(
- q_reshaped, kv_paged, out,
- qo_indptr, kv_indptr,
- page_ids, seq_lens, 1,
- page_size=1, nhead_kv=_NKV, sm_scale=_SM_SC,
- intra_batch_mode=False,
- )
return out
scrolls · 140 diff lines total

Best evidence level for this revision: reported

JSON