Skip to content
KernelIndex
Search⌘K

submission 750700

pongtsu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-750700?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
9.41µs
#182 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:abf7931b5aaf885ac987cfc6d44094462852513acffa694af8d89e41621ec1ea
license declaredunknown
license concludedunknown
authorspongtsu
imported2026-08-15

Techniques

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

num-warps = 4num_warps=4,
split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
stages = 2num_stages=2,
tile-k = 256BLOCK_SIZE_K=256,
tile-m = 8BLOCK_SIZE_M=8,
tile-n = 128BLOCK_SIZE_N=128,

Kernel source

submission.py432 lines
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
    _gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk

from task import input_t, output_t

_buffers = {}
_ASM_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_FUSED_M_THRESHOLD = 64


def _get_fused_config(M, N, K):
    """Our own configs — keep BSK=256 (BSK=512 caused regression), try other tweaks."""
    if K > 4096:
        # K=7168: try splitK=8
        return dict(
            BLOCK_SIZE_M=8,
            BLOCK_SIZE_N=128,
            BLOCK_SIZE_K=256,
            GROUP_SIZE_M=1,
            num_warps=4,
            num_stages=2,
            waves_per_eu=2,
            matrix_instr_nonkdim=16,
            cache_modifier=".cg",
            NUM_KSPLIT=8,
        )
    if M <= 4:
        # Try waves_per_eu=2 (vs 0), cache_modifier=None (vs .cg)
        return dict(
            BLOCK_SIZE_M=4,
            BLOCK_SIZE_N=128,
            BLOCK_SIZE_K=256,
            GROUP_SIZE_M=1,
            num_warps=4,
            num_stages=2,
            waves_per_eu=2,
            matrix_instr_nonkdim=16,
            cache_modifier=None,
            NUM_KSPLIT=1,
        )
    elif M <= 8:
        return dict(
            BLOCK_SIZE_M=8,
            BLOCK_SIZE_N=128,
            BLOCK_SIZE_K=256,
            GROUP_SIZE_M=1,
            num_warps=4,
            num_stages=2,
            waves_per_eu=2,
            matrix_instr_nonkdim=16,
            cache_modifier=None,
            NUM_KSPLIT=1,
        )
    elif M <= 32 and K <= 1024:
        return dict(
            BLOCK_SIZE_M=8,
            BLOCK_SIZE_N=128,
            BLOCK_SIZE_K=256,
            GROUP_SIZE_M=1,
            num_warps=4,
            num_stages=2,
            waves_per_eu=2,
            matrix_instr_nonkdim=16,
            cache_modifier=None,
            NUM_KSPLIT=1,
        )
    elif M <= 32:
        return dict(
            BLOCK_SIZE_M=32,
            BLOCK_SIZE_N=64,
            BLOCK_SIZE_K=512,
            GROUP_SIZE_M=1,
            num_warps=8,
            num_stages=1,
            waves_per_eu=2,
            matrix_instr_nonkdim=16,
            cache_modifier=None,
            NUM_KSPLIT=1,
        )
    else:
        # M=64: try waves_per_eu=4 for more latency hiding
        return dict(
            BLOCK_SIZE_M=16,
            BLOCK_SIZE_N=128,
            BLOCK_SIZE_K=256,
            GROUP_SIZE_M=1,
            num_warps=4,
            num_stages=2,
            waves_per_eu=4,
            matrix_instr_nonkdim=16,
            cache_modifier=".cg",
            NUM_KSPLIT=1,
        )


