Skip to content
KernelIndex
Search⌘K

submission 736894

NinoHeather · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

my_submission_refact.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-736894?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
22.1µs
#4 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a67cc2fd250671ae326070458a9d2a5c5713d73e44c1f1ec304e8bd1b63a4f39
license declaredunknown
license concludedunknown
authorsNinoHeather
imported2026-08-15

Kernel source

my_submission_refact.py321 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

import os
import torch
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
from aiter import mla_decode_stage1_asm_fwd, mla_reduce_v1
from aiter.mla import mla_decode_fwd


HAS_STATIC_QUANT = False
try:
    # Prefer JIT quant op when available.
    from aiter.jit.module_quant import static_per_tensor_quant
    HAS_STATIC_QUANT = True
except Exception:
    try:
        from aiter.ops.quant import static_per_tensor_quant
        HAS_STATIC_QUANT = True
    except Exception:
        pass

try:
    from aiter.ops.quant import dynamic_per_tensor_quant
    HAS_DYNAMIC_QUANT = True
except ImportError:
    HAS_DYNAMIC_QUANT = False


FP8_T = aiter_dtypes.fp8
BF16_T = torch.bfloat16
N_HEADS = 16
N_KV_HEADS = 1
QK_DIM = 576
V_DIM = 512
ATTN_SCALE = 1.0 / (QK_DIM ** 0.5)
STATIC_Q_SCALE = torch.tensor([0.1], dtype=torch.float32, device="cuda")


# (batch_size, kv_seq_len) -> (page_size, num_kv_splits, is_sparse, kv_granularity_override)
SHAPE_POLICY = {
    (4, 1024): (2, 8, False, None),
    (4, 8192): (8, 16, False, None),
    (32, 1024): (2, 8, False, None),
    (32, 8192): (8, 12, False, None),
    (64, 1024): (2, 1, False, 4),
    (64, 8192): (8, 12, False, None),
    (256, 1024): (2, 1, True, 2),
    (256, 8192): (8, 8, False, None),
}

NP_8K_PAGE = {
    4: 2048,
    32: 1024,
    64: 512,
    256: 2048,
}

FWD_A8W8_1K_SHAPES = {
    (4, 1024),
    (32, 1024),
}


def _build_shape_state(
    bs: int,
    kv_len: int,
    page_size: int,
    splits: int,
    sparse: bool,
    kv_granularity_override: int | None,
) -> dict:
    pages_per_req = kv_len // page_size
    total_pages = bs * pages_per_req

    qo_prefix = torch.arange(bs + 1, dtype=torch.int32, device="cuda")
    kv_pages_prefix = torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req
    kv_last_page_len = torch.full((bs,), page_size, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")

    meta_info = get_mla_metadata_info_v1(
        bs,
        1,
        N_HEADS,
        FP8_T,
        FP8_T,
        is_sparse=sparse,
        fast_mode=True,
        num_kv_splits=splits,
        intra_batch_mode=True,
    )
    work_meta = [torch.empty(sz, dtype=dt, device="cuda") for sz, dt in meta_info]
    wm, wi, wis, ri, rfm, rpm = work_meta

    kv_granularity = kv_granularity_override if kv_granularity_override is not None else max(page_size, 16)
    get_mla_metadata_v1(
        qo_prefix,
        kv_pages_prefix,
        kv_last_page_len,
        N_HEADS // N_KV_HEADS,
        N_KV_HEADS,
        False,
        wm,
        wis,
        wi,
        ri,
        rfm,
        rpm,
        page_size=page_size,
        kv_granularity=kv_granularity,
        max_seqlen_qo=1,
        uni_seqlen_qo=1,
        fast_mode=True,
        max_split_per_batch=splits,
        intra_batch_mode=True,
        dtype_q=FP8_T,
        dtype_kv=FP8_T,
    )

    partial_rows = rpm.size(0)
    return {
        "page_size": page_size,
        "splits": splits,
        "qo_prefix": qo_prefix,
        "kv_pages_prefix": kv_pages_prefix,
        "kv_last_page_len": kv_last_page_len,
        "kv_indices": kv_indices,
        "wm": wm,
        "wi": wi,
        "wis": wis,
        "ri": ri,
        "rfm": rfm,
        "rpm": rpm,
        "partial_logits": torch.empty((partial_rows, 1, N_HEADS, V_DIM), dtype=torch.float32, device="cuda"),
        "partial_lse": torch.empty((partial_rows, 1, N_HEADS, 1), dtype=torch.float32, device="cuda"),
        "out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
        "q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
        "q_scale": STATIC_Q_SCALE.clone(),
    }


def _build_fwd_a8w8_1k_state(bs: int, kv_len: int, page_size: int = 2) -> dict:
    pages_per_req = kv_len // page_size
    total_pages = bs * pages_per_req
    return {
        "page_size": page_size,
        "kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
        "kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,
        "kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),
        "out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
        "q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
        "q_scale": STATIC_Q_SCALE.clone(),
    }


def _build_np8k_state(bs: int, kv_len: int, page_size: int) -> dict:
    pages_per_req = kv_len // page_size
    total_pages = bs * pages_per_req
    return {
        "page_size": page_size,
        "kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
        "kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,
        "kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),
        "num_kv_splits_indptr": torch.arange(bs + 1, dtype=torch.int32, device="cuda"),
        "split_logits": torch.empty((bs, 1, N_HEADS, V_DIM), dtype=torch.float32, device="cuda"),
        "split_lse": torch.empty((bs, 1, N_HEADS, 1), dtype=torch.float32, device="cuda"),
        "out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
        "q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
        "q_scale": STATIC_Q_SCALE.clone(),
    }


# Eager state init for all leaderboard shapes.
SHAPE_STATE = {
    shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2], cfg[3])
    for shape, cfg in SHAPE_POLICY.items()
}
BF16_STATE = {
    shape: _build_fwd_a8w8_1k_state(shape[0], shape[1], 2)
    for shape in FWD_A8W8_1K_SHAPES
}
NP8K_STATE = {
    (bs, 8192): _build_np8k_state(bs, 8192, ps)
    for bs, ps in NP_8K_PAGE.items()
}


