Skip to content
KernelIndex
Search⌘K

submission 573481

n8_gr8_ · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-573481?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
13.2µs
#412 of 1143
2026-03-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:353c6f3b25e74782fda6335e85280cdd16ac5b4b1f37ae6cb473c62e422aef55
license declaredunknown
license concludedunknown
authorsn8_gr8_
imported2026-08-26

Techniques

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

num-warps = 1sm, sn, BLOCK_M=BLOCK_M, BLOCK_N=sn_po2, num_warps=1,
stages = 1NUM_STAGES=NSC, num_warps=NW, waves_per_eu=0, num_stages=1,
tile-m = 32BLOCK_M = 32

Kernel source

submission.py250 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os, sys

os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_LOG_LEVEL"] = "ERROR"

import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from task import input_t, output_t

_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0

_fqs_ok = False
try:
    from aiter.ops.triton._triton_kernels.quant.quant import (
        _dynamic_mxfp4_quant_kernel,
        _mxfp4_quant_op,
    )
    _fqs_ok = True
except Exception:
    pass


@triton.jit
def _e8m0_unshuffle_kernel(
    src_ptr, dst_ptr,
    sm, sn,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    n = tl.arange(0, BLOCK_N)
    i0 = m // 32
    i1 = (m // 16) % 2
    i2 = m % 16
    i3 = n // 8
    i4 = (n // 4) % 2
    i5 = n % 4
    shuffled_idx = (
        i0[:, None] * (sn * 32) +
        i3[None, :] * 256 +
        i5[None, :] * 64 +
        i2[:, None] * 4 +
        i4[None, :] * 2 +
        i1[:, None]
    )
    mask = (m < sm)[:, None] & (n < sn)[None, :]
    vals = tl.load(src_ptr + shuffled_idx, mask=mask)
    dst_offs = m[:, None] * sn + n[None, :]
    tl.store(dst_ptr + dst_offs, vals, mask=mask)


if _fqs_ok:
    @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_xfp4_m_in, stride_xfp4_n_in,
        M, N, N_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_xfp4_m = tl.cast(stride_xfp4_m_in, tl.int64)
        stride_xfp4_n = tl.cast(stride_xfp4_n_in, tl.int64)
        QBS: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE
        NQB: 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, QBS)

            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_xfp4_m + out_offs_n[None, :] * stride_xfp4_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 * NQB + tl.arange(0, NQB)
            shuffled_offs = (
                (bs_m // 32)[:, None] * (32 * N_pad) +
                (bs_n // 8)[None, :] * 256 +
                (bs_n % 4)[None, :] * 64 +
                (bs_m % 16)[:, None] * 4 +
                ((bs_n // 4) % 2)[None, :] * 2 +
                ((bs_m // 16) % 2)[:, None]
            )

            if EVEN_M_N:
                tl.store(bs_shuffled_ptr + shuffled_offs, bs_e8m0)
            else:
                N_scale = (N + QBS - 1) // QBS
                bs_mask = (bs_m < M)[:, None] & (bs_n < N_scale)[None, :]
                tl.store(bs_shuffled_ptr + shuffled_offs, bs_e8m0, mask=bs_mask)


_fused_out = {}
_unshuffle_bufs = {}
_quant_bufs = {}

_ASM_M_THRESHOLD = 64

def _cfg(bm, bn, bk, gm, nw, ns, wpe, mi, ks, cm=None):
    return {
        "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
        "GROUP_SIZE_M": gm, "num_warps": nw, "num_stages": ns,
        "waves_per_eu": wpe, "matrix_instr_nonkdim": mi,
        "NUM_KSPLIT": ks, "cache_modifier": cm,
    }

_CONFIGS = {
    (4, 512, 2880):    _cfg(4, 128, 512, 1, 4, 1, 2, 16, 1),
    (16, 7168, 2112):  _cfg(8, 128, 512, 1, 4, 2, 2, 16, 7),
    (32, 512, 4096):   _cfg(8, 128, 512, 1, 4, 1, 2, 16, 2),
    (32, 512, 2880):   _cfg(8, 128, 512, 1, 4, 1, 2, 16, 3),
    (64, 2048, 7168):  _cfg(8, 128, 512, 1, 4, 1, 2, 16, 1),
    (256, 1536, 3072): _cfg(8, 128, 512, 1, 4, 1, 2, 16, 1),
}

_bscale_ptr = -1
_bscale_out = None


def _fast_unshuffle_bscale(B_scale_sh, N, K):
    global _bscale_ptr, _bscale_out
    ptr = B_scale_sh.data_ptr()
    if ptr == _bscale_ptr:
        return _bscale_out
    QBS = 32
    sm = N
    sn = (K + QBS - 1) // QBS
    bkey = (sm, sn)
    dst = _unshuffle_bufs.get(bkey)
    if dst is None:
        dst = torch.empty(sm, sn, dtype=torch.uint8, device=B_scale_sh.device)
        _unshuffle_bufs[bkey] = dst
    sn_po2 = triton.next_power_of_2(sn)
    BLOCK_M = 32
    grid = (triton.cdiv(sm, BLOCK_M),)
    _e8m0_unshuffle_kernel[grid](
        B_scale_sh.view(torch.uint8).reshape(-1), dst.view(-1),
        sm, sn, BLOCK_M=BLOCK_M, BLOCK_N=sn_po2, num_warps=1,
    )
    _bscale_ptr = ptr
    _bscale_out = dst
    return dst


def _fast_quant_shuffle(x):
    M, N = x.shape
    QBS = 32
    N_scale = (N + QBS - 1) // QBS
    M_pad = triton.cdiv(M, 32) * 32
    N_pad = triton.cdiv(N_scale, 8) * 8

    bkey = (M, N)
    bufs = _quant_bufs.get(bkey)
    if bufs is not None:
        x_fp4, bs_shuffled = bufs
    else:
        x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
        needs_pad = M_pad != M or N_pad != N_scale
        if needs_pad:
            bs_shuffled = torch.zeros(M_pad * N_pad, dtype=torch.uint8, device=x.device)
        else:
            bs_shuffled = torch.empty(M_pad * N_pad, dtype=torch.uint8, device=x.device)
        _quant_bufs[bkey] = (x_fp4, bs_shuffled)

    if M <= 32:
        NI, BSM, BSN, NW, NSC = 1, min(4, triton.next_power_of_2(M)), 32, 1, 1
    elif M <= 64:
        NI, BSM, BSN, NW, NSC = 1, 4, 128, 1, 1
    else:
        NI, BSM, BSN, NW, NSC = 1, 8, 128, 2, 1

    if N <= 1024:
        NI, NSC, NW = 1, 1, 4
        BSN = max(32, min(256, triton.next_power_of_2(N)))
        BSM = min(8, triton.next_power_of_2(M))

    grid = (triton.cdiv(M, BSM), triton.cdiv(N, BSN * NI))
    _fused_quant_shuffle_kernel[grid](
        x, x_fp4, bs_shuffled,
        *x.stride(), *x_fp4.stride(),
        M=M, N=N, N_pad=N_pad,
        MXFP4_QUANT_BLOCK_SIZE=QBS, SCALING_MODE=0,
        NUM_ITER=NI, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
        NUM_STAGES=NSC, num_warps=NW, waves_per_eu=0, num_stages=1,
    )
    return x_fp4, bs_shuffled.view(M_pad, N_pad)


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape

    if M >= _ASM_M_THRESHOLD:
        if _fqs_ok:
            x_fp4, bs_shuffled = _fast_quant_shuffle(A)
        else:
            x_fp4, bs = dynamic_mxfp4_quant(A)
            bs_shuffled = e8m0_shuffle(bs)
        aq = x_fp4.view(_fp4x2)
        asc = bs_shuffled.view(_fp8_e8m0)
        return aiter.gemm_a4w4(aq, B_shuffle, asc, B_scale_sh, dtype=_bf16, bpreshuffle=True)

    bq_u8 = B_q.view(torch.uint8)
    N_b = bq_u8.shape[0]
    bsc = _fast_unshuffle_bscale(B_scale_sh, N_b, K)
    mn = M << 14 | N_b
    out = _fused_out.get(mn)
    if out is None:
        out = torch.empty(M, N_b, dtype=_bf16, device=A.device)
        _fused_out[mn] = out
    cfg = _CONFIGS.get((M, K, N_b))
    return gemm_a16wfp4(A, bq_u8, bsc, y=out, config=cfg)
scrolls · 250 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