Skip to content
KernelIndex
Search⌘K

submission 679633

GodZmk · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:88418b82c95588989706b05ff07a8af98405111020ee5fb5de7a8f5c21c7d1d0
license declaredunknown
license concludedunknown
authorsGodZmk
imported2026-08-26

Techniques

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

fp4FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
num-warps = 4num_warps=4,
tile-m = 4BLOCK_M=4 rows * GROUP_SIZE//2=16 cols = 64 elements = 1 full wavefront.

Kernel source

submission.py160 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.
BLOCK_M rows processed per program to fill AMD wave64 (64 threads) completely:
  BLOCK_M=4 rows * GROUP_SIZE//2=16 cols = 64 elements = 1 full wavefront.
"""
import torch
import triton
import triton.language as tl

from task import input_t, output_t


@triton.jit
def _mxfp4_quant_kernel(
    x_ptr,
    out_ptr,
    scale_ptr,
    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: [BLOCK_M] at row_offs_1d * (K//32) + group_id
    scale_ptrs = scale_ptr + row_offs_1d * (K // GROUP_SIZE) + group_id
    tl.store(scale_ptrs, e8m0_exp, mask=row_mask_1d)

    # E2M1 bit-manipulation encoding matching ROCm/aiter _mxfp4_quant_op exactly
    # val_to_add = ((1-127)<<23) + (1<<21) - 1 = -1054867457
    # DENORM_MASK_INT = 149<<23 = 0x4A800000; as float = 2^22 = 4194304.0

    # 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)


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

    BLOCK_M = 16
    out = torch.empty(M, K // 2, dtype=torch.uint8, device=x.device)
    scale = torch.empty(M, K // 32, dtype=torch.uint8, device=x.device)

    grid = ((M + BLOCK_M - 1) // BLOCK_M, K // 32)
    _mxfp4_quant_kernel[grid](
        x, out, scale,
        M, K,
        x.stride(0),
        BLOCK_M=BLOCK_M,
        num_warps=4,
    )
    return out, scale


def custom_kernel(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes
    from aiter.utility.fp4_utils import e8m0_shuffle

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

    # Quantize A with the Triton kernel
    A_fp4, A_scale = _triton_mxfp4_quant(A)

    # Shuffle scale to match gemm_a4w4's expected layout
    A_scale_sh = e8m0_shuffle(A_scale.view(torch.uint8))

    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 · 160 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