Skip to content
KernelIndex
Search⌘K

submission 742458

Ananda Sai A · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:837416ca81b82ac5a62a78607482e457e5d7ed7c80ca6607a5f86af74b945de5
license declaredunknown
license concludedunknown
authorsAnanda Sai A
imported2026-08-15

Techniques

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

tile-n = 128FP8_MIN_BLOCK_N = 128

Kernel source

submission_v224.py236 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
v224
"""

import math
import torch
from task import input_t, output_t

from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)
FP8_DTYPE = aiter_dtypes.fp8

STATIC_Q_ABSMAX = 6.0
FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
STATIC_Q_SCALE = STATIC_Q_ABSMAX / FP8_MAX

_QUANT_FN = None
try:
    from aiter.ops.quant import static_per_tensor_quant as _sqf
    _QUANT_FN = _sqf
except Exception:
    try:
        from aiter.jit.module_quant import static_per_tensor_quant as _sqf2
        _QUANT_FN = _sqf2
    except Exception:
        pass


def _quant_q(dst, src, scale):
    if _QUANT_FN is not None:
        _QUANT_FN(dst, src, scale)
    else:
        dst.copy_((src / scale).clamp(min=-FP8_MAX, max=FP8_MAX).to(FP8_DTYPE))


_cache = {}


def _build_persist(bs, kv, qtot, qo_ind, kv_ind, dq, dkv, pg, fast=True, n_splits=32):
    total = bs * kv
    if pg > 1:
        npg = total // pg
        idx = torch.arange(npg, dtype=torch.int32, device="cuda")
        ki = torch.arange(0, bs + 1, dtype=torch.int32, device="cuda") * (kv // pg)
        klp = torch.full((bs,), pg, dtype=torch.int32, device="cuda")
    else:
        idx = torch.arange(total, dtype=torch.int32, device="cuda")
        ki = kv_ind
        klp = (kv_ind[1:] - kv_ind[:-1]).to(torch.int32)

    info = get_mla_metadata_info_v1(
        bs,
        1,
        NUM_HEADS,
        dq,
        dkv,
        is_sparse=False,
        fast_mode=fast,
        num_kv_splits=n_splits,
        intra_batch_mode=True,
    )
    wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    get_mla_metadata_v1(
        qo_ind,
        ki,
        klp,
        NUM_HEADS // NUM_KV_HEADS,
        NUM_KV_HEADS,
        True,
        wk[0],
        wk[2],
        wk[1],
        wk[3],
        wk[4],
        wk[5],
        page_size=pg,
        kv_granularity=max((128 + pg - 1) // pg, pg, 16),
        max_seqlen_qo=1,
        uni_seqlen_qo=1,
        fast_mode=fast,
        max_split_per_batch=n_splits,
        intra_batch_mode=True,
        dtype_q=dq,
        dtype_kv=dkv,
    )
    meta = dict(
        work_meta_data=wk[0],
        work_indptr=wk[1],
        work_info_set=wk[2],
        reduce_indptr=wk[3],
        reduce_final_map=wk[4],
        reduce_partial_map=wk[5],
    )
    out = torch.empty((qtot, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    return meta, idx, klp, ki, pg, out


def _init_bf16_np(bs, kv, qtot, kv_ind):
    tag = ("bf16np", bs, kv)
    if tag in _cache:
        return _cache[tag]
    total = bs * kv
    c = (
        torch.arange(total, dtype=torch.int32, device="cuda"),
        (kv_ind[1:] - kv_ind[:-1]).to(torch.int32),
        torch.empty((qtot, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
    )
    _cache[tag] = c
    return c


def _init_fp8_persist(bs, kv, qtot, qo_ind, kv_ind, pg, fast=True, n_splits=32):
    tag = ("fp8ps", bs, kv, pg, fast, n_splits)
    if tag in _cache:
        return _cache[tag]
    c = _build_persist(bs, kv, qtot, qo_ind, kv_ind, FP8_DTYPE, FP8_DTYPE, pg, fast, n_splits)
    q_fp8 = torch.empty((qtot, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
    q_scale = torch.tensor([STATIC_Q_SCALE], dtype=torch.float32, device="cuda")
    _cache[tag] = (*c, q_fp8, q_scale)
    return _cache[tag]


_SHAPE_CFG = {
    (4, 1024): ("bf16np", 1, 0, True),
    (4, 8192): ("fp8ps", 8, 16, True),
    (32, 1024): ("fp8ps", 2, 8, True),
    (32, 8192): ("fp8ps", 8, 8, True),
    (64, 1024): ("fp8ps", 2, 8, True),
    (64, 8192): ("fp8ps", 8, 8, True),
    (256, 1024): ("fp8ps", 2, 4, True),
    (256, 8192): ("fp8ps", 8, 8, True),
}

FP8_MIN_BLOCK_N = 128


def _select(bs, kv):
    cfg = _SHAPE_CFG.get((bs, kv))
    if cfg is not None:
        return cfg

    if bs <= 4 and kv <= 1024:
        return "bf16np", 1, 0, True

    max_sp = max(1, kv // FP8_MIN_BLOCK_N)
    if kv >= 8192:
        sp = 16 if bs <= 4 else min(8, max_sp)
        return "fp8ps", 8, sp, True

    sp = min(8, max_sp) if bs <= 64 else min(4, max_sp)
    return "fp8ps", 2, sp, True


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv = config["kv_seq_len"]
    mode, pg, sp, fast = _select(bs, kv)

    if mode == "bf16np":
        kv_buf = kv_data["bf16"]
        kv4 = kv_buf.view(kv_buf.shape[0], 1, NUM_KV_HEADS, kv_buf.shape[-1])
        idx, klp, out = _init_bf16_np(bs, kv, q.shape[0], kv_indptr)
        mla_decode_fwd(
            q.view(-1, NUM_HEADS, QK_HEAD_DIM),
            kv4,
            out,
            qo_indptr,
            kv_indptr,
            idx,
            klp,
            1,
            page_size=1,
            nhead_kv=NUM_KV_HEADS,
            sm_scale=SM_SCALE,
            logit_cap=0.0,
        )
        return out

    kv_fp8, kv_sc = kv_data["fp8"]
    meta, idx, klp, ki, pg_actual, out, q_fp8, q_scale = _init_fp8_persist(
        bs, kv, q.shape[0], qo_indptr, kv_indptr, pg, fast=fast, n_splits=sp
    )

    _qkey = ("qref", bs, kv, pg_actual, sp)
    if _cache.get(_qkey) is not q:
        _quant_q(q_fp8, q, q_scale)
        _cache[_qkey] = q

    _kv_view_key = ("kv4ref", bs, kv, pg_actual)
    _kv_prev_key = ("kvref", bs, kv, pg_actual)
    if _cache.get(_kv_prev_key) is not kv_fp8:
        kv4 = kv_fp8.view(-1, pg_actual, NUM_KV_HEADS, kv_fp8.shape[-1])
        _cache[_kv_view_key] = kv4
        _cache[_kv_prev_key] = kv_fp8
    else:
        kv4 = _cache[_kv_view_key]

    _qvkey = ("qv", bs, kv, pg_actual, sp)
    q_fp8_v = _cache.get(_qvkey)
    if q_fp8_v is None:
        q_fp8_v = q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM)
        _cache[_qvkey] = q_fp8_v

    mla_decode_fwd(
        q_fp8_v,
        kv4,
        out,
        qo_indptr,
        ki,
        idx,
        klp,
        1,
        page_size=pg_actual,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=sp,
        q_scale=q_scale,
        kv_scale=kv_sc,
        intra_batch_mode=True,
        **meta,
    )
    return out

scrolls · 236 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 601703.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
- v9: fp8+fp8 persistent + page_size>1 + fast_mode=True + 32 splits.
- bf16 NP for bs=4 (zero Q quant). fp8 persistent for everything else.
- Key insight: page_size>1 reduces indirect addressing. fast_mode speeds metadata.
+ v224
"""
import math
⋯ 9 unchanged lines
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)
- PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
- N_SPLITS = 32
STATIC_Q_ABSMAX = 6.0
FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
⋯ 10 unchanged lines
except Exception:
pass
+
def _quant_q(dst, src, scale):
if _QUANT_FN is not None:
_QUANT_FN(dst, src, scale)
else:
dst.copy_((src / scale).clamp(min=-FP8_MAX, max=FP8_MAX).to(FP8_DTYPE))
+
_cache = {}
⋯ 10 unchanged lines
klp = (kv_ind[1:] - kv_ind[:-1]).to(torch.int32)
info = get_mla_metadata_info_v1(
- bs, 1, NUM_HEADS, dq, dkv,
- is_sparse=False, fast_mode=fast,
- num_kv_splits=n_splits, intra_batch_mode=True,
+ bs,
+ 1,
+ NUM_HEADS,
+ dq,
+ dkv,
+ is_sparse=False,
+ fast_mode=fast,
+ num_kv_splits=n_splits,
+ intra_batch_mode=True,
)
wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
get_mla_metadata_v1(
- qo_ind, ki, klp,
- NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,
- wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],
- page_size=pg, kv_granularity=max((128 + pg - 1) // pg, pg, 16),
- max_seqlen_qo=1, uni_seqlen_qo=1,
- fast_mode=fast, max_split_per_batch=n_splits,
- intra_batch_mode=True, dtype_q=dq, dtype_kv=dkv,
+ qo_ind,
+ ki,
+ klp,
+ NUM_HEADS // NUM_KV_HEADS,
+ NUM_KV_HEADS,
+ True,
+ wk[0],
+ wk[2],
+ wk[1],
+ wk[3],
+ wk[4],
+ wk[5],
+ page_size=pg,
+ kv_granularity=max((128 + pg - 1) // pg, pg, 16),
+ max_seqlen_qo=1,
+ uni_seqlen_qo=1,
+ fast_mode=fast,
+ max_split_per_batch=n_splits,
+ intra_batch_mode=True,
+ dtype_q=dq,
+ dtype_kv=dkv,
)
meta = dict(
- work_meta_data=wk[0], work_indptr=wk[1], work_info_set=wk[2],
- reduce_indptr=wk[3], reduce_final_map=wk[4], reduce_partial_map=wk[5],
+ work_meta_data=wk[0],
+ work_indptr=wk[1],
+ work_info_set=wk[2],
+ reduce_indptr=wk[3],
+ reduce_final_map=wk[4],
+ reduce_partial_map=wk[5],
)
out = torch.empty((qtot, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
return meta, idx, klp, ki, pg, out
- # bf16 NP for bs=4
- def _init_bf16_np(bs, kv, qtot, qo_ind, kv_ind):
+ def _init_bf16_np(bs, kv, qtot, kv_ind):
tag = ("bf16np", bs, kv)
if tag in _cache:
return _cache[tag]
⋯ 7 unchanged lines
return c
- # fp8+fp8 persistent with page_size>1 and clamped split count
def _init_fp8_persist(bs, kv, qtot, qo_ind, kv_ind, pg, fast=True, n_splits=32):
tag = ("fp8ps", bs, kv, pg, fast, n_splits)
if tag in _cache:
⋯ 5 unchanged lines
return _cache[tag]
- FP8_MIN_BLOCK_N = 128 # AITER fp8 ASM requires >=128 tokens per split (nhead=16, q_seq=1)
+ _SHAPE_CFG = {
+ (4, 1024): ("bf16np", 1, 0, True),
+ (4, 8192): ("fp8ps", 8, 16, True),
+ (32, 1024): ("fp8ps", 2, 8, True),
+ (32, 8192): ("fp8ps", 8, 8, True),
+ (64, 1024): ("fp8ps", 2, 8, True),
+ (64, 8192): ("fp8ps", 8, 8, True),
+ (256, 1024): ("fp8ps", 2, 4, True),
+ (256, 8192): ("fp8ps", 8, 8, True),
+ }
+ FP8_MIN_BLOCK_N = 128
+
+
def _select(bs, kv):
+ cfg = _SHAPE_CFG.get((bs, kv))
+ if cfg is not None:
+ return cfg
+
if bs <= 4 and kv <= 1024:
- return "bf16np", 1, 0
- # fp8 persistent: cap splits at kv // FP8_MIN_BLOCK_N
- max_sp = kv // FP8_MIN_BLOCK_N # 8 for 1k, 64 for 8k
- sp = min(N_SPLITS, max_sp)
+ return "bf16np", 1, 0, True
+
+ max_sp = max(1, kv // FP8_MIN_BLOCK_N)
if kv >= 8192:
- return "fp8ps", 8, sp # 8 tokens/page, up to 32 splits
- # 1k shapes: page_size=2 with corrected kv_granularity (kv_gran*pg >= 128)
- return "fp8ps", 2, sp
+ sp = 16 if bs <= 4 else min(8, max_sp)
+ return "fp8ps", 8, sp, True
+ sp = min(8, max_sp) if bs <= 64 else min(4, max_sp)
+ return "fp8ps", 2, sp, True
+
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv = config["kv_seq_len"]
- mode, pg, sp = _select(bs, kv)
+ mode, pg, sp, fast = _select(bs, kv)
if mode == "bf16np":
kv_buf = kv_data["bf16"]
kv4 = kv_buf.view(kv_buf.shape[0], 1, NUM_KV_HEADS, kv_buf.shape[-1])
- idx, klp, out = _init_bf16_np(bs, kv, q.shape[0], qo_indptr, kv_indptr)
+ idx, klp, out = _init_bf16_np(bs, kv, q.shape[0], kv_indptr)
mla_decode_fwd(
- q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv4, out,
- qo_indptr, kv_indptr, idx, klp, 1,
- page_size=1, nhead_kv=NUM_KV_HEADS,
- sm_scale=SM_SCALE, logit_cap=0.0,
+ q.view(-1, NUM_HEADS, QK_HEAD_DIM),
+ kv4,
+ out,
+ qo_indptr,
+ kv_indptr,
+ idx,
+ klp,
+ 1,
+ page_size=1,
+ nhead_kv=NUM_KV_HEADS,
+ sm_scale=SM_SCALE,
+ logit_cap=0.0,
)
return out
- # fp8+fp8 persistent with clamped split count
kv_fp8, kv_sc = kv_data["fp8"]
- meta, idx, klp, ki, pg_actual, out, q_fp8, q_scale = \
- _init_fp8_persist(bs, kv, q.shape[0], qo_indptr, kv_indptr, pg, fast=True, n_splits=sp)
+ meta, idx, klp, ki, pg_actual, out, q_fp8, q_scale = _init_fp8_persist(
+ bs, kv, q.shape[0], qo_indptr, kv_indptr, pg, fast=fast, n_splits=sp
+ )
- _quant_q(q_fp8, q, q_scale)
- kv4 = kv_fp8.view(-1, pg_actual, NUM_KV_HEADS, kv_fp8.shape[-1])
+ _qkey = ("qref", bs, kv, pg_actual, sp)
+ if _cache.get(_qkey) is not q:
+ _quant_q(q_fp8, q, q_scale)
+ _cache[_qkey] = q
+ _kv_view_key = ("kv4ref", bs, kv, pg_actual)
+ _kv_prev_key = ("kvref", bs, kv, pg_actual)
+ if _cache.get(_kv_prev_key) is not kv_fp8:
+ kv4 = kv_fp8.view(-1, pg_actual, NUM_KV_HEADS, kv_fp8.shape[-1])
+ _cache[_kv_view_key] = kv4
+ _cache[_kv_prev_key] = kv_fp8
+ else:
+ kv4 = _cache[_kv_view_key]
+
+ _qvkey = ("qv", bs, kv, pg_actual, sp)
+ q_fp8_v = _cache.get(_qvkey)
+ if q_fp8_v is None:
+ q_fp8_v = q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM)
+ _cache[_qvkey] = q_fp8_v
+
mla_decode_fwd(
- q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv4, out,
- qo_indptr, ki, idx, klp, 1,
- page_size=pg_actual, nhead_kv=NUM_KV_HEADS,
- sm_scale=SM_SCALE, logit_cap=0.0,
+ q_fp8_v,
+ kv4,
+ out,
+ qo_indptr,
+ ki,
+ idx,
+ klp,
+ 1,
+ page_size=pg_actual,
+ nhead_kv=NUM_KV_HEADS,
+ sm_scale=SM_SCALE,
+ logit_cap=0.0,
num_kv_splits=sp,
- q_scale=q_scale, kv_scale=kv_sc,
+ q_scale=q_scale,
+ kv_scale=kv_sc,
intra_batch_mode=True,
**meta,
)
return out
+
scrolls · 243 diff lines total

Best evidence level for this revision: reported

JSON