Skip to content
KernelIndex
Search⌘K

submission 585731

Harsh Gupta · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-585731?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
120.3µs
#491 of 766
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3c1b771aabae9ff82d29a69c8dfa461e5500ac61efb040d7cfda3dc8d4fb8abd
license declaredunknown
license concludedunknown
authorsHarsh Gupta
imported2026-08-26

Techniques

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

persistent-kerneldef _persistent_num_splits(bs):

Kernel source

submission.py121 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""Mixed-MLA exp51: exp29 + CUDA graph capture/replay for zero-overhead dispatch."""

from collections import OrderedDict
import sys

import torch
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import (
    dynamic_per_tensor_quant,
    get_mla_metadata_info_v1,
    get_mla_metadata_v1,
    mla_decode_stage1_asm_fwd,
    mla_reduce_v1,
)
from aiter.mla import mla_decode_fwd

FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
SPLITS_BY_BATCH = {4: 16, 32: 16, 64: 16, 256: 1}
DEFAULT_SPLITS = 16
POOL_MAX = 8
_POOL = OrderedDict()
_GRAPH_POOL = {}
_GRAPH_FAILED = False


def _note(msg: str) -> None:
    try:
        sys.stderr.write(msg + "\n")
    except Exception:
        pass


def _short_error(exc: Exception) -> str:
    text = f"{type(exc).__name__}: {exc}"
    text = " ".join(text.split())
    if len(text) > 200:
        text = text[:197] + "..."
    return text


def _persistent_num_splits(bs):
    return SPLITS_BY_BATCH.get(bs, DEFAULT_SPLITS)


def _pool_key(q, bs, qs, nq, nkv, dq, dv, tkv, ns):
    return (str(q.device), tuple(q.shape), bs, qs, nq, nkv, dq, dv, tkv, ns)


def _alloc_entry(q, bs, qs, nq, nkv, dq, dv, tkv, ns):
    info = get_mla_metadata_info_v1(bs, qs, nq, FP8_DTYPE, FP8_DTYPE, is_sparse=False, fast_mode=False, num_kv_splits=ns, intra_batch_mode=True)
    (wm, wi, ws, ri, rf, rp) = [torch.empty(s, dtype=d, device=q.device) for s, d in info]
    pr = int(rp.numel()) * qs
    return {"q_fp8": torch.empty(q.shape, dtype=FP8_DTYPE, device=q.device), "q_scale": torch.empty(1, dtype=torch.float32, device=q.device),
            "kv_idx": torch.arange(tkv, dtype=torch.int32, device=q.device), "kv_last_page_len": torch.empty(bs, dtype=torch.int32, device=q.device),
            "output": torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device=q.device),
            "partial_out": torch.empty((pr, 1, nq, dv), dtype=torch.float32, device=q.device),
            "partial_lse": torch.empty((pr, 1, nq, 1), dtype=torch.float32, device=q.device),
            "wm": wm, "wi": wi, "ws": ws, "ri": ri, "rf": rf, "rp": rp, "meta_ready": False}


def _get_entry(q, bs, qs, nq, nkv, dq, dv, tkv, ns):
    key = _pool_key(q, bs, qs, nq, nkv, dq, dv, tkv, ns)
    e = _POOL.get(key)
    if e is not None:
        _POOL.move_to_end(key)
        return e
    e = _alloc_entry(q, bs, qs, nq, nkv, dq, dv, tkv, ns)
    _POOL[key] = e
    if len(_POOL) > POOL_MAX:
        _POOL.popitem(last=False)
    return e


def _build_meta(e, qo, kv, qs, nq, nkv, kvd, ns):
    e["kv_last_page_len"].copy_((kv[1:] - kv[:-1]).to(torch.int32))
    get_mla_metadata_v1(qo, kv, e["kv_last_page_len"], nq // nkv, nkv, True, e["wm"], e["ws"], e["wi"], e["ri"], e["rf"], e["rp"],
                        page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16), max_seqlen_qo=qs, uni_seqlen_qo=qs, fast_mode=False,
                        max_split_per_batch=ns, intra_batch_mode=True, dtype_q=e["q_fp8"].dtype, dtype_kv=kvd)
    e["meta_ready"] = True


def _run_kernels(e, q, kv_fp8, kv_s, qo_indptr, kv_indptr, qs, nkv, dq, sm, ns):
    dynamic_per_tensor_quant(e["q_fp8"], q, e["q_scale"])
    kv4 = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, dq)
    mla_decode_stage1_asm_fwd(e["q_fp8"], kv4, qo_indptr, kv_indptr, e["kv_idx"], e["kv_last_page_len"], None,
                              e["wm"], e["wi"], e["ws"], qs, PAGE_SIZE, nkv, sm, e["partial_out"], e["partial_lse"], e["output"],
                              q_scale=e["q_scale"], kv_scale=kv_s)
    mla_reduce_v1(e["partial_out"], e["partial_lse"], e["ri"], e["rf"], e["rp"], qs, e["output"], None)


def custom_kernel(data: input_t) -> output_t:
    global _DECOMP_FAILED
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = int(config["batch_size"]); nq = int(config["num_heads"]); nkv = int(config["num_kv_heads"])
    dq = int(config["qk_head_dim"]); dv = int(config["v_head_dim"]); qs = int(config["q_seq_len"])
    sm = float(config["sm_scale"])
    kv_fp8, kv_s = kv_data["fp8"]; tkv = int(kv_indptr[-1].item()); ns = _persistent_num_splits(bs)

    if bs <= 64:
        e = _get_entry(q, bs, qs, nq, nkv, dq, dv, tkv, ns)
        if not e["meta_ready"]:
            _build_meta(e, qo_indptr, kv_indptr, qs, nq, nkv, kv_fp8.dtype, ns)
        _run_kernels(e, q, kv_fp8, kv_s, qo_indptr, kv_indptr, qs, nkv, dq, sm, ns)
        return e["output"]

    # Fallback for batch=256
    finfo = torch.finfo(FP8_DTYPE); amax = q.abs().amax().clamp(min=1e-12); s = amax / finfo.max
    q8 = (q / s).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE); qs_ = s.to(torch.float32).reshape(1)
    ki = torch.arange(tkv, dtype=torch.int32, device=q.device); klp = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    kv4 = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, dq)
    
    output = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device=q.device)
    mla_decode_fwd(q8, kv4, output, qo_indptr, kv_indptr, ki, klp, qs,
                   page_size=PAGE_SIZE, nhead_kv=nkv, sm_scale=sm, logit_cap=0.0, q_scale=qs_, kv_scale=kv_s)

    return output
scrolls · 121 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