@triton.heuristics(
    {
        "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
        and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
    }
)
@triton.jit
def _fused_mxfp4_quant_shuffle_kernel(
    x_ptr,
    x_fp4_ptr,
    bs_ptr,
    stride_x_m_in,
    stride_x_n_in,
    stride_x_fp4_m_in,
    stride_x_fp4_n_in,
    M,
    N,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    EVEN_M_N: tl.constexpr,
    SCALING_MODE: tl.constexpr,
    SCALE_N_PAD: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    stride_x_m = tl.cast(stride_x_m_in, tl.int64)
    stride_x_n = tl.cast(stride_x_n_in, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
        if EVEN_M_N:
            x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
        else:
            x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
            x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(
                tl.float32
            )

        out_tensor, bs_e8m0 = _mxfp4_quant_op(
            x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
        )

        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = (
            out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
        )
        if EVEN_M_N:
            tl.store(x_fp4_ptr + out_offs, out_tensor, cache_modifier=".wt")
        else:
            out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
            tl.store(
                x_fp4_ptr + out_offs, out_tensor, mask=out_mask, cache_modifier=".wt"
            )

        # Inline E8M0 scale shuffle (matches e8m0_shuffle permutation)
        bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
        bs_offs_0 = bs_offs_m[:, None] // 32
        bs_offs_1 = bs_offs_m[:, None] % 32
        bs_offs_2 = bs_offs_1 % 16
        bs_offs_1 = bs_offs_1 // 16
        bs_offs_3 = bs_offs_n[None, :] // 8
        bs_offs_4 = bs_offs_n[None, :] % 8
        bs_offs_5 = bs_offs_4 % 4
        bs_offs_4 = bs_offs_4 // 4
        bs_offs = (
            bs_offs_1
            + bs_offs_4 * 2
            + bs_offs_2 * 2 * 2
            + bs_offs_5 * 2 * 2 * 16
            + bs_offs_3 * 2 * 2 * 16 * 4
            + bs_offs_0 * 2 * 16 * SCALE_N_PAD
        )
        bs_mask_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
        bs_e8m0 = tl.where(bs_mask_valid, bs_e8m0, 127)
        SCALE_M_PAD = (M + 255) // 256 * 256
        bs_mask = (bs_offs_m < SCALE_M_PAD)[:, None] & (bs_offs_n < SCALE_N_PAD)[
            None, :
        ]
        tl.store(
            bs_ptr + bs_offs, bs_e8m0.to(tl.uint8), mask=bs_mask, cache_modifier=".cg"
        )


def _get_or_create_buffers(M, K, N, device):
    key = (M, K, N)
    if key not in _buffers:
        if M <= _FUSED_M_THRESHOLD:
            config = _get_fused_config(M, N, K)
            K_kernel = K // 2
            BSK = config["BLOCK_SIZE_K"]
            BSN = max(config["BLOCK_SIZE_N"], 32)
            BSM = config["BLOCK_SIZE_M"]

            if config["NUM_KSPLIT"] > 1:
                SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(
                    K_kernel, BSK, config["NUM_KSPLIT"]
                )
                grid_size = NUM_KSPLIT * triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
                y_pp = torch.empty(
                    (NUM_KSPLIT, M, N), dtype=torch.float32, device=device
                )
                ACTUAL_KSPLIT = triton.cdiv(K_kernel, (SPLITK_BLOCK_SIZE // 2))
                _buffers[key] = {
                    "mode": "fused_splitk",
                    "out": torch.empty((M, N), dtype=torch.bfloat16, device=device),
                    "B_w": None,
                    "B_sc": None,
                    "grid_size": grid_size,
                    "K_kernel": K_kernel,
                    "y_pp": y_pp,
                    "BLOCK_SIZE_M": BSM,
                    "BLOCK_SIZE_N": BSN,
                    "BLOCK_SIZE_K": BSK,
                    "SPLITK_BLOCK_SIZE": SPLITK_BLOCK_SIZE,
                    "GROUP_SIZE_M": config["GROUP_SIZE_M"],
                    "NUM_KSPLIT": NUM_KSPLIT,
                    "num_warps": config["num_warps"],
                    "num_stages": config["num_stages"],
                    "waves_per_eu": config["waves_per_eu"],
                    "matrix_instr_nonkdim": config["matrix_instr_nonkdim"],
                    "cache_modifier": config["cache_modifier"],
                    "ACTUAL_KSPLIT": ACTUAL_KSPLIT,
                    "MAX_KSPLIT": triton.next_power_of_2(NUM_KSPLIT),
                    "reduce_grid": (triton.cdiv(M, 16), triton.cdiv(N, 64)),
                }
            else:
                SPLITK_BLOCK_SIZE = 2 * K_kernel
                grid_size = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
                _buffers[key] = {
                    "mode": "fused_direct",
                    "out": torch.empty((M, N), dtype=torch.bfloat16, device=device),
                    "B_w": None,
                    "B_sc": None,
                    "grid_size": grid_size,
                    "K_kernel": K_kernel,
                    "BLOCK_SIZE_M": BSM,
                    "BLOCK_SIZE_N": BSN,
                    "BLOCK_SIZE_K": BSK,
                    "SPLITK_BLOCK_SIZE": SPLITK_BLOCK_SIZE,
                    "GROUP_SIZE_M": config["GROUP_SIZE_M"],
                    "NUM_KSPLIT": 1,
                    "num_warps": config["num_warps"],
                    "num_stages": config["num_stages"],
                    "waves_per_eu": config["waves_per_eu"],
                    "matrix_instr_nonkdim": config["matrix_instr_nonkdim"],
                    "cache_modifier": config["cache_modifier"],
                }
        else:
            MXFP4_QUANT_BLOCK_SIZE = 32
            SCALE_N_valid = triton.cdiv(K, MXFP4_QUANT_BLOCK_SIZE)
            SCALE_M = triton.cdiv(M, 256) * 256
            SCALE_N = triton.cdiv(SCALE_N_valid, 8) * 8
            BLOCK_SIZE_M = triton.cdiv(min(32, triton.next_power_of_2(M)), 32) * 32
            BLOCK_SIZE_N = 64
            grid = (triton.cdiv(M, BLOCK_SIZE_M), triton.cdiv(K, BLOCK_SIZE_N))
            padded_M = (M + 31) // 32 * 32
            _buffers[key] = {
                "mode": "two_phase",
                "x_fp4": torch.empty((M, K // 2), dtype=torch.uint8, device=device),
                "blockscale": torch.empty(
                    (SCALE_M, SCALE_N), dtype=torch.uint8, device=device
                ),
                "gemm_out": torch.empty(
                    (padded_M, N), dtype=torch.bfloat16, device=device
                ),
                "SCALE_N": SCALE_N,
                "BLOCK_SIZE_M": BLOCK_SIZE_M,
                "BLOCK_SIZE_N": BLOCK_SIZE_N,
                "grid": grid,
                "M": M,
            }
    return _buffers[key]


def custom_kernel(data: input_t) -> output_t:
    A, _, _, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B_shuffle.shape[0]

    buf = _get_or_create_buffers(M, K, N, A.device)

    # Lazy reshape B weights and scales (only on first call or if B changes)
    if buf.get("mode") in ("fused_splitk", "fused_direct"):
        b_ptr = B_shuffle.data_ptr()
        if buf["B_w"] is None or buf.get("_b_ptr") != b_ptr:
            buf["B_w"] = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
            bs_shape = B_scale_sh.shape
            buf["B_sc"] = B_scale_sh.view(torch.uint8).reshape(
                bs_shape[0] // 32, bs_shape[1] * 32
            )
            buf["_b_ptr"] = b_ptr
            # Recover actual N from preshuffle layout
            actual_N = buf["B_w"].shape[0] * 16
            if actual_N != N:
                N = actual_N

    if buf["mode"] == "fused_splitk":
        out = buf["out"]
        y_pp = buf["y_pp"]
        _gemm_a16wfp4_preshuffle_kernel[(buf["grid_size"],)](
            A,
            buf["B_w"],
            y_pp,
            buf["B_sc"],
            M,
            N,
            buf["K_kernel"],
            A.stride(0),
            A.stride(1),
            buf["B_w"].stride(0),
            buf["B_w"].stride(1),
            y_pp.stride(0),
            y_pp.stride(1),
            y_pp.stride(2),
            buf["B_sc"].stride(0),
            buf["B_sc"].stride(1),
            BLOCK_SIZE_M=buf["BLOCK_SIZE_M"],
            BLOCK_SIZE_N=buf["BLOCK_SIZE_N"],
            BLOCK_SIZE_K=buf["BLOCK_SIZE_K"],
            GROUP_SIZE_M=buf["GROUP_SIZE_M"],
            NUM_KSPLIT=buf["NUM_KSPLIT"],
            SPLITK_BLOCK_SIZE=buf["SPLITK_BLOCK_SIZE"],
            num_warps=buf["num_warps"],
            num_stages=buf["num_stages"],
            waves_per_eu=buf["waves_per_eu"],
            matrix_instr_nonkdim=buf["matrix_instr_nonkdim"],
            PREQUANT=True,
            cache_modifier=buf["cache_modifier"],
        )
        _gemm_afp4wfp4_reduce_kernel[buf["reduce_grid"]](
            y_pp,
            out,
            M,
            N,
            y_pp.stride(0),
            y_pp.stride(1),
            y_pp.stride(2),
            out.stride(0),
            out.stride(1),
            16,
            64,
            buf["ACTUAL_KSPLIT"],
            buf["MAX_KSPLIT"],
        )
        return out

    elif buf["mode"] == "fused_direct":
        out = buf["out"]
        _gemm_a16wfp4_preshuffle_kernel[(buf["grid_size"],)](
            A,
            buf["B_w"],
            out,
            buf["B_sc"],
            M,
            N,
            buf["K_kernel"],
            A.stride(0),
            A.stride(1),
            buf["B_w"].stride(0),
            buf["B_w"].stride(1),
            0,
            out.stride(0),
            out.stride(1),
            buf["B_sc"].stride(0),
            buf["B_sc"].stride(1),
            BLOCK_SIZE_M=buf["BLOCK_SIZE_M"],
            BLOCK_SIZE_N=buf["BLOCK_SIZE_N"],
            BLOCK_SIZE_K=buf["BLOCK_SIZE_K"],
            GROUP_SIZE_M=buf["GROUP_SIZE_M"],
            NUM_KSPLIT=buf["NUM_KSPLIT"],
            SPLITK_BLOCK_SIZE=buf["SPLITK_BLOCK_SIZE"],
            num_warps=buf["num_warps"],
            num_stages=buf["num_stages"],
            waves_per_eu=buf["waves_per_eu"],
            matrix_instr_nonkdim=buf["matrix_instr_nonkdim"],
            PREQUANT=True,
            cache_modifier=buf["cache_modifier"],
        )
        return out

    else:  # two_phase for M > 64
        _fused_mxfp4_quant_shuffle_kernel[buf["grid"]](
            A,
            buf["x_fp4"],
            buf["blockscale"],
            *A.stride(),
            *buf["x_fp4"].stride(),
            M=M,
            N=K,
            BLOCK_SIZE_M=buf["BLOCK_SIZE_M"],
            BLOCK_SIZE_N=buf["BLOCK_SIZE_N"],
            NUM_ITER=1,
            NUM_STAGES=1,
            MXFP4_QUANT_BLOCK_SIZE=32,
            SCALING_MODE=0,
            SCALE_N_PAD=buf["SCALE_N"],
            num_warps=2,
            waves_per_eu=0,
            num_stages=1,
        )
        gemm_a4w4_asm(
            buf["x_fp4"].view(dtypes.fp4x2),
            B_shuffle,
            buf["blockscale"].view(dtypes.fp8_e8m0),
            B_scale_sh,
            buf["gemm_out"],
            _ASM_KERNEL_32x128,
            None,
            1.0,
            0.0,
            True,
            log2_k_split=0,
        )
        return buf["gemm_out"][:M]
scrolls · 432 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