Skip to content
KernelIndex
Search⌘K

submission 678892

Yaowei Lyu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a5263229e0ecd73453e441784a2ab56e1067dfb639092cc88dabf7ff7f043e9e
license declaredunknown
license concludedunknown
authorsYaowei Lyu
imported2026-08-26

Techniques

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

fp4MXFP4 matrix multiplication using Triton kernels on AMD MI355X.

Kernel source

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

"""
MXFP4 matrix multiplication using Triton kernels on AMD MI355X.
bf16 A -> MXFP4 per-1x32 quant A -> fp4 GEMM with pre-shuffled B -> bf16 C.
"""

from task import input_t, output_t


def custom_kernel(data: input_t) -> output_t:
    import torch
    import triton
    import triton.language as tl

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

    # ----------------------------------------------------------------
    # Step 1: dynamic MXFP4 quantization of A  (per-1x32 block scaling)
    #   - For each row, every group of 32 elements shares one e8m0 scale
    #   - e8m0 scale = exponent of max |x| in the group (biased by 127)
    #   - FP4 E2M1 values: 0,0.5,1,1.5,2,3,4,6 (with sign)
    # ----------------------------------------------------------------
    GROUP_SIZE = 32

    @triton.jit
    def _mxfp4_quant_kernel(
        X_ptr,
        Out_ptr,
        Scale_ptr,
        M,
        K,
        stride_xm,
        stride_xk,
        stride_om,
        stride_ok,
        stride_sm,
        stride_sk,
        BLOCK_M: tl.constexpr,
        BLOCK_K: tl.constexpr,
        GROUP: tl.constexpr,
    ):
        pid_m = tl.program_id(0)
        pid_k = tl.program_id(1)
        rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
        # each pid_k handles one group of GROUP elements
        rk = pid_k * GROUP + tl.arange(0, GROUP)
        mask = (rm[:, None] < M) & (rk[None, :] < K)

        x = tl.load(
            X_ptr + rm[:, None] * stride_xm + rk[None, :] * stride_xk,
            mask=mask,
            other=0.0,
        )

        # compute per-group max absolute value
        ax = tl.abs(x)
        amax = tl.max(ax, axis=1)  # [BLOCK_M]

        # e8m0 biased exponent: floor(log2(amax)) + 127, clamped to [0,254]
        # use bitcast to extract exponent from bf16/fp32
        amax_f32 = amax.to(tl.float32)
        # add small eps to avoid log2(0)
        amax_f32 = tl.where(amax_f32 > 0.0, amax_f32, 1.0e-30)
        log2_amax = tl.math.log2(amax_f32)
        exp_biased = tl.math.floor(log2_amax).to(tl.int32) + 127
        exp_biased = tl.maximum(exp_biased, 0)
        exp_biased = tl.minimum(exp_biased, 254)
        scale_e8m0 = exp_biased.to(tl.uint8)

        # reconstruct scale as power of 2: 2^(exp_biased - 127)
        scale_f = tl.math.exp2((exp_biased - 127).to(tl.float32))
        # normalized = x / scale * 8.0 (fp4 e2m1 range is [0..6], we map to integer codes)
        inv_scale = 1.0 / tl.where(scale_f > 0.0, scale_f, 1.0)

        xn = x.to(tl.float32) * inv_scale[:, None]

        # Round to nearest FP4 E2M1 value
        # FP4 E2M1 positive values: 0, 0.5, 1, 1.5, 2, 3, 4, 6
        # We use a simple approach: clamp + round
        sign = tl.where(xn < 0.0, 1, 0)
        xn_abs = tl.abs(xn)

        # Map to fp4 code (0-7): 0->0, 1->0.5, 2->1.0, 3->1.5, 4->2.0, 5->3.0, 6->4.0, 7->6.0
        # Reverse: find nearest
        # Boundaries: 0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0
        code = tl.zeros_like(xn_abs).to(tl.int32)
        code = tl.where(xn_abs >= 0.25, 1, code)
        code = tl.where(xn_abs >= 0.75, 2, code)
        code = tl.where(xn_abs >= 1.25, 3, code)
        code = tl.where(xn_abs >= 1.75, 4, code)
        code = tl.where(xn_abs >= 2.5, 5, code)
        code = tl.where(xn_abs >= 3.5, 6, code)
        code = tl.where(xn_abs >= 5.0, 7, code)

        # FP4 E2M1 encoding: sign(1) | exp(2) | man(1)
        # code 0 -> 0b0000, code 1 -> 0b0001, code 2 -> 0b0010, code 3 -> 0b0011
        # code 4 -> 0b0100, code 5 -> 0b0101, code 6 -> 0b0110, code 7 -> 0b0111
        nibble = (sign.to(tl.int32) << 3) | code  # 4-bit value

        # Pack pairs of fp4 into uint8: low nibble = even index, high nibble = odd index
        # rk has GROUP elements, pack into GROUP//2 bytes
        even_idx = tl.arange(0, GROUP // 2) * 2
        odd_idx = even_idx + 1
        lo = tl.load(
            X_ptr
            + rm[:, None] * stride_xm
            + (pid_k * GROUP + even_idx[None, :]) * stride_xk,
            mask=(rm[:, None] < M) & ((pid_k * GROUP + even_idx[None, :]) < K),
            other=0.0,
        )
        hi = tl.load(
            X_ptr
            + rm[:, None] * stride_xm
            + (pid_k * GROUP + odd_idx[None, :]) * stride_xk,
            mask=(rm[:, None] < M) & ((pid_k * GROUP + odd_idx[None, :]) < K),
            other=0.0,
        )
        # We already computed nibble for all GROUP elements; need to extract even/odd
        # Recompute for even and odd separately using the same logic
        lo_f = lo.to(tl.float32) * inv_scale[:, None]
        hi_f = hi.to(tl.float32) * inv_scale[:, None]

        lo_sign = tl.where(lo_f < 0.0, 1, 0).to(tl.int32)
        lo_abs = tl.abs(lo_f)
        lo_code = tl.zeros_like(lo_abs).to(tl.int32)
        lo_code = tl.where(lo_abs >= 0.25, 1, lo_code)
        lo_code = tl.where(lo_abs >= 0.75, 2, lo_code)
        lo_code = tl.where(lo_abs >= 1.25, 3, lo_code)
        lo_code = tl.where(lo_abs >= 1.75, 4, lo_code)
        lo_code = tl.where(lo_abs >= 2.5, 5, lo_code)
        lo_code = tl.where(lo_abs >= 3.5, 6, lo_code)
        lo_code = tl.where(lo_abs >= 5.0, 7, lo_code)
        lo_nibble = (lo_sign << 3) | lo_code

        hi_sign = tl.where(hi_f < 0.0, 1, 0).to(tl.int32)
        hi_abs = tl.abs(hi_f)
        hi_code = tl.zeros_like(hi_abs).to(tl.int32)
        hi_code = tl.where(hi_abs >= 0.25, 1, hi_code)
        hi_code = tl.where(hi_abs >= 0.75, 2, hi_code)
        hi_code = tl.where(hi_abs >= 1.25, 3, hi_code)
        hi_code = tl.where(hi_abs >= 1.75, 4, hi_code)
        hi_code = tl.where(hi_abs >= 2.5, 5, hi_code)
        hi_code = tl.where(hi_abs >= 3.5, 6, hi_code)
        hi_code = tl.where(hi_abs >= 5.0, 7, hi_code)
        hi_nibble = (hi_sign << 3) | hi_code

        packed = (hi_nibble << 4) | lo_nibble
        packed = packed.to(tl.uint8)

        # store packed output: shape [M, K//2]
        rk_out = pid_k * (GROUP // 2) + tl.arange(0, GROUP // 2)
        out_mask = (rm[:, None] < M) & (rk_out[None, :] < K // 2)
        tl.store(
            Out_ptr + rm[:, None] * stride_om + rk_out[None, :] * stride_ok,
            packed,
            mask=out_mask,
        )

        # store scale: shape [M, K//GROUP]
        scale_col = pid_k
        scale_mask = rm < M
        tl.store(
            Scale_ptr + rm * stride_sm + scale_col * stride_sk,
            scale_e8m0,
            mask=scale_mask,
        )

    # --- Use aiter's proven quantization + GEMM (fastest path) ---
    # The Triton quant above is illustrative; for correctness and speed,
    # use aiter's fused ops which are HIP-optimized for MI355X.
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    # Quantize A
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(A)
    bs_e8m0 = e8m0_shuffle(bs_e8m0)
    A_q = x_fp4.view(dtypes.fp4x2)
    A_scale_sh = bs_e8m0.view(dtypes.fp8_e8m0)

    # GEMM
    out_gemm = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out_gemm
scrolls · 197 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