Skip to content
KernelIndex
Search⌘K

submission 646717

zaiji100 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ee3c6255cbed9b3693f4a730b0db0c7c2fb83bf1c8eb1240db41ecedc57b4c2f
license declaredunknown
license concludedunknown
authorszaiji100
imported2026-08-26

Kernel source

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

import torch
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
import aiter; from aiter.mla import mla_decode_fwd as _warmup_fn
FP8 = aiter_dtypes.fp8; BF16 = torch.bfloat16; FP32 = torch.float32
_SM = float(1.0 / (576 ** 0.5)); _c = {}

def _build(bs, tq, tkv, qsl, qoi, kvi, dev, kvd, qd, fast, ns):
    klp = (kvi[1:] - kvi[:-1]).to(torch.int32)
    kidx = torch.arange(tkv, dtype=torch.int32, device=dev)
    info = get_mla_metadata_info_v1(bs, qsl, 16, qd, kvd,
        is_sparse=False, fast_mode=fast, num_kv_splits=ns, intra_batch_mode=True)
    bufs = [torch.empty(s, dtype=t, device=dev) for s, t in info]
    wm, wi, wis, ri, rfm, rpm = bufs
    get_mla_metadata_v1(qoi, kvi, klp, 16, 1, True, wm, wis, wi, ri, rfm, rpm,
        page_size=1, kv_granularity=16, max_seqlen_qo=qsl, uni_seqlen_qo=qsl,
        fast_mode=fast, max_split_per_batch=ns, intra_batch_mode=True,
        dtype_q=qd, dtype_kv=kvd)
    np_ = rpm.size(0)
    lg = torch.empty((np_ * qsl, 1, 16, 512), dtype=FP32, device=dev)
    ls = torch.empty((np_ * qsl, 1, 16, 1), dtype=FP32, device=dev)
    qi_buf = torch.empty((tq, 16, 576), dtype=FP8, device=dev) if qd == FP8 else None
    qs_buf = torch.empty(1, dtype=FP32, device=dev) if qd == FP8 else None
    return (klp, kidx, wm, wi, wis, ri, rfm, rpm, lg, ls, qi_buf, qs_buf)

def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]; tq = q.shape[0]; qsl = config["q_seq_len"]
    kvf, kvs = kv_data["fp8"]; tkv = kvf.shape[0]

    # Optimal strategy per config size:
    # bs≤4: a16w16 (bf16 KV) + fast_mode — fastest for tiny batch
    # bs≤32: a16w8 (fp8 KV) + fast_mode — fast metadata
    # bs>32, tkv<2M: a16w8 + normal mode
    # bs>32, tkv≥2M: a8w8 (fp8 Q) + normal mode
    if bs <= 4:
        kvd = BF16; qd = BF16; fast = True; ns = 32; use_bf16_kv = True
    elif bs <= 32:
        kvd = FP8; qd = BF16; fast = True; ns = 32; use_bf16_kv = False
    elif tkv < 2000000:
        kvd = FP8; qd = BF16; fast = False; ns = 32; use_bf16_kv = False
    else:
        kvd = FP8; qd = FP8; fast = False; ns = 16; use_bf16_kv = False

    k = (bs, tq, tkv, qsl, kvd, qd, fast)
    c = _c.get(k)
    if c is None: c = _build(bs, tq, tkv, qsl, qo_indptr, kv_indptr, q.device, kvd, qd, fast, ns); _c[k] = c
    klp, kidx, wm, wi, wis, ri, rfm, rpm, lg, ls, qi_buf, qs_buf = c

    o = torch.empty((tq, 16, 512), dtype=BF16, device=q.device)
    if use_bf16_kv:
        kv_tensor = kv_data['bf16'].view(tkv, 1, 1, -1)
        kv_scale = None
    else:
        kv_tensor = kvf.view(tkv, 1, 1, -1)
        kv_scale = kvs
    if qd == FP8:
        aiter.dynamic_per_tensor_quant(qi_buf, q, qs_buf)
        qi, qs = qi_buf, qs_buf
    else:
        qi, qs = q, None
    aiter.mla_decode_stage1_asm_fwd(qi, kv_tensor, qo_indptr, kv_indptr, kidx, klp,
        None, wm, wi, wis, qsl, 1, 1, _SM, lg, ls, o, qs, kv_scale)
    aiter.mla_reduce_v1(lg, ls, ri, rfm, rpm, qsl, o, None)
    return o
scrolls · 70 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