Skip to content
KernelIndex
Search⌘K

submission 614123

anairdrop · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b89789b775d9a31fad1accb4b58678ced00a3141b226ba4f4d6c446dc8a8f4aa
license declaredunknown
license concludedunknown
authorsanairdrop
imported2026-08-26

Kernel source

submission.py144 lines
"""
Optimized MLA decode: bypass ref_kernel overhead by calling mla_decode_fwd
directly with aggressively cached buffers.

Profiled overhead breakdown (bs=4, kvseq=1024):
  fp8_quant:  31µs  ← uses aiter.per_tensor_quant (was 50µs manual)
  kv_indices:  5µs  ← cached
  kv_last_pg:  8µs  ← cached
  metadata:   13µs  ← cached
  output_alloc: 3µs
  mla_decode: 21µs  ← the actual kernel
"""
import torch
import aiter
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

# MLA constants
PAGE_SIZE = 1
SM_SCALE = 1.0 / (576 ** 0.5)
V_HEAD_DIM = 512
QK_HEAD_DIM = 576
FP8_DTYPE = aiter_dtypes.fp8

# Caches
_cache = {}


def _quantize_fp8(tensor):
    """Use aiter.per_tensor_quant for ~23% faster FP8 quantization."""
    flat = tensor.view(-1, tensor.shape[-1])
    fp8_flat, scale = aiter.per_tensor_quant(flat, quant_dtype=FP8_DTYPE)
    return fp8_flat.view(tensor.shape), scale


def _get_or_build_cache(bs, qseq, nq, nkv, total_kv_len, num_splits,
                        q_dtype, kv_dtype, qo_indptr, kv_indptr):
    key = (bs, qseq, nq, nkv, total_kv_len, num_splits)
    if key in _cache:
        return _cache[key]

    # Build all cached tensors
    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)

    # Metadata
    info = get_mla_metadata_info_v1(
        bs, qseq, nq, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=False,
        num_kv_splits=num_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,
        nq // nkv, nkv, True,
        wm, wis, wi, ri, rfm, rpm,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=qseq,
        uni_seqlen_qo=qseq,
        fast_mode=False,
        max_split_per_batch=num_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    cached = {
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "meta": {
            "work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
            "reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm,
        },
    }
    _cache[key] = cached
    return cached


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"]
    qseq = config["q_seq_len"]
    kvseq = config["kv_seq_len"]

    # FP8 quantize Q (~31µs via aiter.per_tensor_quant)
    q_fp8, q_scale = _quantize_fp8(q)

    # FP8 KV from input
    kv_fp8, kv_scale = kv_data["fp8"]

    # Select num_kv_splits
    total_kv = bs * kvseq
    if total_kv >= 128 * 1024:
        num_splits = 32
    elif total_kv >= 16 * 1024:
        num_splits = 16
    else:
        num_splits = 8

    # Avoid GPU→CPU sync: compute total_kv_len directly (uniform kv lengths)
    total_kv_len = bs * kvseq

    # Get all cached buffers (kv_indices, kv_last_page_len, metadata)
    c = _get_or_build_cache(
        bs, qseq, nq, nkv, total_kv_len, num_splits,
        q_fp8.dtype, kv_fp8.dtype, qo_indptr, kv_indptr,
    )

    kv_buffer_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])

    # Reuse cached output tensor
    o_key = (bs, nq, V_HEAD_DIM)
    if o_key not in _cache or _cache[o_key].shape[0] != q.shape[0]:
        _cache[o_key] = torch.empty((q.shape[0], nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    o = _cache[o_key]

    mla_decode_fwd(
        q_fp8.view(-1, nq, QK_HEAD_DIM),
        kv_buffer_4d,
        o,
        qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
        qseq,
        page_size=PAGE_SIZE,
        nhead_kv=nkv,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=num_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **c["meta"],
    )

    return o
scrolls · 144 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