Skip to content
KernelIndex
Search⌘K

submission 716406

janice.jiayao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:728eeb60d43b97b9d7c8fee625e47523fb97ff17db5354ae8d1db17a30c68c94
license declaredunknown
license concludedunknown
authorsjanice.jiayao
imported2026-08-26

Kernel source

submission.py154 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
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

FP8_DTYPE = aiter_dtypes.fp8
_finfo = torch.finfo(FP8_DTYPE)
FP8_MAX = _finfo.max
FP8_MIN = _finfo.min
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1

_cache = {}


def get_config(batch_size: int, kv_seq_len: int):
    """
    Dispatch table fitted from MI355X probe data.

    Best observed steady-state configs:
      bs=4,   kv=1024 -> a16w8 split=32  fast_mode=True  intra_batch=False
      bs=4,   kv=8192 -> a16w8 split=256 fast_mode=False intra_batch=True
      bs=32,  kv=1024 -> a16w8 split=256 fast_mode=True  intra_batch=False
      bs=32,  kv=8192 -> a8w8  split=256 fast_mode=False intra_batch=True
      bs=64,  kv=1024 -> a16w8 split=32  fast_mode=False intra_batch=True
      bs=64,  kv=8192 -> a8w8  split=32  fast_mode=False intra_batch=True
      bs=256, kv=1024 -> a8w8  split=32  fast_mode=False intra_batch=True
      bs=256, kv=8192 -> a8w8  split=256 fast_mode=False intra_batch=True
    """
    if kv_seq_len <= 1024:
        if batch_size <= 4:
            return False, 32, True, False
        if batch_size <= 32:
            return False, 256, True, False
        if batch_size <= 64:
            return False, 32, False, True
        return True, 32, False, True

    if batch_size <= 4:
        return False, 256, False, True
    if batch_size <= 32:
        return True, 256, False, True
    if batch_size <= 64:
        return True, 32, False, True
    return True, 256, False, True


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

    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    total_kv_len = batch_size * kv_seq_len

    use_a8w8, num_kv_splits, fast_mode, intra_batch_mode = get_config(batch_size, kv_seq_len)
    cache_key = (batch_size, kv_seq_len, use_a8w8, num_kv_splits, fast_mode, intra_batch_mode)

    kv_buffer_fp8, kv_scale = kv_data["fp8"]

    if cache_key not in _cache:
        q_dtype = FP8_DTYPE if use_a8w8 else torch.bfloat16
        kv_dtype = FP8_DTYPE

        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
        kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")

        info = get_mla_metadata_info_v1(
            batch_size, 1, 16, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=fast_mode,
            num_kv_splits=num_kv_splits, intra_batch_mode=intra_batch_mode,
        )
        work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        (work_metadata, work_indptr, work_info_set,
         reduce_indptr, reduce_final_map, reduce_partial_map) = work

        get_mla_metadata_v1(
            qo_indptr, kv_indptr, kv_last_page_len,
            16, 1, True, work_metadata, work_info_set, work_indptr,
            reduce_indptr, reduce_final_map, reduce_partial_map,
            page_size=1, kv_granularity=16,
            max_seqlen_qo=1, uni_seqlen_qo=1,
            fast_mode=fast_mode, topk=-1, max_split_per_batch=num_kv_splits,
            intra_batch_mode=intra_batch_mode, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )

        logits = torch.empty((reduce_partial_map.size(0), 1, 16, 512), dtype=torch.float32, device="cuda")
        attn_lse = torch.empty((reduce_partial_map.size(0), 1, 16, 1), dtype=torch.float32, device="cuda")
        o = torch.empty((batch_size, 16, 512), dtype=torch.bfloat16, device="cuda")

        entry = {
            "use_a8w8": use_a8w8,
            "kv_indices": kv_indices,
            "kv_last_page_len": kv_last_page_len,
            "work_metadata": work_metadata,
            "work_indptr": work_indptr,
            "work_info_set": work_info_set,
            "reduce_indptr": reduce_indptr,
            "reduce_final_map": reduce_final_map,
            "reduce_partial_map": reduce_partial_map,
            "logits": logits,
            "attn_lse": attn_lse,
            "o": o,
        }

        if use_a8w8:
            entry["q_fp8"] = torch.empty((batch_size, 16, 576), dtype=FP8_DTYPE, device="cuda")
            entry["q_scratch"] = torch.empty((batch_size, 16, 576), dtype=torch.bfloat16, device="cuda")
            entry["q_amax"] = torch.empty((), dtype=q.dtype, device="cuda")
            entry["q_scale"] = torch.empty((1,), dtype=torch.float32, device="cuda")

        _cache[cache_key] = entry

    entry = _cache[cache_key]

    if entry["use_a8w8"]:
        q_scratch = entry["q_scratch"]
        q_amax = entry["q_amax"]
        q_scale = entry["q_scale"]
        q_fp8 = entry["q_fp8"]

        torch.abs(q, out=q_scratch)
        torch.amax(q_scratch, out=q_amax)
        q_amax.clamp_(min=1e-12)
        q_scale.copy_(q_amax)
        q_scale.div_(FP8_MAX)
        torch.div(q, q_scale, out=q_scratch)
        q_scratch.clamp_(min=FP8_MIN, max=FP8_MAX)
        q_fp8.copy_(q_scratch)

        aiter.mla_decode_stage1_asm_fwd(
            q_fp8.view(-1, 16, 576), kv_buffer_fp8.view(-1, 1, 1, 576),
            qo_indptr, kv_indptr, entry["kv_indices"], entry["kv_last_page_len"], None,
            entry["work_metadata"], entry["work_indptr"], entry["work_info_set"],
            1, PAGE_SIZE, 1, SM_SCALE, entry["logits"], entry["attn_lse"], entry["o"], q_scale, kv_scale,
        )
    else:
        aiter.mla_decode_stage1_asm_fwd(
            q.view(-1, 16, 576), kv_buffer_fp8.view(-1, 1, 1, 576),
            qo_indptr, kv_indptr, entry["kv_indices"], entry["kv_last_page_len"], None,
            entry["work_metadata"], entry["work_indptr"], entry["work_info_set"],
            1, PAGE_SIZE, 1, SM_SCALE, entry["logits"], entry["attn_lse"], entry["o"], None, kv_scale,
        )

    aiter.mla_reduce_v1(
        entry["logits"], entry["attn_lse"],
        entry["reduce_indptr"], entry["reduce_final_map"], entry["reduce_partial_map"],
        1, entry["o"], None,
    )

    return entry["o"]
scrolls · 154 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