Skip to content
KernelIndex
Search⌘K

submission 697219

parcadei · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-697219?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
57.4µs
#214 of 766
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:52d89d34c1e6a72bea2e93c2d3df1b6d9775d9b10a0515155fa309b3f0535a95
license declaredunknown
license concludedunknown
authorsparcadei
imported2026-08-15

Techniques

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

persistent-kernel- S2: Persistent FP8 (gran=128, ns=32) + C++ reduce

Kernel source

submission.py327 lines
"""
Mixed-MLA submission: aiter ASM with hybrid dispatch.

Architecture:
- S1: BF16 16-split + C++ reduce (no Q copy, no FP8 overhead)
- S2: Persistent FP8 (gran=128, ns=32) + C++ reduce
- S3: BF16 4-split + C++ reduce
- S4: FP8 8-split + C++ reduce
- S5: Persistent FP8 (gran=64, ns=32) + C++ reduce
- S6: FP8 4-split + C++ reduce
- S7/S8: NP1 FP8 (1-split, no reduce)
- Closure-based dispatch: pre-built callables per shape capture all tensor refs
"""
from __future__ import annotations

import math
import os
from typing import Any

os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

import torch
import torch.nn.functional as F

from task import input_t, output_t

PAGE_SIZE = 1
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)

# ---- Aiter lazy-load --------------------------------------------------------

_AITER_API: dict[str, Any] | None = None
_AITER_IMPORT_ERROR: Exception | None = None
_FN_STAGE1: Any = None
_FN_METADATA: Any = None
_FN_REDUCE: Any = None
_FP8_DTYPE: Any = None
_META_INFO_FN: Any = None


def _load_aiter() -> dict[str, Any]:
    global _AITER_API, _AITER_IMPORT_ERROR, _FN_STAGE1, _FN_METADATA, _FN_REDUCE, _FP8_DTYPE, _META_INFO_FN
    if _AITER_API is not None:
        return _AITER_API
    if _AITER_IMPORT_ERROR is not None:
        raise _AITER_IMPORT_ERROR
    try:
        import aiter
        from aiter import dtypes as aiter_dtypes
        from aiter.ops.attention import get_mla_metadata_info_v1
        _AITER_API = {"ok": True}
        _FP8_DTYPE = aiter_dtypes.fp8
        _FN_STAGE1 = aiter.mla_decode_stage1_asm_fwd
        _FN_METADATA = aiter.get_mla_metadata_v1
        _FN_REDUCE = aiter.mla_reduce_v1
        _META_INFO_FN = get_mla_metadata_info_v1
        return _AITER_API
    except Exception as exc:
        _AITER_IMPORT_ERROR = exc
        raise


