Skip to content
KernelIndex
Search⌘K

submission 720763

mingkai_37292 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:56b4697fcbdad7c711fd3942a97698a83fa2e19dacf49d3ceca40d89f14e6d55
license declaredunknown
license concludedunknown
authorsmingkai_37292
imported2026-08-26

Techniques

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

autotunethe e8m0_shuffle step. BLOCK_M is autotuned per (M, K) shape to minimize wasted work
fp4FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
num-warps = 1triton.Config({"BLOCK_M": 4}, num_warps=1),
tile-m = 4BLOCK_M: tl.constexpr = 4,

Kernel source

submission.py202 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quantization is implemented as a custom Triton kernel.
Scales are written directly into the shuffled layout expected by gemm_a4w4, fusing
the e8m0_shuffle step. BLOCK_M is autotuned per (M, K) shape to minimize wasted work
on small-M shapes and maximize occupancy on larger ones.
"""
import torch
import triton
import triton.language as tl

from task import input_t, output_t


@triton.autotune(
    configs=[
        triton.Config({"BLOCK_M": 4},  num_warps=1),
        triton.Config({"BLOCK_M": 8},  num_warps=2),
        triton.Config({"BLOCK_M": 16}, num_warps=4),
        triton.Config({"BLOCK_M": 32}, num_warps=4),
    ],
    key=["M", "K"],
)
@triton.jit
def _mxfp4_quant_kernel(
    x_ptr,
    out_ptr,
    scale_sh_ptr,   # pre-shuffled scale buffer: flat [M_padded * (K//32)]
    M, K,
    stride_xm,
    GROUP_SIZE: tl.constexpr = 32,
    BLOCK_M: tl.constexpr = 4,
):
    row_start = tl.program_id(0) * BLOCK_M
    group_id = tl.program_id(1)

    k_start = group_id * GROUP_SIZE

    # 1D row offsets and 2D tensors
    row_offs_1d = row_start + tl.arange(0, BLOCK_M)              # [BLOCK_M]
    row_offs = row_offs_1d[:, None]                               # [BLOCK_M, 1]
    half_offs = tl.arange(0, GROUP_SIZE // 2)[None, :]            # [1, 16]

    k_even = k_start + half_offs * 2      # [1, 16]
    k_odd  = k_start + half_offs * 2 + 1  # [1, 16]

    row_mask_1d = row_offs_1d < M         # [BLOCK_M]
    row_mask = row_mask_1d[:, None]       # [BLOCK_M, 1] — broadcasts over columns
    mask_e = row_mask & (k_even < K)      # [BLOCK_M, 16]
    mask_o = row_mask & (k_odd  < K)      # [BLOCK_M, 16]

    x_even = tl.load(x_ptr + row_offs * stride_xm + k_even, mask=mask_e, other=0.0).to(tl.float32)
    x_odd  = tl.load(x_ptr + row_offs * stride_xm + k_odd,  mask=mask_o, other=0.0).to(tl.float32)

    # E8M0 scale per row: reduce over 16 cols -> [BLOCK_M]
    abs_max_1d = tl.maximum(tl.max(tl.abs(x_even), axis=1),
                             tl.max(tl.abs(x_odd),  axis=1))   # [BLOCK_M]

    # Match reference _mxfp4_quant_op scale exactly via bitwise rounding
    abs_max_1d = tl.maximum(abs_max_1d, 1e-38).to(tl.float32)
    abs_max_int = abs_max_1d.to(tl.int32, bitcast=True)
    abs_max_rounded = ((abs_max_int + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000).to(tl.float32, bitcast=True)
    scale_e8m0_unbiased = tl.floor(tl.math.log2(abs_max_rounded)).to(tl.int32) - 2
    scale_e8m0_unbiased = tl.minimum(tl.maximum(scale_e8m0_unbiased, -127), 127)
    e8m0_exp = (scale_e8m0_unbiased + 127).to(tl.uint8)       # [BLOCK_M]
    quant_scale = tl.math.exp2(-scale_e8m0_unbiased.to(tl.float32))[:, None]  # [BLOCK_M, 1]

    # Store scales directly into shuffled layout (fusing e8m0_shuffle).
    sn = K // GROUP_SIZE
    j = group_id
    sh_off = (
        (row_offs_1d // 32) * (32 * sn)
        + (j // 8) * 256
        + (j % 4) * 64
        + (row_offs_1d % 16) * 4
        + ((j // 4) % 2) * 2
        + (row_offs_1d // 16) % 2
    )
    tl.store(scale_sh_ptr + sh_off, e8m0_exp, mask=row_mask_1d)

    # E2M1 bit-manipulation encoding matching ROCm/aiter _mxfp4_quant_op exactly
    # Even elements -> lo nibbles
    xs_e = x_even * quant_scale
    xs_e_uint = xs_e.to(tl.int32, bitcast=True).to(tl.uint32)
    s_e = xs_e_uint & 0x80000000
    xs_e_pos_uint = xs_e_uint ^ s_e
    xs_e_pos = xs_e_pos_uint.to(tl.float32, bitcast=True)
    sat_e = xs_e_pos >= 6.0
    den_e = xs_e_pos < 1.0
    mant_odd_e = (xs_e_pos_uint >> 22) & 1
    norm_e = ((xs_e_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_e.to(tl.int32)) >> 22
    norm_e = norm_e.to(tl.uint8)
    den_val_e = (xs_e_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
    den_val_e = den_val_e.to(tl.uint8)
    q_e = tl.full(xs_e.shape, 7, dtype=tl.uint8)
    q_e = tl.where(~sat_e, norm_e, q_e)
    q_e = tl.where(den_e, den_val_e, q_e)
    sign_e_lp = (s_e >> 28).to(tl.uint8)
    lo = (q_e | sign_e_lp) & 0xF

    # Odd elements -> hi nibbles
    xs_o = x_odd * quant_scale
    xs_o_uint = xs_o.to(tl.int32, bitcast=True).to(tl.uint32)
    s_o = xs_o_uint & 0x80000000
    xs_o_pos_uint = xs_o_uint ^ s_o
    xs_o_pos = xs_o_pos_uint.to(tl.float32, bitcast=True)
    sat_o = xs_o_pos >= 6.0
    den_o = xs_o_pos < 1.0
    mant_odd_o = (xs_o_pos_uint >> 22) & 1
    norm_o = ((xs_o_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_o.to(tl.int32)) >> 22
    norm_o = norm_o.to(tl.uint8)
    den_val_o = (xs_o_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
    den_val_o = den_val_o.to(tl.uint8)
    q_o = tl.full(xs_o.shape, 7, dtype=tl.uint8)
    q_o = tl.where(~sat_o, norm_o, q_o)
    q_o = tl.where(den_o, den_val_o, q_o)
    sign_o_lp = (s_o >> 28).to(tl.uint8)
    hi = ((q_o | sign_o_lp) & 0xF) << 4

    # Pack two FP4 nibbles per byte: lo nibble = even index, hi nibble = odd index
    packed = lo | hi   # [BLOCK_M, 16]

    # Store packed output: [BLOCK_M, 16] at row_offs*(K//2) + k_start//2 + half_offs
    out_base = row_offs * (K // 2) + k_start // 2 + half_offs  # [BLOCK_M, 16]
    tl.store(out_ptr + out_base, packed.to(tl.uint8), mask=mask_e)


# Module-level workspace cache
_quant_ws: dict = {}
_gemm_ws: dict = {}


def _triton_mxfp4_quant(x: torch.Tensor):
    """x: [M, K] bf16 -> (fp4_packed [M, K//2] uint8, scale_sh [M_padded, K//32] uint8 in shuffled layout)"""
    M, K = x.shape
    assert K % 32 == 0, "K must be a multiple of 32 for per-1x32 MXFP4 quant"
    x = x.contiguous()

    key = (M, K)
    if key not in _quant_ws:
        M_padded = (M + 255) // 256 * 256
        _quant_ws[key] = (
            torch.empty(M, K // 2, dtype=torch.uint8, device=x.device),
            torch.empty(M_padded * (K // 32), dtype=torch.uint8, device=x.device),
        )
    out, scale_sh_flat = _quant_ws[key]
    M_padded = (M + 255) // 256 * 256

    def grid(meta):
        return ((M + meta["BLOCK_M"] - 1) // meta["BLOCK_M"], K // 32)

    _mxfp4_quant_kernel[grid](
        x, out, scale_sh_flat,
        M, K,
        x.stride(0),
    )
    return out, scale_sh_flat.view(M_padded, K // 32)


def custom_kernel(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes

    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()

    M, K = A.shape

    # Quantize A; scale is returned already in the shuffled layout
    A_fp4, A_scale_sh = _triton_mxfp4_quant(A)

    if M <= 64:
        N = B.shape[0]  # B is [N, K] weight matrix
        gemm_key = (M, N, K)
        if gemm_key not in _gemm_ws:
            _gemm_ws[gemm_key] = torch.empty(M, N, dtype=torch.bfloat16, device=A.device)
        out = _gemm_ws[gemm_key]
        aiter.gemm_a4w4_asm(
            A_fp4.view(dtypes.fp4x2),
            B_shuffle,
            A_scale_sh.view(dtypes.fp8_e8m0),
            B_scale_sh,
            out,
            "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
            bpreshuffle=True,
            log2_k_split=None,
        )
        return out

    out_gemm = aiter.gemm_a4w4(
        A_fp4.view(dtypes.fp4x2),
        B_shuffle,
        A_scale_sh.view(dtypes.fp8_e8m0),
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out_gemm
scrolls · 202 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