Skip to content
KernelIndex
Search⌘K

submission 674864

yanchaomei · 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.

submission_mla_bf16.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-674864?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
73.7µs
#355 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:38566e524a90171c76a9a2e6be241337ca45be283d9eeaae1129f1735a90704a
license declaredunknown
license concludedunknown
authorsyanchaomei
imported2026-08-26

Kernel source

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

"""
MLA Decode — BF16 path (no FP8 quant overhead).
Uses mla_dec_stage1_bf16_a16w16_subQ16_mqa16.co kernel.
Trade: 2x more KV bandwidth vs 5µs saved on Q quant.
For small batches, quant overhead dominates → BF16 wins.
"""

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,
    mla_decode_stage1_asm_fwd, mla_reduce_v1,
    per_tensor_quant_hip,
)
from aiter.mla import mla_decode_fwd

FP8 = aiter_dtypes.fp8
BF16 = torch.bfloat16
NH, NKV, QKD, VD = 16, 1, 576, 512
SM = 1.0 / (QKD ** 0.5)
CU = 256
_c = {}


def _splits_fp8(bs, kl):
    oh = 84.1
    best_s, best = 1, float("-inf")
    for s in range(1, 17):
        w = (bs * s + CU - 1) // CU
        score = bs * s / (w * CU) * kl / (kl + oh * s)
        if score > best:
            best, best_s = score, s
    return min(best_s, max(1, (kl + 127) // 128))


def _splits_bf16(bs, kl):
    oh = 42.0  # BF16 has lower overhead per split (bigger tiles)
    best_s, best = 1, float("-inf")
    for s in range(1, 17):
        w = (bs * s + CU - 1) // CU
        score = bs * s / (w * CU) * kl / (kl + oh * s)
        if score > best:
            best, best_s = score, s
    return min(best_s, max(1, (kl + 127) // 128))


def _init_fp8(bs, qs, kl, qoi, kvi):
    ns = _splits_fp8(bs, kl)
    ki = torch.arange(bs*kl, dtype=torch.int32, device="cuda")
    klp = (kvi[1:]-kvi[:-1]).to(torch.int32)
    info = get_mla_metadata_info_v1(bs, qs, NH, FP8, FP8,
        is_sparse=False, fast_mode=False, num_kv_splits=ns, intra_batch_mode=True)
    w = [torch.empty(s, dtype=t, device="cuda") for s,t in info]
    wm,wi,wis,ri,rfm,rpm = w
    get_mla_metadata_v1(qoi,kvi,klp, NH,NKV,True, wm,wis,wi,ri,rfm,rpm,
        page_size=1, kv_granularity=16, max_seqlen_qo=qs, uni_seqlen_qo=qs,
        fast_mode=False, max_split_per_batch=ns, intra_batch_mode=True,
        dtype_q=FP8, dtype_kv=FP8)
    np_ = rpm.numel()
    return {"ns":ns, "wm":wm,"wi":wi,"wis":wis,"ri":ri,"rfm":rfm,"rpm":rpm,
            "ki":ki,"klp":klp,
            "sd":torch.empty((np_*qs,1,NH,VD),dtype=torch.float32,device="cuda"),
            "sl":torch.empty((np_*qs,1,NH,1),dtype=torch.float32,device="cuda")}


def _init_bf16(bs, qs, kl, qoi, kvi):
    ns = _splits_bf16(bs, kl)
    ki = torch.arange(bs*kl, dtype=torch.int32, device="cuda")
    klp = (kvi[1:]-kvi[:-1]).to(torch.int32)
    # BF16 dtype for metadata
    info = get_mla_metadata_info_v1(bs, qs, NH, BF16, BF16,
        is_sparse=False, fast_mode=False, num_kv_splits=ns, intra_batch_mode=True)
    w = [torch.empty(s, dtype=t, device="cuda") for s,t in info]
    wm,wi,wis,ri,rfm,rpm = w
    get_mla_metadata_v1(qoi,kvi,klp, NH,NKV,True, wm,wis,wi,ri,rfm,rpm,
        page_size=1, kv_granularity=16, max_seqlen_qo=qs, uni_seqlen_qo=qs,
        fast_mode=False, max_split_per_batch=ns, intra_batch_mode=True,
        dtype_q=BF16, dtype_kv=BF16)
    np_ = rpm.numel()
    return {"ns":ns, "wm":wm,"wi":wi,"wis":wis,"ri":ri,"rfm":rfm,"rpm":rpm,
            "ki":ki,"klp":klp,
            "sd":torch.empty((np_*qs,1,NH,VD),dtype=torch.float32,device="cuda"),
            "sl":torch.empty((np_*qs,1,NH,1),dtype=torch.float32,device="cuda")}


# Use BF16 for small batches (quant overhead dominates), FP8 for large
def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs, kl = config["batch_size"], config["kv_seq_len"]
    qs, qt = config.get("q_seq_len",1), q.shape[0]

    # BF16 wins for kv=1024 with bs=32/64 (saves Q quant overhead)
    use_bf16 = (kl <= 1024 and 16 <= bs <= 64)

    if use_bf16:
        key = ("bf16", bs, kl)
        if key not in _c:
            _c[key] = _init_bf16(bs, qs, kl, qo_indptr, kv_indptr)
        c = _c[key]
        kv_bf16 = kv_data["bf16"]
        kv4 = kv_bf16.view(kv_bf16.shape[0], 1, NKV, kv_bf16.shape[-1])
        q_bf16 = q.view(-1, NH, QKD)
        out = torch.empty((qt, NH, VD), dtype=torch.bfloat16, device="cuda")

        mla_decode_stage1_asm_fwd(
            q_bf16, kv4, qo_indptr, kv_indptr,
            c["ki"], c["klp"], None,
            c["wm"], c["wi"], c["wis"],
            qs, 1, NKV, SM,
            c["sd"], c["sl"], out,
            None, None,  # No scales for BF16
        )
        mla_reduce_v1(c["sd"], c["sl"], c["ri"], c["rfm"], c["rpm"], qs, out, None)
        return out
    else:
        key = ("fp8", bs, kl)
        if key not in _c:
            _c[key] = _init_fp8(bs, qs, kl, qo_indptr, kv_indptr)
        c = _c[key]
        kf, ks = kv_data["fp8"]
        qf, qsc = per_tensor_quant_hip(q.view(-1,NH,QKD), quant_dtype=FP8)
        qsc = qsc.reshape(1)
        kv4 = kf.view(kf.shape[0],1,NKV,kf.shape[-1])
        out = torch.empty((qt, NH, VD), dtype=torch.bfloat16, device="cuda")

        mla_decode_stage1_asm_fwd(
            qf.view(-1,NH,QKD), kv4, qo_indptr, kv_indptr,
            c["ki"], c["klp"], None,
            c["wm"], c["wi"], c["wis"],
            qs, 1, NKV, SM,
            c["sd"], c["sl"], out, qsc, ks,
        )
        mla_reduce_v1(c["sd"], c["sl"], c["ri"], c["rfm"], c["rpm"], qs, out, None)
        return out
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