def _quantize_query_to_fp8(q_bf16: torch.Tensor, q_fp8: torch.Tensor, q_scale: torch.Tensor) -> None:
    if HAS_STATIC_QUANT:
        static_per_tensor_quant(q_fp8, q_bf16, STATIC_Q_SCALE)
        return

    if HAS_DYNAMIC_QUANT:
        dynamic_per_tensor_quant(q_fp8, q_bf16.view_as(q_fp8), q_scale)
        return

    finfo = torch.finfo(FP8_T)
    amax = q_bf16.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    q_fp8.copy_((q_bf16 / scale).clamp(finfo.min, finfo.max).to(FP8_T))
    q_scale.fill_(scale.item())


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_len = config["kv_seq_len"]
    total_q = q.shape[0]
    shape = (bs, kv_len)

    if shape in FWD_A8W8_1K_SHAPES:
        fwd_state = BF16_STATE[shape]
        kv_fp8, kv_scale = kv_data["fp8"]
        page_size = fwd_state["page_size"]
        kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
        q_fp8 = fwd_state["q_fp8"]
        q_scale = fwd_state["q_scale"]
        _quantize_query_to_fp8(q, q_fp8, q_scale)
        out = fwd_state["out_bf16"]
        mla_decode_fwd(
            q=q_fp8,
            kv_buffer=kv_4d,
            o=out,
            qo_indptr=qo_indptr,
            kv_indptr=fwd_state["kv_pages_prefix"],
            kv_indices=fwd_state["kv_indices"],
            kv_last_page_lens=fwd_state["kv_last_page_lens"],
            max_seqlen_q=1,
            page_size=page_size,
            nhead_kv=1,
            sm_scale=ATTN_SCALE,
            q_scale=q_scale,
            kv_scale=kv_scale,
            intra_batch_mode=True,
        )
        return out[:total_q]

    if shape in NP8K_STATE:
        np_state = NP8K_STATE[shape]
        kv_fp8, kv_scale = kv_data["fp8"]
        page_size = np_state["page_size"]
        kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
        q_fp8 = np_state["q_fp8"]
        q_scale = np_state["q_scale"]
        _quantize_query_to_fp8(q, q_fp8, q_scale)

        out = np_state["out_bf16"]
        split_logits = np_state["split_logits"]
        split_lse = np_state["split_lse"]
        out.zero_()
        split_lse.fill_(float("-inf"))
        mla_decode_stage1_asm_fwd(
            q_fp8,
            kv_4d,
            qo_indptr,
            np_state["kv_pages_prefix"],
            np_state["kv_indices"],
            np_state["kv_last_page_lens"],
            np_state["num_kv_splits_indptr"],
            None,
            None,
            None,
            1,
            page_size,
            N_KV_HEADS,
            ATTN_SCALE,
            split_logits,
            split_lse,
            out,
            q_scale,
            kv_scale,
        )
        return out[:total_q]

    state = SHAPE_STATE[shape]
    kv_fp8, kv_scale = kv_data["fp8"]
    page_size = state["page_size"]
    kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)

    q_fp8 = state["q_fp8"]
    q_scale = state["q_scale"]
    _quantize_query_to_fp8(q, q_fp8, q_scale)

    mla_decode_stage1_asm_fwd(
        q_fp8,
        kv_4d,
        state["qo_prefix"],
        state["kv_pages_prefix"],
        state["kv_indices"],
        state["kv_last_page_len"],
        None,
        state["wm"],
        state["wi"],
        state["wis"],
        1,
        page_size,
        N_KV_HEADS,
        ATTN_SCALE,
        state["partial_logits"],
        state["partial_lse"],
        state["out_bf16"],
        q_scale,
        kv_scale,
    )

    if state["splits"] > 1:
        mla_reduce_v1(
            state["partial_logits"],
            state["partial_lse"],
            state["ri"],
            state["rfm"],
            state["rpm"],
            1,
            state["out_bf16"],
            None,
        )

    return state["out_bf16"][:total_q]
