Skip to content
KernelIndex
Search⌘K

submission 586016

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f769edf86c91fcb08b8159b9cdb163db6363088df97d9da609b83798f91a1432
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Kernel source

submission.py109 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""v697 - Direct stage1+reduce with PAGE_SIZE=1, pre-allocated split buffers."""

import os

os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")

import torch
from task import input_t, output_t

import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

try:
    from aiter.jit.module_quant import static_per_tensor_quant
except Exception:
    try:
        from aiter.ops.quant import static_per_tensor_quant
    except Exception:
        static_per_tensor_quant = None

NUM_HEADS = 16
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
NUM_KV_SPLITS = 32

FP8_DTYPE = aiter_dtypes.fp8
_FP8_FINFO = torch.finfo(FP8_DTYPE)

_cache = {}


def _quantize_q_fp8(q, q_fp8_buf):
    amax = q.abs().amax().clamp(min=1e-12)
    scale = (amax / _FP8_FINFO.max).reshape(1).to(torch.float32)
    if static_per_tensor_quant is not None:
        static_per_tensor_quant(q_fp8_buf, q, scale)
    else:
        q_fp8_buf.copy_(
            (q / scale).clamp(min=_FP8_FINFO.min, max=_FP8_FINFO.max).to(FP8_DTYPE)
        )
    return scale


def _get_cache(dev, qo_indptr, kv_indptr, bs, kvlen):
    key = (dev.index, bs, kvlen)
    c = _cache.get(key)
    if c is not None:
        return c

    total_kv = bs * kvlen
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
    kv_lpl = torch.ones(bs, dtype=torch.int32, device=dev)
    out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
    q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)

    # Metadata matching reference: PAGE_SIZE=1, is_causal=True, fast_mode=False
    info = get_mla_metadata_info_v1(
        bs, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=False,
        num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=True,
    )
    bufs = [torch.empty(s, dtype=t, device=dev) for s, t in info]
    wmd, wi, wis, ri, rfm, rpm = bufs
    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_lpl,
        16, 1, True,
        wmd, wis, wi, ri, rfm, rpm,
        page_size=1, kv_granularity=16,
        max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
        intra_batch_mode=True,
        dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
    )

    # Pre-allocate split buffers (size from reduce_partial_map)
    pt = int(rpm.numel())
    split_out = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
    split_lse = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)

    c = (kv_indices, kv_lpl, out, q_fp8, wmd, wi, wis, ri, rfm, rpm, split_out, split_lse)
    _cache[key] = c
    return c


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = int(config["batch_size"])
    kvlen = int(config["kv_seq_len"])
    sm_scale = float(config["sm_scale"])

    cache = _get_cache(q.device, qo_indptr, kv_indptr, bs, kvlen)
    kv_indices, kv_lpl, out, q_fp8, wmd, wi, wis, ri, rfm, rpm, split_out, split_lse = cache

    q_scale = _quantize_q_fp8(q, q_fp8)
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_buf = kv_fp8.view(-1, 1, 1, QK_HEAD_DIM)

    aiter.mla_decode_stage1_asm_fwd(
        q_fp8, kv_buf, qo_indptr, kv_indptr, kv_indices, kv_lpl, None,
        wmd, wi, wis, 1, 1, 1, sm_scale, split_out, split_lse, out, q_scale, kv_scale,
    )
    aiter.mla_reduce_v1(split_out, split_lse, ri, rfm, rpm, 1, out, None)
    return out
scrolls · 109 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