def _torch_fallback(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_bf16 = kv_data["bf16"]
    sm_scale = float(config.get("sm_scale", SM_SCALE))
    kv_lora_rank = int(config["kv_lora_rank"])
    out_chunks = []
    batch_size = qo_indptr.shape[0] - 1
    for batch_idx in range(batch_size):
        q_start = int(qo_indptr[batch_idx].item())
        q_end = int(qo_indptr[batch_idx + 1].item())
        kv_start = int(kv_indptr[batch_idx].item())
        kv_end = int(kv_indptr[batch_idx + 1].item())
        qi = q[q_start:q_end].float().permute(1, 0, 2)
        kv_rows = kv_bf16[kv_start:kv_end, 0].float()
        scores = torch.matmul(qi * sm_scale, kv_rows.T)
        probs = F.softmax(scores, dim=-1)
        values = kv_rows[:, :kv_lora_rank]
        out = torch.matmul(probs, values).permute(1, 0, 2)
        out_chunks.append(out.to(torch.bfloat16))
    return torch.cat(out_chunks, dim=0)


# ---- Entry point ------------------------------------------------------------

_SHAPE_DISPATCH: dict[tuple, Any] = {}

# Shape -> num_splits for FP8 multi-split + C++ reduce
_FP8_SPLITS: dict[tuple[int, int], int] = {
    (32, 8192): 8,    # S4: FP8 8-split + C++ reduce
    (64, 8192): 4,    # S6: FP8 4-split + C++ reduce
}

# Shapes using NP1 (1-split, no reduce, stage1 writes directly to output)
_NP1_SHAPES: set[tuple[int, int]] = {
    (256, 1024),      # S7
    (256, 8192),      # S8
}

# Shape -> num_splits for BF16 multi-split + C++ reduce
_BF16_SPLITS: dict[tuple[int, int], int] = {
    (4, 1024): 16,    # S1: BF16 16-split + C++ reduce
    (32, 1024): 4,    # S3: BF16 4-split — FP8 non-persistent fails ranked
}

# Shape -> (kv_granularity, num_splits) for persistent mode + C++ reduce
_PERSISTENT_CONFIG: dict[tuple[int, int], tuple[int, int]] = {
    (4, 8192): (128, 32),   # S2: persistent, gran=128
    (64, 1024): (64, 32),   # S5: persistent, gran=64
}


def _build_dispatch(key, q, data):
    """Cold path: pre-allocate tensors and build a closure for the hot path."""
    if _FN_STAGE1 is None:
        _load_aiter()
    config = data[4]
    bs, kv_len = key
    nh = int(config["num_heads"])
    qs = int(config["q_seq_len"])
    total_kv = bs * kv_len
    device = q.device

    kv_fp8 = data[1]["fp8"][0]
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    output = torch.empty((q.shape[0], nh, V_HEAD_DIM), dtype=torch.bfloat16, device=device)

    rows = q.numel() // q.shape[-1]
    q_flat_shape = (rows, q.shape[-1])
    q_fp8_buf = torch.empty(q_flat_shape, dtype=_FP8_DTYPE, device=device)
    q_shaped = q_fp8_buf.view(-1, nh, QK_HEAD_DIM)
    # Cached unit scale for cast-only FP8 quantization.
    # Persistent ASM kernel does NOT modify q_scale in-place (verified: reference
    # caches a single unit_scale tensor across calls without resetting).
    # Non-persistent ASM kernel DOES modify q_scale — use fill_(1.0) there.
    q_scale = torch.ones(1, dtype=torch.float32, device=device)

    qo_indptr = data[2].clone()
    kv_indptr = data[3].clone()
    kv_view_shape = (kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)

    bf16_splits = _BF16_SPLITS.get(key)
    persistent_config = _PERSISTENT_CONFIG.get(key)
    fp8_splits = _FP8_SPLITS.get(key)
    is_np1 = key in _NP1_SHAPES

    if bf16_splits is not None:
        # BF16 multi-split + C++ reduce (S1, S3)
        total_q = q.shape[0]
        num_splits = bf16_splits
        splits_indptr = torch.arange(0, (bs + 1) * num_splits, num_splits, dtype=torch.int, device=device)
        kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
        np_logits = torch.empty((total_q, num_splits, nh, V_HEAD_DIM), dtype=torch.float32, device=device)
        np_lse = torch.empty((total_q, num_splits, nh, 1), dtype=torch.float32, device=device)
        kv_bf16 = data[1]["bf16"]
        kv_bf16_view = (kv_bf16.shape[0], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)

        np_total = total_q * num_splits
        reduce_indptr = torch.arange(0, np_total + num_splits, num_splits, dtype=torch.int32, device=device)[:total_q + 1]
        reduce_partial_map = torch.arange(np_total, dtype=torch.int32, device=device)
        rl_view = np_logits.reshape(np_total, 1, nh, V_HEAD_DIM)
        rls_view = np_lse.reshape(np_total, 1, nh, 1)

        def _dispatch_bf16(data):
            q_bf16 = data[0].view(-1, nh, QK_HEAD_DIM)
            kv_v = data[1]["bf16"].view(kv_bf16_view)
            _FN_STAGE1(
                q_bf16, kv_v,
                qo_indptr, kv_indptr, kv_indices,
                kv_last, splits_indptr, None, None, None,
                qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
                np_logits, np_lse, output, None, None,
            )
            _FN_REDUCE(rl_view, rls_view, reduce_indptr, None, reduce_partial_map, qs, output, None)
            return output

        _SHAPE_DISPATCH[key] = _dispatch_bf16
        return _dispatch_bf16

    elif persistent_config is not None:
        # PERSISTENT mode: metadata + persistent ASM kernel + C++ reduce (S2, S5)
        kv_gran, num_splits = persistent_config
        total_q = q.shape[0]
        kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

        meta_info = _META_INFO_FN(
            batch_size=bs,
            max_seqlen_qo=qs,
            num_head_qo=nh,
            q_dtype=_FP8_DTYPE,
            kv_dtype=_FP8_DTYPE,
            is_sparse=False,
            fast_mode=True,
            num_kv_splits=num_splits,
            intra_batch_mode=True,
        )

        work_metadata_ptrs = torch.empty(meta_info[0][0], dtype=meta_info[0][1], device=device)
        work_indptr = torch.empty(meta_info[1][0], dtype=meta_info[1][1], device=device)
        work_info_set = torch.empty(meta_info[2][0], dtype=meta_info[2][1], device=device)
        reduce_indptr = torch.empty(meta_info[3][0], dtype=meta_info[3][1], device=device)
        reduce_final_map = torch.empty(meta_info[4][0], dtype=meta_info[4][1], device=device)
        reduce_partial_map = torch.empty(meta_info[5][0], dtype=meta_info[5][1], device=device)

        _FN_METADATA(
            qo_indptr, kv_indptr, kv_last,
            nh, NUM_KV_HEADS, True,
            work_metadata_ptrs, work_info_set, work_indptr,
            reduce_indptr, reduce_final_map, reduce_partial_map,
            PAGE_SIZE, kv_gran,
            qs, qs, True, -1, num_splits, True,
            _FP8_DTYPE, _FP8_DTYPE,
        )

        n_partials = reduce_partial_map.shape[0]
        logits = torch.empty(
            (n_partials * qs, 1, nh, V_HEAD_DIM),
            dtype=torch.float32, device=device,
        )
        attn_lse = torch.empty(
            (n_partials * qs, 1, nh, 1),
            dtype=torch.float32, device=device,
        )

        def _dispatch_persistent(data):
            q_fp8_buf.copy_(data[0].view(q_flat_shape))
            # No fill_(1.0) needed: persistent ASM kernel does not modify q_scale
            fp8 = data[1]["fp8"]
            kv_v = fp8[0].view(kv_view_shape)
            _FN_STAGE1(
                q_shaped, kv_v,
                qo_indptr, kv_indptr, kv_indices,
                kv_last, None,
                work_metadata_ptrs, work_indptr, work_info_set,
                qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
                logits, attn_lse, output, q_scale, fp8[1],
            )
            _FN_REDUCE(
                logits, attn_lse,
                reduce_indptr, reduce_final_map, reduce_partial_map,
                qs, output, None,
            )
            return output

        _SHAPE_DISPATCH[key] = _dispatch_persistent
        return _dispatch_persistent

    elif fp8_splits is not None:
        # FP8 multi-split + C++ reduce (S4, S6)
        total_q = q.shape[0]
        num_splits = fp8_splits
        splits_indptr = torch.arange(0, (bs + 1) * num_splits, num_splits, dtype=torch.int, device=device)
        kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
        np_logits = torch.empty((total_q, num_splits, nh, V_HEAD_DIM), dtype=torch.float32, device=device)
        np_lse = torch.empty((total_q, num_splits, nh, 1), dtype=torch.float32, device=device)

        np_total = total_q * num_splits
        reduce_indptr_fp8 = torch.arange(0, np_total + num_splits, num_splits, dtype=torch.int32, device=device)[:total_q + 1]
        reduce_partial_map_fp8 = torch.arange(np_total, dtype=torch.int32, device=device)
        rl_view = np_logits.reshape(np_total, 1, nh, V_HEAD_DIM)
        rls_view = np_lse.reshape(np_total, 1, nh, 1)

        def _dispatch_fp8(data):
            q_fp8_buf.copy_(data[0].view(q_flat_shape))
            fp8 = data[1]["fp8"]
            kv_v = fp8[0].view(kv_view_shape)
            _FN_STAGE1(
                q_shaped, kv_v,
                qo_indptr, kv_indptr, kv_indices,
                kv_last, splits_indptr, None, None, None,
                qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
                np_logits, np_lse, output, q_scale, fp8[1],
            )
            _FN_REDUCE(rl_view, rls_view, reduce_indptr_fp8, None, reduce_partial_map_fp8, qs, output, None)
            return output

        _SHAPE_DISPATCH[key] = _dispatch_fp8
        return _dispatch_fp8

    elif is_np1:
        # NP1: 1-split, no reduce needed, stage1 writes bf16 directly to output
        total_q = q.shape[0]
        indptr = torch.arange(0, bs + 1, dtype=torch.int, device=device)
        kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
        logits = output.view(total_q, 1, nh, V_HEAD_DIM)
        attn_lse = torch.empty((total_q, 1, nh, 1), dtype=torch.float32, device=device)

        def _dispatch_np1(data):
            q_fp8_buf.copy_(data[0].view(q_flat_shape))
            fp8 = data[1]["fp8"]
            kv_v = fp8[0].view(kv_view_shape)
            _FN_STAGE1(
                q_shaped, kv_v,
                qo_indptr, kv_indptr, kv_indices,
                kv_last, indptr, None, None, None,
                qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
                logits, attn_lse, output, q_scale, fp8[1],
            )
            return output

        _SHAPE_DISPATCH[key] = _dispatch_np1
        return _dispatch_np1

    else:
        raise RuntimeError(f"Shape {key} not in dispatch table")


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    q = data[0]
    if not q.is_cuda:
        return _torch_fallback(data)

    config = data[4]
    key = (int(config["batch_size"]), int(config["kv_seq_len"]))

    dispatch = _SHAPE_DISPATCH.get(key)
    if dispatch is not None:
        return dispatch(data)

    dispatch = _build_dispatch(key, q, data)
    return dispatch(data)
scrolls · 327 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