Skip to content
KernelIndex
Search⌘K

submission 687741

Jianian-Xu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:da6f4eaa0eea0df12725635c61d219193fadc53989dfe978fe45ec38df130f16
license declaredunknown
license concludedunknown
authorsJianian-Xu
imported2026-08-15

Techniques

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

persistent-kernelFor kv<=1024: non-persistent mode with num_kv_splits=1, skips reduce entirely (2 launches).
split-k_SINGLE_SPLIT_KV_THRESHOLD = 0 # 0 = never use single-split

Kernel source

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

import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'

"""
MLA decode — bypasses mla_decode_fwd, calls stage1+reduce directly.
For kv<=1024: non-persistent mode with num_kv_splits=1, skips reduce entirely (2 launches).
For kv>1024: persistent mode with splits + reduce (3 launches).
"""

import torch
from task import input_t, output_t
from utils import make_match_reference

import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter import dynamic_per_tensor_quant as _dpq
from aiter import static_per_tensor_quant as _spq

# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM   # 576
V_HEAD_DIM = KV_LORA_RANK                        # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 16
FP8_DTYPE = aiter_dtypes.fp8

# Threshold: kv_per_seq <= this uses non-persistent single-split (skip reduce)
# DISABLED: non-persistent with 1 split has worse CU utilization than persistent
_SINGLE_SPLIT_KV_THRESHOLD = 0  # 0 = never use single-split

# ---------------------------------------------------------------------------
# Caches
# ---------------------------------------------------------------------------
_metadata_cache: dict = {}
_kv_indices_cache: dict = {}
_kv_last_page_len_cache: dict = {}
_kv_indptr_cache: dict = {}
_qo_indptr_cache: dict = {}
_splits_indptr_cache: dict = {}
_fast_path_cache: dict = {}


def _get_cached_kv_indices(n):
    if n not in _kv_indices_cache:
        _kv_indices_cache[n] = torch.arange(n, dtype=torch.int32, device="cuda")
    return _kv_indices_cache[n]

