Skip to content
KernelIndex
Search⌘K

submission 703780

Navid Khazaee · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mixed-mla_68ae08f_benchmark.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-703780?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
74.2µs
#361 of 766
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7bcafeb5cba8cfcf84bd2bfe14783ff7273edbc2f0f4d666434e7d103911d500
license declaredunknown
license concludedunknown
authorsNavid Khazaee
imported2026-08-26

Kernel source

mixed-mla_68ae08f_benchmark.py140 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

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

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

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_FP8_MAX = torch.finfo(FP8_DTYPE).max
_FP8_MIN = torch.finfo(FP8_DTYPE).min

# ---------------------------------------------------------------------------
# All caches keyed by shape tuples — never by data_ptr
# ---------------------------------------------------------------------------
_meta_cache = {}
_meta_filled = {}
_kv_idx_cache = {}
_out_cache = {}
_klp_cache = {}

# ---------------------------------------------------------------------------
# Metadata: allocate once, fill once per shape key
# No .item() calls — all shape info from config dict (pure Python, no GPU sync)
# ---------------------------------------------------------------------------
def _get_meta(bs, kv_seq, mql, nq, nkv, qdt, kvdt, qoi, kvi, nks):
    ck = (bs, kv_seq, mql, nq, nkv, str(qdt), str(kvdt), nks)

    if ck not in _meta_cache:
        info = get_mla_metadata_info_v1(
            bs, mql, nq, qdt, kvdt,
            is_sparse=False, fast_mode=True,
            num_kv_splits=nks, intra_batch_mode=True,
        )
        _meta_cache[ck] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]

    wm, wi, wis, ri, rfm, rpm = _meta_cache[ck]

    if ck not in _meta_filled:
        klp_key = (bs, kv_seq)
        if klp_key not in _klp_cache:
            _klp_cache[klp_key] = torch.full((bs,), kv_seq, dtype=torch.int32, device="cuda")
        klp = _klp_cache[klp_key]

        get_mla_metadata_v1(
            qoi, kvi, klp, nq // nkv, nkv, True,
            wm, wis, wi, ri, rfm, rpm,
            page_size=PAGE_SIZE, kv_granularity=64,
            max_seqlen_qo=mql, uni_seqlen_qo=mql,
            fast_mode=True,
            max_split_per_batch=nks,
            intra_batch_mode=True, dtype_q=qdt, dtype_kv=kvdt,
        )
        _meta_filled[ck] = True

    return {"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
            "reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm}


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    mql = config["q_seq_len"]
    kv_seq = config["kv_seq_len"]
    total_q = q.shape[0]
    tot_kv = bs * kv_seq  # NO .item() — pure Python, no GPU sync

    # Output buffer cache
    okey = (total_q, nq, dv)
    if okey not in _out_cache:
        _out_cache[okey] = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
    o = _out_cache[okey]

    # KV indices cache
    if tot_kv not in _kv_idx_cache:
        _kv_idx_cache[tot_kv] = torch.arange(tot_kv, dtype=torch.int32, device="cuda")
    kv_idx = _kv_idx_cache[tot_kv]

    # Per-shape config: (num_kv_splits, use_bf16)
    # bf16 skips fp8 quant overhead (~15us from kernel launches)
    # bf16 wins when compute/latency-bound; fp8 wins when bandwidth-bound
    if bs <= 4:
        nks, use_bf16 = (16, True) if kv_seq <= 1024 else (32, True)
    elif bs <= 32:
        nks, use_bf16 = (8, True) if kv_seq <= 1024 else (16, False)
    elif bs <= 64:
        nks, use_bf16 = (8, True) if kv_seq <= 1024 else (4, False)
    else:
        nks, use_bf16 = (4 if kv_seq <= 1024 else 2, False)

    if use_bf16:
        kv_buf = kv_data["bf16"]
        q_in, q_scale, kv_scale = q, None, None
    else:
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_buf = kv_fp8
        # Simplified fp8 quant: removed unnecessary clamps (save 2 kernel launches)
        # By construction |q/scale| <= FP8_MAX, so clamping is redundant
        scale = q.abs().amax() / _FP8_MAX
        q_in = (q / scale).to(FP8_DTYPE)
        q_scale = scale.to(torch.float32).reshape(1)

    kv4d = kv_buf.view(tot_kv, PAGE_SIZE, nkv, kv_buf.shape[-1])

    meta = _get_meta(bs, kv_seq, mql, nq, nkv, q_in.dtype, kv_buf.dtype,
                     qo_indptr, kv_indptr, nks)

    # KV last page length — cached, no .item()
    klp_key = (bs, kv_seq)
    if klp_key not in _klp_cache:
        _klp_cache[klp_key] = torch.full((bs,), kv_seq, dtype=torch.int32, device="cuda")
    klp = _klp_cache[klp_key]

    mla_decode_fwd(
        q_in.view(-1, nq, dq), kv4d, o,
        qo_indptr, kv_indptr, kv_idx, klp, mql,
        page_size=PAGE_SIZE, nhead_kv=nkv,
        sm_scale=SM_SCALE, logit_cap=0.0,
        num_kv_splits=nks,
        q_scale=q_scale, kv_scale=kv_scale,
        intra_batch_mode=True, **meta,
    )
    return o


check_implementation = make_match_reference(custom_kernel, rtol=1e-01, atol=1e-01)
scrolls · 140 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