Skip to content
KernelIndex
Search⌘K

submission 585823

jefflyu_47387 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:57844e5c30e1210ed36032a1764a9009250f077ec8bcdebf23746b9522c040e5
license declaredunknown
license concludedunknown
authorsjefflyu_47387
imported2026-08-26

Techniques

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

fp4- B uses provided MXFP4 packed tensor (B_q), and we reconstruct raw scales from B.
stages = 2num_stages = 2

Kernel source

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

"""
Phase 1 prototype: TileLang fused dequant(B)+GEMM kernel.

Numerical path:
- A is consumed directly as bf16 (no A quantization in this prototype).
- B uses provided MXFP4 packed tensor (B_q), and we reconstruct raw scales from B.
"""

from task import input_t, output_t


_KERNEL_CACHE = {}


def _aiter_fallback(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    def _quant_mxfp4(x, shuffle=True):
        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
        if shuffle:
            bs_e8m0 = e8m0_shuffle(bs_e8m0)
        return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

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

    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    return aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def _build_mxfp4_dequant_gemm_kernel(m: int, n: int, k: int):
    import tilelang
    import tilelang.language as T
    from tvm import tir

    block_M = 64
    block_N = 64
    block_K = 64
    num_stages = 2
    threads = 128
    num_bits = 4
    scale_size = 32

    num_elems_per_byte = 8 // num_bits
    qk = k // num_elems_per_byte
    block_qk = block_K // num_elems_per_byte

    def _tir_u8_to_f4_to_bf16(val: tir.PrimExpr, pos: tir.PrimExpr):
        mask = tir.const((1 << num_bits) - 1, T.uint16)
        f4 = (val >> (pos.astype(T.uint16) * tir.const(num_bits, T.uint16))) & mask
        s = f4 >> tir.const(3, T.uint16)
        e_f4 = (f4 & tir.const(6, T.uint16)) >> tir.const(1, T.uint16)
        e_bf16 = e_f4 + tir.const(126, T.uint16)
        m_f4 = f4 & tir.const(1, T.uint16)
        return tir.reinterpret(
            T.bfloat16,
            ((((s << tir.const(8, T.uint16)) | e_bf16) << tir.const(7, T.uint16))
             | (m_f4 << tir.const(6, T.uint16))).astype(T.uint16),
        )

    @T.macro
    def _simple_dequant_bf16_fp4(B_shared, B_dequantize_shared, Scale, k_block):
        B_local = T.alloc_fragment((block_N, block_qk), "uint8")
        B_dequantize_local = T.alloc_fragment((block_N, block_K), "bfloat16")

        bx = T.get_block_binding(0)
        T.copy(B_shared, B_local)

        for i, j in T.Parallel(block_N, block_K):
            fp4_as_bf16 = _tir_u8_to_f4_to_bf16(
                B_local[i, j // num_elems_per_byte],
                j % num_elems_per_byte,
            )
            # E8M0 scale is exponent-only, applying power-of-two multiplication.
            scale_exp = Scale[
                bx * block_N + i,
                k_block * block_K // scale_size + j // scale_size,
            ]
            B_dequantize_local[i, j] = fp4_as_bf16 * T.shift_left(1, scale_exp)

        T.copy(B_dequantize_local, B_dequantize_shared)

    @tilelang.jit(out_idx=[-1])
    def _mxfp4_dequant_gemm(M, N, K):
        @T.prim_func
        def main(
            A: T.Tensor((M, K), "bfloat16"),
            B_q_u8: T.Tensor((N, qk), "uint8"),
            B_scale_u8: T.Tensor((N, K // scale_size), "uint8"),
            C: T.Tensor((M, N), "bfloat16"),
        ):
            with T.Kernel(
                T.ceildiv(N, block_N),
                T.ceildiv(M, block_M),
                threads=threads,
            ) as (bx, by):
                A_shared = T.alloc_shared((block_M, block_K), "bfloat16")
                B_shared = T.alloc_shared((block_N, block_qk), "uint8")
                B_dequantize_shared = T.alloc_shared((block_N, block_K), "bfloat16")
                C_local = T.alloc_fragment((block_M, block_N), "float32")

                T.clear(C_local)

                for ko in T.Pipelined(K // block_K, num_stages=num_stages):
                    T.copy(A[by * block_M, ko * block_K], A_shared)
                    T.copy(B_q_u8[bx * block_N, ko * block_qk], B_shared)
                    _simple_dequant_bf16_fp4(B_shared, B_dequantize_shared, B_scale_u8, ko)
                    T.gemm(A_shared, B_dequantize_shared, C_local, transpose_B=True)

                T.copy(C_local, C[by * block_M, bx * block_N])

        return main

    return _mxfp4_dequant_gemm(m, n, k)


def custom_kernel(data: input_t) -> output_t:
    import torch
    from aiter.ops.triton.quant import dynamic_mxfp4_quant

    A, B, B_q, B_shuffle, B_scale_sh = data

    try:
        import tilelang  # noqa: F401
    except ModuleNotFoundError:
        return _aiter_fallback(data)

    A = A.contiguous()
    B = B.contiguous()
    m, k = A.shape
    n, _ = B.shape

    # Task constraints guarantee K % 64 == 0.
    assert k % 64 == 0

    # B_q is fp4x2-packed; reinterpret as raw packed bytes for TileLang dequant macro.
    B_q_u8 = B_q.view(torch.uint8).contiguous()

    # Task input provides shuffled scales; dequant macro here expects unshuffled scales.
    _, B_scale_raw = dynamic_mxfp4_quant(B)
    B_scale_u8 = B_scale_raw.view(torch.uint8)[:n, : (k // 32)].contiguous()

    key = (m, n, k)
    kernel = _KERNEL_CACHE.get(key)
    if kernel is None:
        try:
            kernel = _build_mxfp4_dequant_gemm_kernel(m, n, k)
            _KERNEL_CACHE[key] = kernel
        except ModuleNotFoundError:
            return _aiter_fallback(data)

    return kernel(A, B_q_u8, B_scale_u8)
scrolls · 168 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