scrolls · 321 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 687679.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """
- Refactored equivalent of my_submission_demo.py.
-
- Behavior/perf intent:
- - identical shape policy and kernel path (a8w8 via stage1_asm + reduce_v1)
- - identical eager buffer/materialization strategy
- - identical quantization fallback chain
- """
-
+ import os
import torch
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
from aiter import mla_decode_stage1_asm_fwd, mla_reduce_v1
+ from aiter.mla import mla_decode_fwd
HAS_STATIC_QUANT = False
try:
- from aiter.ops.quant import static_per_tensor_quant
+ # Prefer JIT quant op when available.
+ from aiter.jit.module_quant import static_per_tensor_quant
HAS_STATIC_QUANT = True
- except ImportError:
- pass
+ except Exception:
+ try:
+ from aiter.ops.quant import static_per_tensor_quant
+ HAS_STATIC_QUANT = True
+ except Exception:
+ pass
try:
from aiter.ops.quant import dynamic_per_tensor_quant
⋯ 12 unchanged lines
STATIC_Q_SCALE = torch.tensor([0.1], dtype=torch.float32, device="cuda")
- # (batch_size, kv_seq_len) -> (page_size, num_kv_splits, is_sparse)
+ # (batch_size, kv_seq_len) -> (page_size, num_kv_splits, is_sparse, kv_granularity_override)
SHAPE_POLICY = {
- (4, 1024): (1, 8, False),
- (4, 8192): (8, 16, False),
- (32, 1024): (1, 8, False),
- (32, 8192): (8, 12, False),
- (64, 1024): (2, 4, False),
- (64, 8192): (8, 12, False),
- (256, 1024): (2, 1, True),
- (256, 8192): (8, 8, False),
+ (4, 1024): (2, 8, False, None),
+ (4, 8192): (8, 16, False, None),
+ (32, 1024): (2, 8, False, None),
+ (32, 8192): (8, 12, False, None),
+ (64, 1024): (2, 1, False, 4),
+ (64, 8192): (8, 12, False, None),
+ (256, 1024): (2, 1, True, 2),
+ (256, 8192): (8, 8, False, None),
}
+ NP_8K_PAGE = {
+ 4: 2048,
+ 32: 1024,
+ 64: 512,
+ 256: 2048,
+ }
- def _build_shape_state(bs: int, kv_len: int, page_size: int, splits: int, sparse: bool) -> dict:
+ FWD_A8W8_1K_SHAPES = {
+ (4, 1024),
+ (32, 1024),
+ }
+
+
+ def _build_shape_state(
+ bs: int,
+ kv_len: int,
+ page_size: int,
+ splits: int,
+ sparse: bool,
+ kv_granularity_override: int | None,
+ ) -> dict:
pages_per_req = kv_len // page_size
total_pages = bs * pages_per_req
⋯ 16 unchanged lines
work_meta = [torch.empty(sz, dtype=dt, device="cuda") for sz, dt in meta_info]
wm, wi, wis, ri, rfm, rpm = work_meta
+ kv_granularity = kv_granularity_override if kv_granularity_override is not None else max(page_size, 16)
get_mla_metadata_v1(
qo_prefix,
kv_pages_prefix,
⋯ 8 unchanged lines
rfm,
rpm,
page_size=page_size,
- kv_granularity=max(page_size, 16),
+ kv_granularity=kv_granularity,
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=True,
⋯ 25 unchanged lines
}
+ def _build_fwd_a8w8_1k_state(bs: int, kv_len: int, page_size: int = 2) -> dict:
+ pages_per_req = kv_len // page_size
+ total_pages = bs * pages_per_req
+ return {
+ "page_size": page_size,
+ "kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
+ "kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,
+ "kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),
+ "out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
+ "q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
+ "q_scale": STATIC_Q_SCALE.clone(),
+ }
+
+
+ def _build_np8k_state(bs: int, kv_len: int, page_size: int) -> dict:
+ pages_per_req = kv_len // page_size
+ total_pages = bs * pages_per_req
+ return {
+ "page_size": page_size,
+ "kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
+ "kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,
+ "kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),
+ "num_kv_splits_indptr": torch.arange(bs + 1, dtype=torch.int32, device="cuda"),
+ "split_logits": torch.empty((bs, 1, N_HEADS, V_DIM), dtype=torch.float32, device="cuda"),
+ "split_lse": torch.empty((bs, 1, N_HEADS, 1), dtype=torch.float32, device="cuda"),
+ "out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
+ "q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
+ "q_scale": STATIC_Q_SCALE.clone(),
+ }
+
+
# Eager state init for all leaderboard shapes.
SHAPE_STATE = {
- shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2])
+ shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2], cfg[3])
for shape, cfg in SHAPE_POLICY.items()
}
+ BF16_STATE = {
+ shape: _build_fwd_a8w8_1k_state(shape[0], shape[1], 2)
+ for shape in FWD_A8W8_1K_SHAPES
+ }
+ NP8K_STATE = {
+ (bs, 8192): _build_np8k_state(bs, 8192, ps)
+ for bs, ps in NP_8K_PAGE.items()
+ }
def _quantize_query_to_fp8(q_bf16: torch.Tensor, q_fp8: torch.Tensor, q_scale: torch.Tensor) -> None:
⋯ 13 unchanged lines
def custom_kernel(data: input_t) -> output_t:
- q, kv_data, _, _, config = data
+ q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
total_q = q.shape[0]
+ shape = (bs, kv_len)
- state = SHAPE_STATE[(bs, kv_len)]
+ if shape in FWD_A8W8_1K_SHAPES:
+ fwd_state = BF16_STATE[shape]
+ kv_fp8, kv_scale = kv_data["fp8"]
+ page_size = fwd_state["page_size"]
+ kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
+ q_fp8 = fwd_state["q_fp8"]
+ q_scale = fwd_state["q_scale"]
+ _quantize_query_to_fp8(q, q_fp8, q_scale)
+ out = fwd_state["out_bf16"]
+ mla_decode_fwd(
+ q=q_fp8,
+ kv_buffer=kv_4d,
+ o=out,
+ qo_indptr=qo_indptr,
+ kv_indptr=fwd_state["kv_pages_prefix"],
+ kv_indices=fwd_state["kv_indices"],
+ kv_last_page_lens=fwd_state["kv_last_page_lens"],
+ max_seqlen_q=1,
+ page_size=page_size,
+ nhead_kv=1,
+ sm_scale=ATTN_SCALE,
+ q_scale=q_scale,
+ kv_scale=kv_scale,
+ intra_batch_mode=True,
+ )
+ return out[:total_q]
+
+ if shape in NP8K_STATE:
+ np_state = NP8K_STATE[shape]
+ kv_fp8, kv_scale = kv_data["fp8"]
+ page_size = np_state["page_size"]
+ kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
+ q_fp8 = np_state["q_fp8"]
+ q_scale = np_state["q_scale"]
+ _quantize_query_to_fp8(q, q_fp8, q_scale)
+
+ out = np_state["out_bf16"]
+ split_logits = np_state["split_logits"]
+ split_lse = np_state["split_lse"]
+ out.zero_()
+ split_lse.fill_(float("-inf"))
+ mla_decode_stage1_asm_fwd(
+ q_fp8,
+ kv_4d,
+ qo_indptr,
+ np_state["kv_pages_prefix"],
+ np_state["kv_indices"],
+ np_state["kv_last_page_lens"],
+ np_state["num_kv_splits_indptr"],
+ None,
+ None,
+ None,
+ 1,
+ page_size,
+ N_KV_HEADS,
+ ATTN_SCALE,
+ split_logits,
+ split_lse,
+ out,
+ q_scale,
+ kv_scale,
+ )
+ return out[:total_q]
+
+ state = SHAPE_STATE[shape]
kv_fp8, kv_scale = kv_data["fp8"]
page_size = state["page_size"]
kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
scrolls · 237 diff lines total

Best evidence level for this revision: reported

JSON