Skip to content
KernelIndex
Search⌘K

submission 645164

Esquie · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-645164?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
#489 of 1143
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ee6e042eff0b8bf08335930717e3071ab7949708201ed7f311244b997361d20b
license declaredunknown
license concludedunknown
authorsEsquie
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM: fused quant+shuffle + explicit ASM kernel selection.
split-kkernel_name, split_k = GEMM_CONFIGS.get((m, n, k), (ASM_32x128, 0))
stages = 1num_warps=c["NW"], waves_per_eu=0, num_stages=1,

Kernel source

submission.py161 lines
"""
MXFP4 GEMM: fused quant+shuffle + explicit ASM kernel selection.
Default path picks 192x128 for small M (wastes 188 rows at M=4).
Force 32x128 for all shapes — matches tuned config from 256-CU CSV.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
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

_cache = {}
SCALE_GROUP_SIZE = 32

ASM_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
ASM_64x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"


@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_shuffled_ptr,
    stride_x_m_in, stride_x_n_in,
    stride_fp4_m_in, stride_fp4_n_in,
    M, N, K_SCALE_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_fp4_m = tl.cast(stride_fp4_m_in, tl.int64)
    stride_fp4_n = tl.cast(stride_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_fp4_m + out_offs_n[None, :] * stride_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)

        rows = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)).to(tl.int64)
        cols = (pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)).to(tl.int64)
        i0 = rows // 32; i1 = (rows // 16) % 2; i2 = rows % 16
        i3 = cols // 8; i4 = (cols // 4) % 2; i5 = cols % 4
        K_SP = tl.cast(K_SCALE_PAD, tl.int64)
        shuffled_idx = (i0[:, None] * (K_SP * 32) + i3[None, :] * 256
                        + i5[None, :] * 64 + i2[:, None] * 4
                        + i4[None, :] * 2 + i1[:, None])
        if EVEN_M_N:
            tl.store(bs_shuffled_ptr + shuffled_idx, bs_e8m0)
        else:
            n_scale = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
            bs_mask = (rows[:, None] < M) & (cols[None, :] < n_scale)
            tl.store(bs_shuffled_ptr + shuffled_idx, bs_e8m0, mask=bs_mask)


def _get_quant_params(m, k):
    if m <= 32:
        NI, BSM, BSN, NW, NS = 1, triton.next_power_of_2(m), 32, 1, 1
    else:
        NI, BSM, BSN, NW, NS = 4, 64, 64, 4, 2
        if k <= 16384:
            BSM, BSN = 32, 128
    if k <= 1024:
        NI, NS, NW = 1, 1, 4
        BSN = max(32, min(256, triton.next_power_of_2(k)))
        BSM = min(8, triton.next_power_of_2(m))
    grid = (triton.cdiv(m, BSM), triton.cdiv(k, BSN * NI))
    return grid, BSM, BSN, NI, NS, NW


# Per-shape GEMM config: (kernel_name, log2_k_split)
# 32x128 is the best ASM kernel for M<=64 per the 256-CU tuned CSV
# 64x128 is better for M=256 per the CSV
GEMM_CONFIGS = {
    (4, 2880, 512):    (ASM_32x128, 0),
    (16, 2112, 7168):  (ASM_32x128, 0),
    (32, 4096, 512):   (ASM_32x128, 0),
    (32, 2880, 512):   (ASM_32x128, 0),
    (64, 7168, 2048):  (ASM_32x128, 0),
    (256, 3072, 1536): (ASM_32x128, 0),
}


def _setup(m, n, k):
    k_scale = k // SCALE_GROUP_SIZE
    m_pad = ((m + 255) // 256) * 256
    k_scale_pad = ((k_scale + 7) // 8) * 8

    x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device="cuda")
    bs_shuffled = torch.empty(m_pad * k_scale_pad, dtype=torch.uint8, device="cuda")
    out = torch.empty(((m + 31) // 32 * 32, n), dtype=torch.bfloat16, device="cuda")

    grid, BSM, BSN, NI, NS, NW = _get_quant_params(m, k)

    kernel_name, split_k = GEMM_CONFIGS.get((m, n, k), (ASM_32x128, 0))

    return {
        "x_fp4": x_fp4, "bs_shuffled": bs_shuffled, "out": out,
        "m_pad": m_pad, "k_scale_pad": k_scale_pad,
        "grid": grid, "BSM": BSM, "BSN": BSN, "NI": NI, "NS": NS, "NW": NW,
        "kernel_name": kernel_name, "split_k": split_k,
    }


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n = B.shape[0]
    key = (m, n, k)

    if key not in _cache:
        _cache[key] = _setup(m, n, k)

    c = _cache[key]

    _fused_quant_shuffle_kernel[c["grid"]](
        A, c["x_fp4"], c["bs_shuffled"],
        A.stride(0), A.stride(1),
        c["x_fp4"].stride(0), c["x_fp4"].stride(1),
        m, k, c["k_scale_pad"],
        BLOCK_SIZE_M=c["BSM"], BLOCK_SIZE_N=c["BSN"],
        NUM_ITER=c["NI"], NUM_STAGES=c["NS"],
        MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
        num_warps=c["NW"], waves_per_eu=0, num_stages=1,
    )

    A_q = c["x_fp4"].view(dtypes.fp4x2)
    A_scale_sh = c["bs_shuffled"].view(c["m_pad"], c["k_scale_pad"]).view(dtypes.fp8_e8m0)

    gemm_a4w4_asm(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        c["out"], c["kernel_name"],
        None, 1.0, 0.0, True, c["split_k"],
    )
    return c["out"][:m]
scrolls · 161 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