Skip to content
KernelIndex
Search⌘K

submission 714956

xiehuanyi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-714956?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
171.9µs
#548 of 766
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:355469c87a5f23471cf7f0ab1cc995d440dfc05a130da6246580736391b2cfc4
license declaredunknown
license concludedunknown
authorsxiehuanyi
imported2026-08-26

Kernel source

submission_v7.py95 lines
"""
MLA Decode v7: Lean version - only cache Q quantization, no KV caching.
v6 timed out due to memory pressure from KV caching. Keep it simple.
"""
import torch
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

SM_SCALE = 1.0 / (576 ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8

_cached_q = None
_cached_q_data = None
_meta_bufs = {}


def custom_kernel(data: input_t) -> output_t:
    global _cached_q, _cached_q_data
    q, kv_data, qo_indptr, kv_indptr, config = data

    bs = config["batch_size"]
    nq, nkv = 16, 1
    dq, dv = 576, 512
    qsl = config["q_seq_len"]
    kvsl = config["kv_seq_len"]

    # Cache Q fp8 quantization
    if _cached_q is q:
        q_fp8, q_sc = _cached_q_data
    else:
        finfo = torch.finfo(FP8_DTYPE)
        amax = q.abs().amax().clamp(min=1e-12)
        sc = amax / finfo.max
        q_fp8 = (q / sc).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
        q_sc = sc.to(torch.float32).reshape(1)
        _cached_q = q
        _cached_q_data = (q_fp8, q_sc)

    kv_fp8, kv_sc = kv_data["fp8"]
    kv4d = kv_fp8.view(kv_fp8.shape[0], 1, 1, kv_fp8.shape[-1])

    # Adaptive splits
    if bs <= 4:
        nks, kvg = (64, 64) if kvsl >= 8192 else (32, 32)
    elif bs <= 32:
        nks, kvg = 32, 32
    elif bs <= 64:
        nks, kvg = (32, 32) if kvsl >= 8192 else (16, 16)
    else:
        nks, kvg = (16, 16) if kvsl >= 8192 else (8, 16)

    tkv = int(kv_indptr[-1].item())
    kv_idx = torch.arange(tkv, dtype=torch.int32, device="cuda")
    kv_lpl = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    # Cache metadata buffers by shape key
    mk = (bs, kvsl, nks, kvg)
    if mk not in _meta_bufs:
        info = get_mla_metadata_info_v1(
            bs, qsl, nq, q_fp8.dtype, kv_fp8.dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=nks, intra_batch_mode=True,
        )
        _meta_bufs[mk] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]

    wm, wi, ws, ri, rfm, rpm = _meta_bufs[mk]

    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_lpl,
        nq, nkv, True,
        wm, ws, wi, ri, rfm, rpm,
        page_size=1, kv_granularity=kvg,
        max_seqlen_qo=qsl, uni_seqlen_qo=qsl,
        fast_mode=False, max_split_per_batch=nks,
        intra_batch_mode=True,
        dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype,
    )

    o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")
    mla_decode_fwd(
        q_fp8.view(-1, nq, dq), kv4d, o,
        qo_indptr, kv_indptr, kv_idx, kv_lpl, qsl,
        page_size=1, nhead_kv=nkv,
        sm_scale=SM_SCALE, logit_cap=0.0,
        num_kv_splits=nks,
        q_scale=q_sc, kv_scale=kv_sc,
        intra_batch_mode=True,
        work_meta_data=wm, work_indptr=wi, work_info_set=ws,
        reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
    )
    return o
scrolls · 95 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