Skip to content
KernelIndex
Search⌘K

submission 746532

kkosey · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:809672bf4e72c30ec3266c2092da1cda3e2c08b18bf911e6b2823b3df5408744
license declaredunknown
license concludedunknown
authorskkosey
imported2026-08-15

Kernel source

submission_e99_noibm_ps8.py169 lines
#!POPCORN gpu MI355X
#!POPCORN leaderboard amd-mixed-mla
"""E93: E88 optimal dispatch + hybrid PAGE_SIZE (PS=4 for kv>1024, PS=2 for kv≤1024).

Combines:
- E88: a16w8 for all except bs≥256 kv>1024 (a8w8), IBM=True for kv>1024 bs≥32
- E89: PAGE_SIZE=4 gives -16% to -32% speedup for kv=8192 shapes

Expected: best of both worlds.
"""

import torch
import aiter
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

FP8_DTYPE = aiter_dtypes.fp8

NUM_KV_HEADS = 1
SM_SCALE = 1.0 / (576 ** 0.5)
QK_HEAD_DIM = 576
V_HEAD_DIM  = 512

_cache = {}

_stage1_op    = torch.ops.aiter.mla_decode_stage1_asm_fwd
_reduce_op    = torch.ops.aiter.mla_reduce_v1
_dpt_quant_op = torch.ops.aiter.dynamic_per_tensor_quant

_MIN_PROGRAMS = 120
_TARGET_GRID  = 512


def _choose_strategy(batch_size, q_seq_len, kv_seq_len):
    total_q = batch_size * q_seq_len
    # bs≤4: always a16w8 NS=1
    if batch_size <= 4:
        return True, 1
    # E98: ALL a16w8 (no a8w8 branch)
    ns = max(1, (_TARGET_GRID + total_q - 1) // total_q)
    return True, min(16, ns)


def _init_shape(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr):
    nkv = NUM_KV_HEADS
    use_a16w8, num_kv_splits = _choose_strategy(batch_size, q_seq_len, kv_seq_len)

    q_dtype  = torch.bfloat16 if use_a16w8 else FP8_DTYPE
    kv_dtype = FP8_DTYPE

    # Hybrid PAGE_SIZE: PS=4 for kv>1024, PS=2 for kv≤1024
    page_size = 8 if kv_seq_len > 1024 else 2

    # IBM: kv>1024 AND bs>=32 (E88 optimal)
    ibm = False

    kv_indptr_pages = (kv_indptr // page_size).to(torch.int32)
    kv_last_page_len = (kv_indptr_pages[1:] - kv_indptr_pages[:-1]).to(torch.int32)

    info = get_mla_metadata_info_v1(
        batch_size, q_seq_len, num_heads, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=False,
        num_kv_splits=num_kv_splits, intra_batch_mode=ibm,
    )
    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_pages, kv_last_page_len,
        num_heads // nkv, nkv, True,
        work_metadata, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=page_size,
        kv_granularity=max(page_size, 16),
        max_seqlen_qo=q_seq_len,
        uni_seqlen_qo=q_seq_len,
        fast_mode=False,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=ibm,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    total_q    = batch_size * q_seq_len
    total_kv   = batch_size * kv_seq_len
    total_pages = total_kv // page_size
    partial_size = reduce_partial_map.size(0) * q_seq_len

    result = {
        "use_a16w8":          use_a16w8,
        "page_size":          page_size,
        "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,
        "kv_last_page_len":   kv_last_page_len,
        "kv_indptr_pages":    kv_indptr_pages,
        "kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
        "logits":   torch.empty((partial_size, 1, num_heads, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
        "attn_lse": torch.empty((partial_size, 1, num_heads, 1),          dtype=torch.float32, device="cuda"),
        "output":   torch.empty((total_q, num_heads, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
    }
    if not use_a16w8:
        result["q_fp8"]   = torch.empty((total_q, num_heads, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
        result["q_scale"] = torch.empty(1, dtype=torch.float32, device="cuda")
    return result


def _get_cache(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr):
    key = (batch_size, q_seq_len, num_heads, kv_seq_len)
    if key not in _cache:
        _cache[key] = _init_shape(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr)
    return _cache[key]


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

    batch_size = config["batch_size"]
    num_heads  = config["num_heads"]
    q_seq_len  = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]

    c = _get_cache(batch_size, q_seq_len, num_heads, kv_seq_len, qo_indptr, kv_indptr)

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    total_kv = batch_size * kv_seq_len
    page_size = c["page_size"]
    kv_buffer_4d = kv_buffer_fp8.view(
        total_kv // page_size, page_size, NUM_KV_HEADS, kv_buffer_fp8.shape[-1]
    )

    o = c["output"]
    kv_indptr_pages = c["kv_indptr_pages"]

    if c["use_a16w8"]:
        _stage1_op(
            q, kv_buffer_4d, qo_indptr, kv_indptr_pages,
            c["kv_indices"], c["kv_last_page_len"], None,
            c["work_metadata"], c["work_indptr"], c["work_info_set"],
            q_seq_len, page_size, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], o,
            None, kv_scale_fp8,
        )
    else:
        q_fp8   = c["q_fp8"]
        q_scale = c["q_scale"]
        _dpt_quant_op(q_fp8, q, q_scale)
        _stage1_op(
            q_fp8, kv_buffer_4d, qo_indptr, kv_indptr_pages,
            c["kv_indices"], c["kv_last_page_len"], None,
            c["work_metadata"], c["work_indptr"], c["work_info_set"],
            q_seq_len, page_size, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], o,
            q_scale, kv_scale_fp8,
        )

    _reduce_op(
        c["logits"], c["attn_lse"],
        c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
        q_seq_len, o, None,
    )

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