Skip to content
KernelIndex
Search⌘K

submission 586726

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

persistent-kernel"""v707 - bf16 for small shapes, 1-split for (64,1K)+(256,1K), persistent for 8K."""

Kernel source

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

"""v707 - bf16 for small shapes, 1-split for (64,1K)+(256,1K), persistent for 8K."""

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 mla as aiter_mla
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

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 _build_nonpersist_1split(dev, bs, kvlen, kv_indptr):
    """Non-persistent mode with 1 split: output written directly, no reduce needed."""
    total_kv = bs * kvlen
    _, num_splits_indptr = aiter_mla.get_meta_param(1, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
    kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
    out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
    # With 1 split, logits shares memory with out
    logits = out.view(bs, 1, NUM_HEADS, V_HEAD_DIM)
    attn_lse = torch.empty((bs, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
    q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)
    return (num_splits_indptr, kv_indices, kv_lpl, out, logits, attn_lse, q_fp8)


def _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):
    """Persistent mode with metadata — uses mla_reduce_v1."""
    total_kv = bs * kvlen
    num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
    kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
    out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)

    info = get_mla_metadata_info_v1(
        bs, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,
    )
    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, False,
        wmd, wis, wi, ri, rfm, rpm,
        page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,
        dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
    )
    pt = int(rpm.numel())
    po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
    pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
    q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)
    return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl, q_fp8)


def _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):
    """Persistent mode with bf16 Q + bf16 KV — no quantization needed."""
    total_kv = bs * kvlen
    num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, torch.bfloat16)
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
    kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
    out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)

    info = get_mla_metadata_info_v1(
        bs, 1, NUM_HEADS, torch.bfloat16, torch.bfloat16,
        is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,
    )
    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, False,
        wmd, wis, wi, ri, rfm, rpm,
        page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,
        dtype_q=torch.bfloat16, dtype_kv=torch.bfloat16,
    )
    pt = int(rpm.numel())
    po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
    pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
    return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl)


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"])
    dev = q.device

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

    key = (dev.index, bs, kvlen)

    # --- bf16 path for small batches: skip Q quantization entirely ---
    if bs <= 32 and kvlen == 1024:
        bkey = ("bf16", *key)
        if bkey not in _cache:
            _cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 16, True)
        c = _cache[bkey]
        kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)
        aiter.mla_decode_stage1_asm_fwd(
            q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,
            c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,
        )
        aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
        return c[2]

    if bs == 4 and kvlen == 8192:
        bkey = ("bf16", *key)
        if bkey not in _cache:
            _cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 64, True)
        c = _cache[bkey]
        kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)
        aiter.mla_decode_stage1_asm_fwd(
            q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,
            c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,
        )
        aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
        return c[2]

    # --- 1-split non-persistent for select shapes: skip reduce ---
    if (bs, kvlen) in ((64, 1024), (256, 1024)):
        nkey = ("np1", *key)
        if nkey not in _cache:
            _cache[nkey] = _build_nonpersist_1split(dev, bs, kvlen, kv_indptr)
        c = _cache[nkey]
        q_scale = _quantize_q_fp8(q, c[6])
        aiter.mla_decode_stage1_asm_fwd(
            c[6], kv_buf, qo_indptr, kv_indptr, c[1], c[2], c[0],
            None, None, None, 1, 1, 1, sm_scale, c[4], c[5], c[3], q_scale, kv_scale,
        )
        return c[3]

    # --- Persistent mode for remaining shapes ---
    pkey = ("persist", *key)
    if pkey not in _cache:
        kv_gran = 64 if kvlen == 8192 else 8
        intra = True
        _cache[pkey] = _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra)
    c = _cache[pkey]
    q_scale = _quantize_q_fp8(q, c[11])
    aiter.mla_decode_stage1_asm_fwd(
        c[11], kv_buf, qo_indptr, kv_indptr, c[0], c[1], None,
        c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], q_scale, kv_scale,
    )
    aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
    return c[2]
