Skip to content
KernelIndex
Search⌘K

submission 658894

bhagawan-yantrion · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:44396cde58c0f067a69122375b9aadaef063571a3870854ebf6626e10bc9b237
license declaredunknown
license concludedunknown
authorsbhagawan-yantrion
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

split-kMLA decode — optimized aiter FP8 with per-shape split-K tuning and full caching.

Kernel source

submission.py82 lines
"""
MLA decode — optimized aiter FP8 with per-shape split-K tuning and full caching.
"""
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
import aiter

FP8 = aiter_dtypes.fp8
SM = 1.0 / (576 ** 0.5)
PAGE = 1
_c = {}


def _choose_splits(bs, kvsl):
    """Per-shape split-K tuning for MI355X (304 CUs).

    8 benchmark shapes:
      bs=4,   kv=1024  →  need many splits (4 WGs is nothing)
      bs=4,   kv=8192  →  many splits, long sequences
      bs=32,  kv=1024  →  moderate splits
      bs=32,  kv=8192  →  moderate splits, long sequences
      bs=64,  kv=1024  →  fewer splits (64 WGs already decent)
      bs=64,  kv=8192  →  some splits for long sequences
      bs=256, kv=1024  →  minimal splits (256 WGs fills CUs)
      bs=256, kv=8192  →  minimal splits

    Rule: target bs * nsplits ≈ 256-512 for CU utilization,
    but also need enough tokens per split (≥64) for MFMA efficiency.
    """
    # Hardcoded per benchmark shape for optimal performance
    if bs <= 4:
        return min(16, max(1, kvsl // 128))  # bs=4: 8 for kv=1024, 16 for kv=8192 (cap reduce overhead)
    if bs <= 32:
        return min(16, max(1, kvsl // 128))  # bs=32: 8 for kv=1024, 16 for kv=8192 (more parallelism)
    if bs <= 64:
        return max(1, min(8, kvsl // 256))   # bs=64: 4 for kv=1024, 8 for kv=8192 (balance)
    return max(1, min(4, kvsl // 512))       # bs=256: 1-2 for kv=1024, 4 for kv=8192


def _aiter_fp8(q, kv_data, qo, kvi, cfg):
    bs = cfg["batch_size"]; nq = cfg["num_heads"]; nkv = cfg["num_kv_heads"]
    dq = cfg["qk_head_dim"]; dv = cfg["v_head_dim"]; qsl = cfg["q_seq_len"]
    kvsl = cfg["kv_seq_len"]; tq = q.shape[0]
    nsplits = _choose_splits(bs, kvsl)
    key = ("fp8", bs, kvsl, nq, qsl, nsplits)
    if key not in _c:
        c = {}
        c["ki"] = torch.arange(bs * kvsl, dtype=torch.int32, device="cuda")
        c["klp"] = torch.full((bs,), kvsl, dtype=torch.int32, device="cuda")
        c["o"] = torch.empty((tq, nq, dv), dtype=torch.bfloat16, device="cuda")
        c["qf"] = torch.empty((tq * nq, dq), dtype=FP8, device="cuda")
        c["qs"] = torch.empty(1, dtype=torch.float32, device="cuda")
        info = get_mla_metadata_info_v1(bs, qsl, nq, FP8, FP8, is_sparse=False,
            fast_mode=False, num_kv_splits=nsplits, 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(qo, kvi, c["klp"], nq // nkv, nkv, True,
            wm, wis, wi, ri, rfm, rpm, page_size=PAGE, kv_granularity=16,
            max_seqlen_qo=qsl, uni_seqlen_qo=qsl, fast_mode=False,
            max_split_per_batch=nsplits, intra_batch_mode=True, dtype_q=FP8, dtype_kv=FP8)
        c["m"] = {"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
                  "reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm}
        c["nsplits"] = nsplits
        _c[key] = c
    c = _c[key]
    aiter.dynamic_per_tensor_quant(c["qf"], q.reshape(tq * nq, dq), c["qs"])
    kvf, kvs = kv_data["fp8"]
    kv4d = kvf.view(kvf.shape[0], PAGE, nkv, kvf.shape[-1])
    mla_decode_fwd(c["qf"].view(tq, nq, dq), kv4d, c["o"], qo, kvi, c["ki"], c["klp"],
        qsl, page_size=PAGE, nhead_kv=nkv, sm_scale=SM, logit_cap=0,
        num_kv_splits=c["nsplits"], q_scale=c["qs"], kv_scale=kvs,
        intra_batch_mode=True, **c["m"])
    return c["o"]


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo, kvi, cfg = data
    return _aiter_fp8(q, kv_data, qo, kvi, cfg)
scrolls · 82 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