Skip to content
KernelIndex
Search⌘K

submission 686644

zwang86 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:45bb2a614e64533f61c85fb1b1b24320f2a1fb6dedc1bca7574fc2ee9b18d8c0
license declaredunknown
license concludedunknown
authorszwang86
imported2026-08-26

Techniques

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

num-warps = 1NUM_ITER, NUM_STAGES, NUM_WARPS = 1, 1, 4
split-k_SK = 15 # splitK
stages = 1NUM_WARPS, NUM_STAGES = 1, 1
tile-n = 1NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N = 1, triton.next_power_of_2(M), 32

Kernel source

submission.py278 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from task import input_t, output_t

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,
    gemm_a4w4_blockscale,
    get_GEMM_config,
)


@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_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,
    scaleN_pad,
    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,
):
    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)
        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)

        bs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)

        d0 = bs_m // 32
        d1 = (bs_m % 32) // 16
        d2 = bs_m % 16
        d3 = bs_n // 8
        d4 = (bs_n % 8) // 4
        d5 = bs_n % 4

        bs_offs = (
            d1[:, None]
            + d4[None, :] * 2
            + d2[:, None] * 4
            + d5[None, :] * 64
            + d3[None, :] * 256
            + d0[:, None] * 32 * scaleN_pad
        )

        if EVEN_M_N:
            tl.store(bs_ptr + bs_offs, bs_e8m0)
        else:
            scaleN_valid = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
            bs_mask = (bs_m[:, None] < M) & (bs_n[None, :] < scaleN_valid)
            bs_val = tl.where(bs_mask, bs_e8m0, 127)
            tl.store(bs_ptr + bs_offs, bs_val, mask=bs_mask)


_cfg = {}

# Config tuple indices for hot-path access
_GRID = 0
_A_FP4 = 1
_BS = 2
_SNP = 3  # scaleN_pad
_NI = 4   # NUM_ITER
_BSM = 5  # BLOCK_SIZE_M
_BSN = 6  # BLOCK_SIZE_N
_NS = 7   # NUM_STAGES
_NW = 8   # NUM_WARPS
_AQ = 9   # a_q pre-computed view
_ASC = 10 # a_sc
_OUT = 11
_OV = 12  # out_view
_UB = 13  # use_blockscale
_KN = 14  # kernelName (or None)
_SK = 15  # splitK
_K = 16
_KHALF = 17


def _init_shape(M, K, N, device):
    QBLOCK = 32

    a_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
    scaleN_valid = triton.cdiv(K, QBLOCK)
    scaleN_pad = triton.cdiv(scaleN_valid, 8) * 8
    scaleM_pad = triton.cdiv(M, 256) * 256
    blockscale = torch.empty(
        (scaleM_pad, scaleN_pad), dtype=torch.uint8, device=device
    )

    a_q = a_fp4.view(dtypes.fp4x2)
    a_sc = blockscale.view(dtypes.fp8_e8m0)

    M_pad = (M + 31) // 32 * 32
    out = torch.empty((M_pad, N), dtype=torch.bfloat16, device=device)
    out_view = out[:M] if M != M_pad else out

    if M <= 32:
        NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N = 1, triton.next_power_of_2(M), 32
        NUM_WARPS, NUM_STAGES = 1, 1
    else:
        NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N = 4, 64, 64
        NUM_WARPS, NUM_STAGES = 4, 2
        if K <= 16384:
            BLOCK_SIZE_M, BLOCK_SIZE_N = 32, 128
    if K <= 1024:
        NUM_ITER, NUM_STAGES, NUM_WARPS = 1, 1, 4
        BLOCK_SIZE_N = max(32, min(256, triton.next_power_of_2(K)))
        BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))

    grid = (triton.cdiv(M, BLOCK_SIZE_M), triton.cdiv(K, BLOCK_SIZE_N * NUM_ITER))

    ck_config = get_GEMM_config(M, N, K)
    use_blockscale = False
    # Force 32x128 tile for all shapes (default 192x128 wastes compute for small M)
    kernelName = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
    # splitK for untuned shapes: improve CU utilization
    tile_m, tile_n, tile_k = 32, 128, 128
    tile_num = ((M_pad + tile_m - 1) // tile_m) * ((N + tile_n - 1) // tile_n)
    splitK = 0
    cu_num = 304  # MI355X
    cus_per_tile = cu_num / max(tile_num, 1)
    while (cus_per_tile >= (1 << (splitK + 1))
           and (1 << (splitK + 1)) * tile_k < 2 * K
           and splitK < 3):
        splitK += 1
    if ck_config is not None:
        kn = ck_config["kernelName"]
        sk = ck_config.get("splitK", None)
        if "_ZN" not in kn:
            use_blockscale = True
            splitK = 0 if sk is None else int(sk)
        else:
            kernelName = kn
            splitK = int(sk) if sk is not None else 0

    return (
        grid,          # 0
        a_fp4,         # 1
        blockscale,    # 2
        scaleN_pad,    # 3
        NUM_ITER,      # 4
        BLOCK_SIZE_M,  # 5
        BLOCK_SIZE_N,  # 6
        NUM_STAGES,    # 7
        NUM_WARPS,     # 8
        a_q,           # 9
        a_sc,          # 10
        out,           # 11
        out_view,      # 12
        use_blockscale,# 13
        kernelName,    # 14
        splitK,        # 15
        K,             # 16
        K // 2,        # 17
    )


_introspection_cached = False


def custom_kernel(data: input_t) -> output_t:
    global _introspection_cached
    A = data[0]
    B_sh = data[3]
    B_sc = data[4]
    M = A.shape[0]
    K = A.shape[1]
    N = B_sh.shape[0]

    c = _cfg.get((M, K, N))
    if c is None:
        c = _init_shape(M, K, N, A.device)
        _cfg[(M, K, N)] = c

    _fused_quant_shuffle_kernel[c[_GRID]](
        A, c[_A_FP4], c[_BS],
        K, 1, c[_KHALF], 1,
        M=M, N=K, scaleN_pad=c[_SNP],
        MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
        NUM_ITER=c[_NI],
        BLOCK_SIZE_M=c[_BSM],
        BLOCK_SIZE_N=c[_BSN],
        NUM_STAGES=c[_NS],
        num_warps=c[_NW],
        waves_per_eu=0, num_stages=1,
    )

    if c[_UB]:
        gemm_a4w4_blockscale(
            c[_AQ], B_sh, c[_ASC], B_sc, c[_OUT], splitK=c[_SK],
        )
    else:
        gemm_a4w4_asm(
            c[_AQ], B_sh, c[_ASC], B_sc, c[_OUT],
            c[_KN] or "", None, 1.0, 0.0, True, c[_SK],
        )

    if not _introspection_cached:
        _introspection_cached = True
        try:
            import inspect
            import typing
            inner = gemm_a4w4_asm.__globals__.get('_gemm_a4w4_asm')
            if inner is not None and hasattr(inner, '__wrapped__'):
                fn = inner.__wrapped__
                fn.__signature__ = inspect.signature(fn)
                _fn_id = id(fn)
                _fn_hints = typing.get_type_hints(fn)
                _orig_gth = typing.get_type_hints

                def _fast_gth(obj, globalns=None, localns=None,
                              include_extras=False):
                    if id(obj) == _fn_id:
                        return _fn_hints
                    return _orig_gth(obj, globalns=globalns,
                                     localns=localns,
                                     include_extras=include_extras)

                typing.get_type_hints = _fast_gth
        except Exception:
            pass

    return c[_OV]
scrolls · 278 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