Skip to content
KernelIndex
Search⌘K

submission 747474

rajeev9 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747474?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
75.0µs
#368 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3ccb707fec89f36a4087341689f3600e076ce14ad948b0666740462fb5437748
license declaredunknown
license concludedunknown
authorsrajeev9
imported2026-08-26

Techniques

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

persistent-kernelMLA Decode: All-persistent FP8 ASM with per-config kv_granularity tuning.

Kernel source

submission.py117 lines
"""
MLA Decode: All-persistent FP8 ASM with per-config kv_granularity tuning.
g=64 for bs<=4/kv<=1024, g=128 for bs<=4/kv>1024, g=32 for bs<=32, g=16 for bs>32.
"""
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.jit.utils.chip_info import get_cu_num

SM_SCALE: float = 1.0 / (576 ** 0.5)
FP8 = dtypes.fp8

_QSC = torch.empty(1, dtype=torch.float32, device="cuda")
_CACHE = {}

_quant = aiter.dynamic_per_tensor_quant
_stage1 = aiter.mla_decode_stage1_asm_fwd
_reduce = aiter.mla_reduce_v1
_get_meta = aiter.get_mla_metadata_v1

def _get_kv_granularity(bs, kv_seq):
    if bs <= 4 and kv_seq <= 1024:
        return 64
    if bs <= 4:
        return 128
    if bs <= 32:
        return 32
    return 16


def _build_cache(bs, kv_seq, nq, nkv, dv, dq, qsl, qo_indptr, kv_indptr):
    cu_num = get_cu_num()
    tkv = bs * kv_seq

    tile_cnt = bs
    max_work = bs + cu_num - 1
    max_split_tiles = min(max_work, (cu_num - 1) * 2)

    device = "cuda"
    work_meta = torch.empty(2, dtype=torch.uint64, device=device)
    work_indptr = torch.empty(cu_num + 1, dtype=torch.int32, device=device)
    work_info_set = torch.empty((max_work, 8), dtype=torch.int32, device=device)
    reduce_indptr = torch.empty(tile_cnt + 1, dtype=torch.int32, device=device)
    reduce_final_map = torch.empty((tile_cnt, 2), dtype=torch.int32, device=device)
    reduce_partial_map = torch.empty(max_split_tiles, dtype=torch.int32, device=device)

    klp = torch.full((bs,), kv_seq, dtype=torch.int32, device=device)
    kv_indices = torch.arange(tkv, dtype=torch.int32, device=device)

    kvg = _get_kv_granularity(bs, kv_seq)
    _get_meta(
        qo_indptr, kv_indptr, klp,
        nq // nkv, nkv, False,
        work_meta, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        1, kvg, qsl, qsl, True, -1, -1, False, FP8, FP8,
    )

    total_q = bs * qsl
    logits = torch.empty((max_split_tiles * qsl, 1, nq, dv), dtype=torch.float32, device=device)
    attn_lse = torch.empty((max_split_tiles * qsl, 1, nq, 1), dtype=torch.float32, device=device)
    o = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device=device)
    qf = torch.empty((total_q, nq * dq), dtype=FP8, device=device)
    qf_3d = qf.view(-1, nq, dq)

    return (
        work_meta, work_indptr, work_info_set,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        klp, kv_indices,
        logits, attn_lse, o, qf, qf_3d,
        qsl, nkv,
    )


@torch.inference_mode()
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"]
    dv = config["v_head_dim"]
    dq = config["qk_head_dim"]
    qsl = config["q_seq_len"]
    kv_seq = config["kv_seq_len"]

    key = (bs, kv_seq)
    if key not in _CACHE:
        _CACHE[key] = _build_cache(bs, kv_seq, nq, nkv, dv, dq, qsl, qo_indptr, kv_indptr)

    (work_meta, work_indptr, work_info_set,
     reduce_indptr, reduce_final_map, reduce_partial_map,
     klp, kv_indices,
     logits, attn_lse, o, qf, qf_3d,
     qsl_c, nkv_c) = _CACHE[key]

    kv_fp8, kv_sc = kv_data["fp8"]

    _quant(qf, q, _QSC)

    _stage1(
        qf_3d, kv_fp8.view(-1, 1, nkv_c, dq),
        qo_indptr, kv_indptr, kv_indices, klp,
        None, work_meta, work_indptr, work_info_set,
        qsl_c, 1, nkv_c, SM_SCALE,
        logits, attn_lse, o, _QSC, kv_sc,
    )

    _reduce(
        logits, attn_lse,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        qsl_c, o, None,
    )

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