Skip to content
KernelIndex
Search⌘K

submission 632302

inference_and_chill · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4dedec7b57c2a888f973cfa03cf904874e2639301cb2545527c1e7ba589a64c9
license declaredunknown
license concludedunknown
authorsinference_and_chill
imported2026-08-15

Kernel source

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

"""MLA decode: FP8 aiter wrapper with static Q quant + metadata reuse."""

from __future__ import annotations

import torch
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
from aiter.ops.quant import dynamic_per_tensor_quant, static_per_tensor_quant
from task import input_t, output_t

# ---------------------------------------------------------------------------
# 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

FP8_DTYPE = aiter_dtypes.fp8

# Per-case tuning: (batch_size, kv_seq_len) -> (num_kv_splits, kv_granularity)
_TUNE = {
    (4, 1024): (16, 16),
    (4, 8192): (16, 16),
    (32, 1024): (8, 16),
    (32, 8192): (24, 16),
    (64, 1024): (8, 16),
    (64, 8192): (24, 16),
    (256, 1024): (8, 16),
    (256, 8192): (24, 16),
}

# ---------------------------------------------------------------------------
# Caches — keyed to avoid per-call allocations
# ---------------------------------------------------------------------------
_case_cache: dict[tuple, dict] = {}

_q_ref: torch.Tensor | None = None
_q_fp8_buf: torch.Tensor | None = None
_q_scale_buf: torch.Tensor | None = None
_use_static_quant: bool = False

_kv_ref: torch.Tensor | None = None
_kv_4d: torch.Tensor | None = None


# ---------------------------------------------------------------------------
# Entry point — FP8 ASM path
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    global _q_ref, _q_fp8_buf, _q_scale_buf, _use_static_quant
    global _kv_ref, _kv_4d

    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]

    # Cached Q FP8 quantization
    if q is not _q_ref:
        if _q_fp8_buf is None or _q_fp8_buf.shape != q.shape:
            _q_fp8_buf = torch.empty(q.shape, dtype=FP8_DTYPE, device="cuda")
            _q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
            _use_static_quant = False

        if _use_static_quant:
            static_per_tensor_quant(_q_fp8_buf, q, _q_scale_buf)
        else:
            dynamic_per_tensor_quant(_q_fp8_buf, q, _q_scale_buf)
            _use_static_quant = True
        _q_ref = q

    # FP8 KV — cache view
    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    if kv_buffer_fp8 is not _kv_ref:
        _kv_4d = kv_buffer_fp8.view(
            kv_buffer_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_buffer_fp8.shape[-1]
        )
        _kv_ref = kv_buffer_fp8

    # Per-case cached data (metadata, indices, output, kv_last_page_len)
    case_key = (batch_size, kv_seq_len)
    cd = _case_cache.get(case_key)
    if cd is None:
        tune = _TUNE.get(case_key, (32, 16))
        num_kv_splits, kv_gran = tune

        total_kv_len = batch_size * kv_seq_len
        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
        kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

        info = get_mla_metadata_info_v1(
            batch_size,
            1,
            NUM_HEADS,
            FP8_DTYPE,
            FP8_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]
        wm, wi, wis, ri, rfm, rpm = work

        get_mla_metadata_v1(
            qo_indptr,
            kv_indptr,
            kv_last_page_len,
            NUM_HEADS // NUM_KV_HEADS,
            NUM_KV_HEADS,
            True,
            wm,
            wis,
            wi,
            ri,
            rfm,
            rpm,
            page_size=PAGE_SIZE,
            kv_granularity=kv_gran,
            max_seqlen_qo=1,
            uni_seqlen_qo=1,
            fast_mode=False,
            max_split_per_batch=num_kv_splits,
            intra_batch_mode=True,
            dtype_q=FP8_DTYPE,
            dtype_kv=FP8_DTYPE,
        )

        total_q = batch_size
        o = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")

        cd = {
            "nks": num_kv_splits,
            "ki": kv_indices,
            "klpl": kv_last_page_len,
            "wm": wm,
            "wi": wi,
            "wis": wis,
            "ri": ri,
            "rfm": rfm,
            "rpm": rpm,
            "o": o,
        }
        _case_cache[case_key] = cd

    o = cd["o"]
    mla_decode_fwd(
        _q_fp8_buf,
        _kv_4d,
        o,
        qo_indptr,
        kv_indptr,
        cd["ki"],
        cd["klpl"],
        1,  # max_seqlen_q
        page_size=PAGE_SIZE,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=cd["nks"],
        q_scale=_q_scale_buf,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        work_meta_data=cd["wm"],
        work_indptr=cd["wi"],
        work_info_set=cd["wis"],
        reduce_indptr=cd["ri"],
        reduce_final_map=cd["rfm"],
        reduce_partial_map=cd["rpm"],
    )
    return o
scrolls · 181 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