def _get_cached_kv_last_page_len(total_kv, bs):
    k = (total_kv, bs)
    if k not in _kv_last_page_len_cache:
        _kv_last_page_len_cache[k] = torch.full((bs,), total_kv // bs, dtype=torch.int32, device="cuda")
    return _kv_last_page_len_cache[k]

def _get_cached_indptr(cache, bs, step):
    k = (bs, step)
    if k not in cache:
        cache[k] = torch.arange(0, (bs + 1) * step, step, dtype=torch.int32, device="cuda")
    return cache[k]

def _get_cached_metadata(bs, max_q_len, nq, nkv, q_dtype, kv_dtype,
                         qo_indptr, kv_indptr, kv_last_page_len,
                         total_kv_len, num_kv_splits):
    ck = (bs, max_q_len, nq, nkv, q_dtype, kv_dtype, num_kv_splits, total_kv_len)
    if ck not in _metadata_cache:
        info = get_mla_metadata_info_v1(
            bs, max_q_len, nq, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=num_kv_splits, intra_batch_mode=True,
        )
        work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        d = {
            "work_meta_data": work[0], "work_indptr": work[1], "work_info_set": work[2],
            "reduce_indptr": work[3], "reduce_final_map": work[4], "reduce_partial_map": work[5],
        }
        _metadata_cache[ck] = (d, [None])

    cached, pop_store = _metadata_cache[ck]
    pk = (qo_indptr.data_ptr(), kv_indptr.data_ptr(), kv_last_page_len.data_ptr())
    if pop_store[0] != pk:
        get_mla_metadata_v1(
            qo_indptr, kv_indptr, kv_last_page_len,
            nq // nkv, nkv, True,
            cached["work_meta_data"], cached["work_info_set"], cached["work_indptr"],
            cached["reduce_indptr"], cached["reduce_final_map"], cached["reduce_partial_map"],
            page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
            max_seqlen_qo=max_q_len, uni_seqlen_qo=max_q_len,
            fast_mode=False, max_split_per_batch=num_kv_splits,
            intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )
        pop_store[0] = pk
    return cached


# ---------------------------------------------------------------------------
# Main kernel
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    kv_input = kv_buffer_fp8

    fp_key = (config["batch_size"], kv_input.shape[0])
    fp = _fast_path_cache.get(fp_key)

    if fp is not None:
        # --- Fast path: always run quant (Q/KV data change in leaderboard) ---
        _spq(fp["q_fp8"], q, fp["q_scale"])
        kv_4d = kv_input.view(fp["kv_view"])

        if fp["single_split"]:
            aiter.mla_decode_stage1_asm_fwd(
                fp["q_fp8_3d"], kv_4d,
                fp["qo_indptr"], fp["kv_indptr"], fp["kv_indices"], fp["kv_last_page_len"],
                fp["splits_indptr"], None, None, None,
                1, PAGE_SIZE, fp["nkv"], SM_SCALE,
                fp["logits_view"], fp["attn_lse"], fp["o"],
                q_scale=fp["q_scale"], kv_scale=kv_scale,
            )
        else:
            aiter.mla_decode_stage1_asm_fwd(
                fp["q_fp8_3d"], kv_4d,
                fp["qo_indptr"], fp["kv_indptr"], fp["kv_indices"], fp["kv_last_page_len"],
                None, fp["work_meta_data"], fp["work_indptr"], fp["work_info_set"],
                1, PAGE_SIZE, fp["nkv"], SM_SCALE,
                fp["logits"], fp["attn_lse"], fp["o"],
                q_scale=fp["q_scale"], kv_scale=kv_scale,
            )
            aiter.mla_reduce_v1(
                fp["logits"], fp["attn_lse"],
                fp["reduce_indptr"], fp["reduce_final_map"], fp["reduce_partial_map"],
                1, fp["o"], None,
            )
        return fp["o"]

    # --- Slow path: first call per config ---
    batch_size, total_kv_len = fp_key
    kv_per_seq = total_kv_len // 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_q = q.shape[0]

    # Decide: single-split non-persistent or multi-split persistent
    single_split = (kv_per_seq <= _SINGLE_SPLIT_KV_THRESHOLD)

    # Common buffers
    kv_indptr_c = _get_cached_indptr(_kv_indptr_cache, batch_size, kv_per_seq)
    qo_indptr_c = _get_cached_indptr(_qo_indptr_cache, batch_size, q_seq_len)
    kv_indices = _get_cached_kv_indices(total_kv_len)
    kv_last_page_len = _get_cached_kv_last_page_len(total_kv_len, batch_size)

    q_fp8_buf = torch.empty(q.shape, dtype=FP8_DTYPE, device="cuda")
    q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
    o_buf = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
    q_fp8_3d = q_fp8_buf.view(-1, nq, dq)
    kv_4d = kv_input.view(total_kv_len, PAGE_SIZE, nkv, kv_input.shape[-1])

    _dpq(q_fp8_buf, q, q_scale_buf)

    fp_entry = {
        "q_fp8": q_fp8_buf, "q_fp8_3d": q_fp8_3d, "q_scale": q_scale_buf,
        "kv_view": (total_kv_len, PAGE_SIZE, nkv, kv_input.shape[-1]),
        "kv_4d": kv_4d,
        "o": o_buf,
        "qo_indptr": qo_indptr_c, "kv_indptr": kv_indptr_c,
        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "nkv": nkv, "single_split": single_split,
        "last_q_ptr": None,
    }

    if single_split:
        # Non-persistent, num_kv_splits=1: stage1 writes directly to o
        splits_indptr = _get_cached_indptr(_splits_indptr_cache, batch_size, 1)
        # logits = o.view(total_q, 1, nq, dv) — view of output, no alloc
        logits_view = o_buf.view(total_q, 1, nq, dv)
        attn_lse = torch.empty((total_q, 1, nq, 1), dtype=torch.float32, device="cuda")

        aiter.mla_decode_stage1_asm_fwd(
            q_fp8_3d, kv_4d,
            qo_indptr_c, kv_indptr_c, kv_indices, kv_last_page_len,
            splits_indptr, None, None, None,  # non-persistent
            q_seq_len, PAGE_SIZE, nkv, SM_SCALE,
            logits_view, attn_lse, o_buf,
            q_scale=q_scale_buf, kv_scale=kv_scale,
        )
        # NO reduce needed

        fp_entry.update({
            "splits_indptr": splits_indptr,
            "logits_view": logits_view,
            "attn_lse": attn_lse,
        })
    else:
        # Persistent mode
        num_splits = 64 if batch_size >= 256 else NUM_KV_SPLITS
        meta = _get_cached_metadata(
            batch_size, q_seq_len, nq, nkv,
            FP8_DTYPE, kv_input.dtype,
            qo_indptr_c, kv_indptr_c, kv_last_page_len,
            total_kv_len, num_splits,
        )
        rpm_size = meta["reduce_partial_map"].size(0)
        logits_buf = torch.empty((rpm_size * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda")
        attn_lse_buf = torch.empty((rpm_size * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda")

        aiter.mla_decode_stage1_asm_fwd(
            q_fp8_3d, kv_4d,
            qo_indptr_c, kv_indptr_c, kv_indices, kv_last_page_len,
            None, meta["work_meta_data"], meta["work_indptr"], meta["work_info_set"],
            q_seq_len, PAGE_SIZE, nkv, SM_SCALE,
            logits_buf, attn_lse_buf, o_buf,
            q_scale=q_scale_buf, kv_scale=kv_scale,
        )
        aiter.mla_reduce_v1(
            logits_buf, attn_lse_buf,
            meta["reduce_indptr"], meta["reduce_final_map"], meta["reduce_partial_map"],
            q_seq_len, o_buf, None,
        )

        fp_entry.update({
            "logits": logits_buf, "attn_lse": attn_lse_buf,
            "work_meta_data": meta["work_meta_data"],
            "work_indptr": meta["work_indptr"],
            "work_info_set": meta["work_info_set"],
            "reduce_indptr": meta["reduce_indptr"],
            "reduce_final_map": meta["reduce_final_map"],
            "reduce_partial_map": meta["reduce_partial_map"],
        })

    _fast_path_cache[fp_key] = fp_entry
    return o_buf
scrolls · 245 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