Skip to content
KernelIndex
Search⌘K

submission 586303

shiyeegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-586303?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
22.1µs
#743 of 1143
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:005e082c8f95b30a505aa1d0efde4831ca69268757bc78b87b738080b1e642b7
license declaredunknown
license concludedunknown
authorsshiyeegao
imported2026-08-26

Techniques

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

num-warps = 4num_warps=4,
split-ksplitk_block_size = (k + num_ksplit - 1) // num_ksplit
stages = 1num_stages=1,

Kernel source

submission.py512 lines
#!POPCORN leaderboard amd-mxfp4-mm

from __future__ import annotations

from typing import Any


_RUNTIME: dict[str, Any] | None = None

_DEFAULT_CONFIGS = (
    (8, 8, 64, 256, 4, 2, 2, 1, 1, 0),
    (31, 16, 64, 256, 4, 2, 2, 1, 1, 0),
    (32, 32, 64, 256, 4, 2, 2, 1, 1, 0),
    (64, 32, 64, 256, 4, 2, 2, 1, 1, 0),
    (128, 32, 64, 256, 4, 2, 2, 1, 1, 0),
    (256, 32, 64, 256, 4, 2, 2, 1, 1, 0),
    (1 << 30, 32, 64, 256, 4, 2, 2, 1, 1, 0),
)

_SPECIAL_CONFIGS = {
    (2112, 7168): (
        (8, 8, 32, 512, 1, 2, 2, 1, 7, 2),
        (31, 16, 32, 512, 1, 4, 1, 4, 14, 2),
        (32, 32, 128, 512, 1, 4, 1, 1, 14, 2),
        (64, 64, 32, 512, 1, 4, 2, 1, 7, 2),
        (128, 64, 32, 1024, 1, 2, 2, 1, 1, 0),
        (256, 64, 32, 1024, 1, 2, 2, 2, 1, 0),
        (1 << 30, 256, 64, 256, 1, 2, 2, 1, 1, 0),
    ),
    (3072, 1536): (
        (8, 8, 32, 512, 1, 4, 2, 1, 1, 2),
        (31, 8, 32, 512, 1, 4, 2, 1, 1, 0),
        (32, 32, 32, 512, 1, 4, 2, 1, 1, 2),
        (64, 64, 32, 512, 1, 2, 2, 1, 1, 2),
        (128, 64, 32, 512, 1, 2, 2, 1, 1, 0),
        (256, 128, 32, 512, 1, 4, 2, 1, 1, 0),
        (1 << 30, 128, 64, 256, 8, 2, 2, 2, 1, 0),
    ),
    (4096, 512): (
        (8, 8, 64, 512, 1, 2, 1, 1, 1, 0),
        (31, 8, 64, 512, 1, 2, 1, 1, 1, 0),
        (32, 32, 64, 512, 1, 4, 1, 1, 1, 0),
        (64, 32, 128, 512, 1, 4, 1, 1, 1, 0),
        (128, 32, 128, 512, 1, 4, 1, 1, 1, 0),
        (256, 256, 256, 512, 1, 4, 1, 1, 1, 0),
        (1 << 30, 64, 256, 512, 1, 4, 1, 1, 1, 0),
    ),
}


def _floor_power_of_two(x: int) -> int:
    if x <= 1:
        return 1
    return 1 << (x.bit_length() - 1)


def _pick_config(m: int, n: int, k: int) -> tuple[int, int, int, int, int, int, int, int, bool, int]:
    table = _SPECIAL_CONFIGS.get((n, k), _DEFAULT_CONFIGS)
    chosen = table[-1]
    for entry in table:
        if m <= entry[0]:
            chosen = entry
            break

    (
        _max_m,
        block_m,
        block_n,
        block_k,
        group_m,
        num_warps,
        num_stages,
        waves_per_eu,
        num_ksplit,
        cache_hint,
    ) = chosen

    splitk_block_size = (k + num_ksplit - 1) // num_ksplit
    if block_k > splitk_block_size:
        block_k = _floor_power_of_two(splitk_block_size)
    block_k = max(block_k, 64)

    even_k = (
        (k % block_k == 0)
        and (splitk_block_size % block_k == 0)
        and (k % splitk_block_size == 0)
    )
    return (
        block_m,
        block_n,
        block_k,
        group_m,
        num_warps,
        num_stages,
        waves_per_eu,
        num_ksplit,
        even_k,
        cache_hint,
    )


