Skip to content
KernelIndex
Search⌘K

submission 753932

guojun21 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f85c6f1dd2fd18915971b696a66d9d3e35ab514213e276b7eb1e067e0e6251b0
license declaredunknown
license concludedunknown
authorsguojun21
imported2026-08-15

Techniques

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

persistent-kernelPaged MLA decode with per-shape routing, persistent ASM kernels,

Kernel source

submission_v3b_kvg128_256_public.py291 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
Paged MLA decode with per-shape routing, persistent ASM kernels,
non-persistent fast path for small batches, and pybind reduce shortcut.
"""

import os as _os
_os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")

import gc as _gc
import sys as _sys

_gc.disable()
_sys.setswitchinterval(1.0)

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

torch.set_grad_enabled(False)

ATTN_HEADS = 16
KV_HEADS = 1
LATENT_DIM = 512
ROPE_DIM = 64
ATTN_DIM = LATENT_DIM + ROPE_DIM   # 576
OUT_DIM = LATENT_DIM                # 512
SOFTMAX_SCALE = 1.0 / (ATTN_DIM ** 0.5)

FP8_DTYPE = aiter_dtypes.fp8

_BENCH_SHAPES = [
    (4, 1024), (4, 8192),
    (32, 1024), (32, 8192),
    (64, 1024), (64, 8192),
    (256, 1024), (256, 8192),
]

# Per-shape config: (page_size, num_splits, use_nonpersistent, kv_granularity, intra_batch)
_SHAPE_CFG = {
    (4, 1024):   (1, 8, False, 128, False),
    (4, 8192):   (8, 32, False, 128, False),
    (32, 1024):  (2, 1, True, 128, False),
    (32, 8192):  (8, 32, False, 32, True),
    (64, 1024):  (2, 8, False, 128, False),
    (64, 8192):  (8, 32, False, 32, False),
    (256, 1024): (2, 8, False, 32, False),
    (256, 8192): (8, 32, False, 128, False),
}

_LAUNCHERS = None


def _build_launcher(batch_size, kv_seq_len, device, stage1_fn, reduce_fn):
    """Pre-build all tensors and metadata for one (batch, kv_len) shape."""
    total_q = batch_size
    nq = ATTN_HEADS
    nkv = KV_HEADS
    dq = ATTN_DIM
    dv = OUT_DIM
    pg, num_splits, use_np, kvg, intra = _SHAPE_CFG.get(
        (batch_size, kv_seq_len), (1, 32, False, 128, False)
    )

    pages_per_item = kv_seq_len // pg
    qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device)
    kv_indptr = torch.arange(
        0, (batch_size + 1) * pages_per_item, pages_per_item,
        dtype=torch.int32, device=device,
    )
    kv_indices = torch.arange(
        batch_size * pages_per_item, dtype=torch.int32, device=device,
    )
    kv_last = torch.full(
        (batch_size,), kv_seq_len, dtype=torch.int32, device=device,
    )

    output = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device=device)
    q_fp8 = torch.empty((total_q, nq, dq), dtype=FP8_DTYPE, device=device)
    q_scale = torch.ones(1, dtype=torch.float32, device=device)

    sm = SOFTMAX_SCALE
    kv_4d_shape = (-1, pg, nkv, dq)
    _copy_q = q_fp8.copy_
    _kv_cache = [None, None]

    # Non-persistent path: single split, stage1 writes output directly
    if use_np and stage1_fn is not None:
        kv_splits_indptr = torch.arange(
            0, batch_size + 1, dtype=torch.int32, device=device,
        )
        logits_alias = output.view(total_q, 1, nq, dv)
        np_attn_lse = torch.empty(
            (total_q, 1, nq, 1), dtype=torch.float32, device=device,
        )

        def launch(q_raw, kv_raw, kv_scale):
            _copy_q(q_raw)
            if _kv_cache[0] is not kv_raw:
                _kv_cache[0] = kv_raw
                _kv_cache[1] = kv_raw.view(kv_4d_shape)
            stage1_fn(
                q_fp8, _kv_cache[1],
                qo_indptr, kv_indptr, kv_indices, kv_last,
                kv_splits_indptr, None, None, None,
                1, pg, nkv, sm,
                logits_alias, np_attn_lse, output, q_scale, kv_scale,
            )
            return output

        return launch

    # Persistent paged mode: pre-built metadata, stage1 + reduce
    q_dtype = FP8_DTYPE
    kv_dtype = FP8_DTYPE
    info = get_mla_metadata_info_v1(
        batch_size, 1, nq, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=False,
        num_kv_splits=num_splits, intra_batch_mode=intra,
    )
    w_meta, w_indptr, w_info, r_indptr, r_final, r_partial = [
        torch.empty(s, dtype=t, device=device) for s, t in info
    ]
    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last,
        nq // nkv, nkv, True,
        w_meta, w_info, w_indptr, r_indptr, r_final, r_partial,
        page_size=pg,
        kv_granularity=kvg,
        max_seqlen_qo=1,
        uni_seqlen_qo=1,
        fast_mode=False,
        max_split_per_batch=num_splits,
        intra_batch_mode=intra,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    partial_count = int(r_partial.size(0))
    logits = torch.empty(
        (partial_count, 1, nq, dv), dtype=torch.float32, device=device,
    )
    attn_lse = torch.empty(
        (partial_count, 1, nq, 1), dtype=torch.float32, device=device,
    )

    if stage1_fn is not None and reduce_fn is not None:
        def launch(q_raw, kv_raw, kv_scale):
            _copy_q(q_raw)
            if _kv_cache[0] is not kv_raw:
                _kv_cache[0] = kv_raw
                _kv_cache[1] = kv_raw.view(kv_4d_shape)
            stage1_fn(
                q_fp8, _kv_cache[1],
                qo_indptr, kv_indptr, kv_indices, kv_last,
                None, w_meta, w_indptr, w_info,
                1, pg, nkv, sm,
                logits, attn_lse, output, q_scale, kv_scale,
            )
            reduce_fn(
                logits, attn_lse,
                r_indptr, r_final, r_partial,
                1, output, None,
            )
            return output
    else:
        meta = {
            "work_meta_data": w_meta,
            "work_indptr": w_indptr,
            "work_info_set": w_info,
            "reduce_indptr": r_indptr,
            "reduce_final_map": r_final,
            "reduce_partial_map": r_partial,
        }

        def launch(q_raw, kv_raw, kv_scale):
            _copy_q(q_raw)
            if _kv_cache[0] is not kv_raw:
                _kv_cache[0] = kv_raw
                _kv_cache[1] = kv_raw.view(kv_4d_shape)
            mla_decode_fwd(
                q_fp8, _kv_cache[1], output,
                qo_indptr, kv_indptr, kv_indices, kv_last,
                1,
                page_size=pg,
                nhead_kv=nkv,
                sm_scale=sm,
                logit_cap=0.0,
                num_kv_splits=num_splits,
                q_scale=q_scale,
                kv_scale=kv_scale,
                intra_batch_mode=intra,
                **meta,
            )
            return output

    return launch


def _init_launchers(device):
    global _LAUNCHERS

    aiter_ns = getattr(torch.ops, "aiter", None)
    stage1_op = getattr(aiter_ns, "mla_decode_stage1_asm_fwd", None) if aiter_ns else None
    reduce_op = getattr(aiter_ns, "mla_reduce_v1", None) if aiter_ns else None
    stage1 = getattr(stage1_op, "default", stage1_op) if stage1_op else None
    reduce = getattr(reduce_op, "default", reduce_op) if reduce_op else None

    try:
        from aiter.jit.core import get_module as _get_mod
        _reduce_mod = _get_mod("module_mla_reduce")
        _reduce_pb = getattr(_reduce_mod, "mla_reduce_v1", None)
        if _reduce_pb is not None:
            reduce = _reduce_pb
            print("[paged-mla] pybind reduce OK", file=_sys.stderr)
    except Exception:
        pass

    launchers = {}
    for bs, kvl in _BENCH_SHAPES:
        launchers[(bs, bs * kvl)] = _build_launcher(bs, kvl, device, stage1, reduce)

    _LAUNCHERS = launchers
    print(
        f"[paged-mla] {len(launchers)} runners, asm={stage1 is not None}",
        file=_sys.stderr,
    )


def _fallback_decode(q, kv_fp8, kv_scale, qo_indptr, kv_indptr, config):
    """General path for shapes not in pre-built table."""
    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    total_kv = kv_fp8.shape[0]
    kv_seq_len = config.get("kv_seq_len")
    device = q.device

    q_fp8 = q.to(FP8_DTYPE)
    q_scale = torch.ones(1, dtype=torch.float32, device=device)

    kv_4d = kv_fp8.view(total_kv, 1, nkv, kv_fp8.shape[-1])
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    if kv_seq_len is not None:
        kv_last = torch.full(
            (batch_size,), kv_seq_len, dtype=torch.int32, device=device,
        )
    else:
        kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device=device)
    mla_decode_fwd(
        q_fp8.view(-1, nq, dq), kv_4d, o,
        qo_indptr, kv_indptr, kv_indices, kv_last,
        q_seq_len,
        page_size=1,
        nhead_kv=nkv,
        sm_scale=SOFTMAX_SCALE,
        logit_cap=0.0,
        num_kv_splits=32,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
    )
    return o


def custom_kernel(data: input_t) -> output_t:
    global _LAUNCHERS

    q, kv_data, qo_indptr, kv_indptr, config = data

    if _LAUNCHERS is None:
        _init_launchers(q.device)

    kv_fp8, kv_scale = kv_data["fp8"]

    launcher = _LAUNCHERS.get((q.shape[0], kv_fp8.shape[0]))
    if launcher is not None:
        return launcher(q, kv_fp8, kv_scale)

    return _fallback_decode(q, kv_fp8, kv_scale, qo_indptr, kv_indptr, config)
scrolls · 291 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