scrolls · 184 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 586016.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """v697 - Direct stage1+reduce with PAGE_SIZE=1, pre-allocated split buffers."""
+ """v707 - bf16 for small shapes, 1-split for (64,1K)+(256,1K), persistent for 8K."""
import os
⋯ 5 unchanged lines
import aiter
from aiter import dtypes as aiter_dtypes
+ from aiter import mla as aiter_mla
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
try:
⋯ 7 unchanged lines
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)
⋯ 13 unchanged lines
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
-
+ def _build_nonpersist_1split(dev, bs, kvlen, kv_indptr):
+ """Non-persistent mode with 1 split: output written directly, no reduce needed."""
total_kv = bs * kvlen
+ _, num_splits_indptr = aiter_mla.get_meta_param(1, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
- kv_lpl = torch.ones(bs, dtype=torch.int32, device=dev)
+ kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
+ # With 1 split, logits shares memory with out
+ logits = out.view(bs, 1, NUM_HEADS, V_HEAD_DIM)
+ attn_lse = torch.empty((bs, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)
+ return (num_splits_indptr, kv_indices, kv_lpl, out, logits, attn_lse, q_fp8)
- # Metadata matching reference: PAGE_SIZE=1, is_causal=True, fast_mode=False
+
+ def _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):
+ """Persistent mode with metadata — uses mla_reduce_v1."""
+ total_kv = bs * kvlen
+ num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)
+ kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
+ kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
+ out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
+
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,
+ is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,
)
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,
+ qo_indptr, kv_indptr, kv_lpl, 16, 1, False,
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,
+ page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,
+ fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,
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)
+ po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
+ pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
+ q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)
+ return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl, q_fp8)
- c = (kv_indices, kv_lpl, out, q_fp8, wmd, wi, wis, ri, rfm, rpm, split_out, split_lse)
- _cache[key] = c
- return c
+ def _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):
+ """Persistent mode with bf16 Q + bf16 KV — no quantization needed."""
+ total_kv = bs * kvlen
+ num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, torch.bfloat16)
+ kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
+ kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
+ out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
+ info = get_mla_metadata_info_v1(
+ bs, 1, NUM_HEADS, torch.bfloat16, torch.bfloat16,
+ is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,
+ )
+ 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, False,
+ wmd, wis, wi, ri, rfm, rpm,
+ page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,
+ fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,
+ dtype_q=torch.bfloat16, dtype_kv=torch.bfloat16,
+ )
+ pt = int(rpm.numel())
+ po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
+ pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
+ return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl)
+
+
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"])
+ dev = q.device
- 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)
+ key = (dev.index, bs, kvlen)
+
+ # --- bf16 path for small batches: skip Q quantization entirely ---
+ if bs <= 32 and kvlen == 1024:
+ bkey = ("bf16", *key)
+ if bkey not in _cache:
+ _cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 16, True)
+ c = _cache[bkey]
+ kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)
+ aiter.mla_decode_stage1_asm_fwd(
+ q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,
+ c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,
+ )
+ aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
+ return c[2]
+
+ if bs == 4 and kvlen == 8192:
+ bkey = ("bf16", *key)
+ if bkey not in _cache:
+ _cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 64, True)
+ c = _cache[bkey]
+ kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)
+ aiter.mla_decode_stage1_asm_fwd(
+ q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,
+ c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,
+ )
+ aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
+ return c[2]
+
+ # --- 1-split non-persistent for select shapes: skip reduce ---
+ if (bs, kvlen) in ((64, 1024), (256, 1024)):
+ nkey = ("np1", *key)
+ if nkey not in _cache:
+ _cache[nkey] = _build_nonpersist_1split(dev, bs, kvlen, kv_indptr)
+ c = _cache[nkey]
+ q_scale = _quantize_q_fp8(q, c[6])
+ aiter.mla_decode_stage1_asm_fwd(
+ c[6], kv_buf, qo_indptr, kv_indptr, c[1], c[2], c[0],
+ None, None, None, 1, 1, 1, sm_scale, c[4], c[5], c[3], q_scale, kv_scale,
+ )
+ return c[3]
+
+ # --- Persistent mode for remaining shapes ---
+ pkey = ("persist", *key)
+ if pkey not in _cache:
+ kv_gran = 64 if kvlen == 8192 else 8
+ intra = True
+ _cache[pkey] = _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra)
+ c = _cache[pkey]
+ q_scale = _quantize_q_fp8(q, c[11])
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,
+ c[11], kv_buf, qo_indptr, kv_indptr, c[0], c[1], None,
+ c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], q_scale, kv_scale,
)
- aiter.mla_reduce_v1(split_out, split_lse, ri, rfm, rpm, 1, out, None)
- return out
+ aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
+ return c[2]
scrolls · 194 diff lines total

Best evidence level for this revision: reported

JSON