def _load_runtime() -> dict[str, Any]:
    global _RUNTIME
    if _RUNTIME is not None:
        return _RUNTIME

    import torch
    import triton
    import triton.language as tl
    from aiter.ops.triton.quant import dynamic_mxfp4_quant

    @triton.jit
    def _remap_xcd(pid, grid_mn, NUM_XCDS: tl.constexpr = 8):
        pids_per_xcd = (grid_mn + NUM_XCDS - 1) // NUM_XCDS
        tall_xcds = grid_mn % NUM_XCDS
        tall_xcds = tl.where(tall_xcds == 0, NUM_XCDS, tall_xcds)
        xcd = pid % NUM_XCDS
        local_pid = pid // NUM_XCDS
        tall_pid = xcd * pids_per_xcd + local_pid
        short_pid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid
        return tl.where(xcd < tall_xcds, tall_pid, short_pid)

    @triton.jit
    def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):
        if GROUP_SIZE_M == 1:
            return pid // num_pid_n, pid % num_pid_n
        num_pid_in_group = GROUP_SIZE_M * num_pid_n
        group_id = pid // num_pid_in_group
        first_pid_m = group_id * GROUP_SIZE_M
        group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
        tl.assume(group_size_m >= 0)
        pid_in_group = pid % num_pid_in_group
        pid_m = first_pid_m + (pid_in_group % group_size_m)
        pid_n = pid_in_group // group_size_m
        return pid_m, pid_n

    @triton.jit
    def _load_preshuffle_weight_full(ptrs, CACHE_HINT: tl.constexpr):
        if CACHE_HINT == 2:
            return tl.load(ptrs, cache_modifier=".cg")
        if CACHE_HINT == 1:
            return tl.load(ptrs, cache_modifier=".ca")
        return tl.load(ptrs)

    @triton.jit
    def _load_preshuffle_weight_masked(ptrs, mask, CACHE_HINT: tl.constexpr):
        if CACHE_HINT == 2:
            return tl.load(ptrs, mask=mask, other=0, cache_modifier=".cg")
        if CACHE_HINT == 1:
            return tl.load(ptrs, mask=mask, other=0, cache_modifier=".ca")
        return tl.load(ptrs, mask=mask, other=0)

    @triton.jit
    def _load_preshuffle_scale(ptrs, CACHE_HINT: tl.constexpr):
        if CACHE_HINT == 2:
            return tl.load(ptrs, cache_modifier=".cg")
        if CACHE_HINT == 1:
            return tl.load(ptrs, cache_modifier=".ca")
        return tl.load(ptrs)

    @triton.jit
    def _mxfp4_preshuffle_kernel(
        a_ptr,
        b_ptr,
        c_ptr,
        a_scales_ptr,
        b_scales_ptr,
        m,
        n,
        k_packed,
        b_rows,
        a_scale_rows,
        b_scale_rows,
        stride_am,
        stride_ak,
        stride_bn,
        stride_bk,
        stride_ck,
        stride_cm,
        stride_cn,
        stride_asm,
        stride_ask,
        stride_bsn,
        stride_bsk,
        BLOCK_M: tl.constexpr,
        BLOCK_N: tl.constexpr,
        BLOCK_K: tl.constexpr,
        GROUP_SIZE_M: tl.constexpr,
        SMALL_M: tl.constexpr,
        NUM_KSPLIT: tl.constexpr,
        SPLITK_BLOCK_SIZE: tl.constexpr,
        EVEN_K: tl.constexpr,
        CACHE_HINT: tl.constexpr,
    ):
        tl.assume(stride_am > 0)
        tl.assume(stride_ak > 0)
        tl.assume(stride_bn > 0)
        tl.assume(stride_bk > 0)
        tl.assume(stride_ck >= 0)
        tl.assume(stride_cm > 0)
        tl.assume(stride_cn > 0)
        tl.assume(stride_asm > 0)
        tl.assume(stride_ask > 0)
        tl.assume(stride_bsn > 0)
        tl.assume(stride_bsk > 0)

        grid_m = tl.cdiv(m, BLOCK_M)
        grid_n = tl.cdiv(n, BLOCK_N)
        grid_mn = grid_m * grid_n

        pid_unified = tl.program_id(axis=0)
        pid_unified = _remap_xcd(pid_unified, grid_mn * NUM_KSPLIT, NUM_XCDS=8)
        pid_k = pid_unified % NUM_KSPLIT
        pid = pid_unified // NUM_KSPLIT

        if NUM_KSPLIT == 1:
            pid_m, pid_n = _pid_grid(pid, grid_m, grid_n, GROUP_SIZE_M=GROUP_SIZE_M)
        else:
            pid_m = pid // grid_n
            pid_n = pid % grid_n

        splitk_packed = SPLITK_BLOCK_SIZE // 2
        k_start = pid_k * splitk_packed
        if k_start < k_packed:
            scale_group: tl.constexpr = 32
            offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
            offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
            offs_am = offs_m % m
            offs_bn = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % b_rows
            offs_bsn = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % b_scale_rows

            acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
            num_k_iter = tl.cdiv(splitk_packed, BLOCK_K // 2)
            offs_k = k_start + tl.arange(0, BLOCK_K // 2)
            offs_k_shuffle = k_start * 16 + tl.arange(0, (BLOCK_K // 2) * 16)

            a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak
            b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk

            if SMALL_M:
                a_scale_k = pid_k * (SPLITK_BLOCK_SIZE // scale_group) + tl.arange(
                    0, BLOCK_K // scale_group
                )
                a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + a_scale_k[None, :] * stride_ask
            else:
                offs_asm = (pid_m * (BLOCK_M // 32) + tl.arange(0, BLOCK_M // 32)) % a_scale_rows
                a_scale_k = pid_k * (SPLITK_BLOCK_SIZE // scale_group) * 32 + tl.arange(
                    0, (BLOCK_K // scale_group) * 32
                )
                a_scale_ptrs = a_scales_ptr + offs_asm[:, None] * stride_asm + a_scale_k[None, :] * stride_ask

            b_scale_k = pid_k * (SPLITK_BLOCK_SIZE // scale_group) * 32 + tl.arange(
                0, (BLOCK_K // scale_group) * 32
            )
            b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + b_scale_k[None, :] * stride_bsk

            for k_iter in range(0, num_k_iter):
                if EVEN_K:
                    a = tl.load(a_ptrs)
                    b = _load_preshuffle_weight_full(b_ptrs, CACHE_HINT=CACHE_HINT)
                else:
                    rem_k = k_packed - (k_start + k_iter * (BLOCK_K // 2))
                    a = tl.load(a_ptrs, mask=tl.arange(0, BLOCK_K // 2)[None, :] < rem_k, other=0)
                    b = _load_preshuffle_weight_masked(
                        b_ptrs,
                        (tl.arange(0, (BLOCK_K // 2) * 16)[None, :] // 16) < rem_k,
                        CACHE_HINT=CACHE_HINT,
                    )

                b = (
                    b.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
                    .permute(0, 1, 4, 2, 3, 5)
                    .reshape(BLOCK_N, BLOCK_K // 2)
                    .trans(1, 0)
                )

                if SMALL_M:
                    a_scales = tl.load(a_scale_ptrs)
                else:
                    a_scales = (
                        tl.load(a_scale_ptrs)
                        .reshape(BLOCK_M // 32, BLOCK_K // scale_group // 8, 4, 16, 2, 2, 1)
                        .permute(0, 5, 3, 1, 4, 2, 6)
                        .reshape(BLOCK_M, BLOCK_K // scale_group)
                    )

                b_scales = (
                    _load_preshuffle_scale(b_scale_ptrs, CACHE_HINT=CACHE_HINT)
                    .reshape(BLOCK_N // 32, BLOCK_K // scale_group // 8, 4, 16, 2, 2, 1)
                    .permute(0, 5, 3, 1, 4, 2, 6)
                    .reshape(BLOCK_N, BLOCK_K // scale_group)
                )

                acc = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc)

                a_ptrs += (BLOCK_K // 2) * stride_ak
                b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
                if SMALL_M:
                    a_scale_ptrs += (BLOCK_K // scale_group) * stride_ask
                else:
                    a_scale_ptrs += BLOCK_K * stride_ask
                b_scale_ptrs += BLOCK_K * stride_bsk

            c = acc.to(c_ptr.type.element_ty)
            c_ptrs = c_ptr + pid_k * stride_ck + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
            mask = (offs_m[:, None] < m) & (offs_n[None, :] < n)
            tl.store(c_ptrs, c, mask=mask, cache_modifier=".wt")

    @triton.jit
    def _splitk_reduce_kernel(
        src_ptr,
        dst_ptr,
        m,
        n,
        actual_ksplit,
        stride_sk,
        stride_sm,
        stride_sn,
        stride_dm,
        stride_dn,
        BLOCK_M: tl.constexpr,
        BLOCK_N: tl.constexpr,
        MAX_KSPLIT: tl.constexpr,
    ):
        pid_m = tl.program_id(axis=0)
        pid_n = tl.program_id(axis=1)

        offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
        offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
        offs_k = tl.arange(0, MAX_KSPLIT)

        src_ptrs = (
            src_ptr
            + offs_k[:, None, None] * stride_sk
            + offs_m[None, :, None] * stride_sm
            + offs_n[None, None, :] * stride_sn
        )
        src_mask = (
            (offs_k[:, None, None] < actual_ksplit)
            & (offs_m[None, :, None] < m)
            & (offs_n[None, None, :] < n)
        )
        vals = tl.load(src_ptrs, mask=src_mask, other=0.0)
        out = tl.sum(vals, axis=0).to(dst_ptr.type.element_ty)

        dst_ptrs = dst_ptr + offs_m[:, None] * stride_dm + offs_n[None, :] * stride_dn
        tl.store(dst_ptrs, out, mask=(offs_m[:, None] < m) & (offs_n[None, :] < n))

    def _as_u8(x: torch.Tensor) -> torch.Tensor:
        y = x.contiguous()
        if y.dtype != torch.uint8:
            y = y.view(torch.uint8)
        return y

    def _pad_scale_cols(x: torch.Tensor) -> torch.Tensor:
        y = _as_u8(x)
        rows, cols = y.shape
        cols_pad = ((cols + 7) // 8) * 8
        if cols_pad == cols:
            return y
        out = torch.zeros((rows, cols_pad), dtype=torch.uint8, device=y.device)
        out[:, :cols] = y
        return out

    def _shuffle_e8m0(x: torch.Tensor) -> torch.Tensor:
        y = _pad_scale_cols(x)
        rows, cols = y.shape
        rows_pad = ((rows + 255) // 256) * 256
        cols_pad = ((cols + 7) // 8) * 8
        out = torch.zeros((rows_pad, cols_pad), dtype=torch.uint8, device=y.device)
        out[:rows, :cols] = y
        return (
            out.view(rows_pad // 32, 2, 16, cols_pad // 8, 2, 4)
            .permute(0, 3, 5, 2, 4, 1)
            .contiguous()
            .view(rows_pad, cols_pad)
        )

    def _view_preshuffle_weight(x: torch.Tensor) -> torch.Tensor:
        y = _as_u8(x)
        rows, cols = y.shape
        return y.view(rows // 16, cols * 16)

    def _view_preshuffle_scale(x: torch.Tensor) -> torch.Tensor:
        y = _as_u8(x)
        rows, cols = y.shape
        return y.view(rows // 32, cols * 32)

    def _prepare_a_scale(x: torch.Tensor, m: int) -> tuple[torch.Tensor, bool]:
        if m < 32:
            return _pad_scale_cols(x), True
        return _view_preshuffle_scale(_shuffle_e8m0(x)), False

    _RUNTIME = {
        "torch": torch,
        "triton": triton,
        "dynamic_mxfp4_quant": dynamic_mxfp4_quant,
        "kernel": _mxfp4_preshuffle_kernel,
        "reduce_kernel": _splitk_reduce_kernel,
        "as_u8": _as_u8,
        "view_preshuffle_weight": _view_preshuffle_weight,
        "view_preshuffle_scale": _view_preshuffle_scale,
        "prepare_a_scale": _prepare_a_scale,
    }
    return _RUNTIME


def custom_kernel(data: tuple[Any, ...]) -> Any:
    rt = _load_runtime()
    torch = rt["torch"]
    triton = rt["triton"]

    a, _b, _b_q, b_shuffle, b_scale_sh = data
    if a.ndim != 2:
        raise RuntimeError(f"A must be 2D, got {tuple(a.shape)}")

    m, k = a.shape
    n = b_shuffle.shape[0]

    a_q, a_scale = rt["dynamic_mxfp4_quant"](a.contiguous())
    a_q_u8 = rt["as_u8"](a_q)
    a_scale_view, small_m = rt["prepare_a_scale"](a_scale, m)
    b_view = rt["view_preshuffle_weight"](b_shuffle)
    b_scale_view = rt["view_preshuffle_scale"](b_scale_sh)

    (
        block_m,
        block_n,
        block_k,
        group_m,
        num_warps,
        num_stages,
        waves_per_eu,
        num_ksplit,
        even_k,
        cache_hint,
    ) = _pick_config(m, n, k)

    grid = (triton.cdiv(m, block_m) * triton.cdiv(n, block_n) * num_ksplit,)
    if num_ksplit == 1:
        out = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
        partial = out
        stride_ck = 0
    else:
        partial = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=a.device)
        out = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
        stride_ck = partial.stride(0)

    rt["kernel"][grid](
        a_q_u8,
        b_view,
        partial,
        a_scale_view,
        b_scale_view,
        m,
        n,
        a_q_u8.shape[1],
        b_view.shape[0],
        a_scale_view.shape[0],
        b_scale_view.shape[0],
        a_q_u8.stride(0),
        a_q_u8.stride(1),
        b_view.stride(0),
        b_view.stride(1),
        stride_ck,
        partial.stride(-2),
        partial.stride(-1),
        a_scale_view.stride(0),
        a_scale_view.stride(1),
        b_scale_view.stride(0),
        b_scale_view.stride(1),
        BLOCK_M=block_m,
        BLOCK_N=block_n,
        BLOCK_K=block_k,
        GROUP_SIZE_M=group_m,
        SMALL_M=small_m,
        NUM_KSPLIT=num_ksplit,
        SPLITK_BLOCK_SIZE=(k + num_ksplit - 1) // num_ksplit,
        EVEN_K=even_k,
        CACHE_HINT=cache_hint,
        num_warps=num_warps,
        num_stages=num_stages,
        waves_per_eu=waves_per_eu,
        matrix_instr_nonkdim=16,
    )

    if num_ksplit == 1:
        return out

    reduce_block_m = 16
    reduce_block_n = 64
    reduce_grid = (triton.cdiv(m, reduce_block_m), triton.cdiv(n, reduce_block_n))
    max_ksplit = 1 << (num_ksplit - 1).bit_length()
    rt["reduce_kernel"][reduce_grid](
        partial,
        out,
        m,
        n,
        num_ksplit,
        partial.stride(0),
        partial.stride(1),
        partial.stride(2),
        out.stride(0),
        out.stride(1),
        BLOCK_M=reduce_block_m,
        BLOCK_N=reduce_block_n,
        MAX_KSPLIT=max_ksplit,
        num_warps=4,
        num_stages=1,
    )
    return out
scrolls · 512 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