Skip to content
KernelIndex
Search⌘K

submission 587635

garrick99 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:226e2edf6bd1c47413dd9bf801d1f758c7c2cb908286695c00823fee734e08b1
license declaredunknown
license concludedunknown
authorsgarrick99
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM v2 — Triton tl.dot_scaled with native fp4 MFMA.
tile-m = 32BLOCK_M = 32
tile-n = 128BLOCK_N = 128

Kernel source

submission.py208 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM v2 — Triton tl.dot_scaled with native fp4 MFMA.

C[M,N] = A[M,K] @ B[N,K]^T where both A and B are MXFP4 quantized.
Uses tl.dot_scaled with lhs_format='e2m1', rhs_format='e2m1' for native
fp4x4 MFMA on MI355X (gfx950).

Falls back to aiter.gemm_a4w4 if dot_scaled fails.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t


# ============================================================================
# Triton MXFP4 GEMM kernel via dot_scaled
# ============================================================================

@triton.jit
def _mxfp4_gemm_kernel(
    A_ptr, B_ptr,              # fp4x2 packed: A(M, K//2), B(N, K//2)
    A_scale_ptr, B_scale_ptr,  # E8M0: A(M_pad, K//32), B(N_pad, K//32)
    C_ptr,                     # bf16: C(M, N)
    M, N, K_PACKED,            # K_PACKED = K // 2
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_asm, stride_ask,
    stride_bsn, stride_bsk,
    stride_cm, stride_cn,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K_PACKED: tl.constexpr,   # packed bytes per K tile (power of 2)
    BLOCK_K_SCALE: tl.constexpr,    # scale blocks per K tile (BLOCK_K_PACKED // 16)
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    m_mask = offs_m < M
    n_mask = offs_n < N

    for k_start in range(0, K_PACKED, BLOCK_K_PACKED):
        offs_k = k_start + tl.arange(0, BLOCK_K_PACKED)
        k_mask = offs_k < K_PACKED

        # Load A tile: (BLOCK_M, BLOCK_K_PACKED) fp4x2
        a = tl.load(
            A_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak,
            mask=m_mask[:, None] & k_mask[None, :], other=0)

        # Load B^T tile: (BLOCK_K_PACKED, BLOCK_N) fp4x2
        # B stored as (N, K//2), we load transposed
        b_t = tl.load(
            B_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn,
            mask=k_mask[:, None] & n_mask[None, :], other=0)

        # Load A scale: (BLOCK_M, BLOCK_K_SCALE) E8M0
        k_scale_start = k_start // 16   # 16 packed bytes per scale block
        offs_ks = k_scale_start + tl.arange(0, BLOCK_K_SCALE)
        a_scale = tl.load(
            A_scale_ptr + offs_m[:, None] * stride_asm + offs_ks[None, :] * stride_ask,
            mask=m_mask[:, None], other=0)

        # Load B scale: (BLOCK_N, BLOCK_K_SCALE) E8M0 — natural (N, K//32) layout
        b_scale = tl.load(
            B_scale_ptr + offs_n[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk,
            mask=n_mask[:, None], other=0)

        # dot_scaled: A_fp4 @ B^T_fp4 with block scaling
        acc = tl.dot_scaled(
            lhs=a,
            rhs=b_t,
            lhs_scale=a_scale,
            rhs_scale=b_scale,
            lhs_format='e2m1',
            rhs_format='e2m1',
            acc=acc,
        )

    # Store C
    c = acc.to(tl.bfloat16)
    tl.store(
        C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn, c,
        mask=m_mask[:, None] & n_mask[None, :])


# ============================================================================
# Aiter fallback
# ============================================================================

def _aiter_gemm(data):
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    A_fp4, A_bs = dynamic_mxfp4_quant(A)
    A_q = A_fp4.view(dtypes.fp4x2)
    A_scale_sh = e8m0_shuffle(A_bs).view(dtypes.fp8_e8m0)

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


# ============================================================================
# Triton path
# ============================================================================

def _triton_gemm(data):
    from aiter.ops.triton.quant import dynamic_mxfp4_quant

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

    # Quantize A (unshuffled)
    A_fp4, A_scale = dynamic_mxfp4_quant(A)

    # Re-quantize B for unshuffled scale (B_q is already unshuffled fp4x2)
    _, B_scale = dynamic_mxfp4_quant(B)

    # View as uint8 for Triton
    a_data = A_fp4
    if a_data.dtype != torch.uint8:
        a_data = a_data.view(torch.uint8)
    b_data = B_q
    if not isinstance(b_data, torch.Tensor):
        b_data = b_data
    if b_data.dtype != torch.uint8:
        b_data = b_data.view(torch.uint8)
    a_sc = A_scale
    if a_sc.dtype != torch.uint8:
        a_sc = a_sc.view(torch.uint8)
    b_sc = B_scale
    if b_sc.dtype != torch.uint8:
        b_sc = b_sc.view(torch.uint8)

    C = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')

    BLOCK_M = 32
    BLOCK_N = 128
    BLOCK_K_PACKED = 64    # 64 packed bytes = 128 logical fp4 elements
    BLOCK_K_SCALE = BLOCK_K_PACKED // 16   # 4 scale blocks

    grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N))
    _mxfp4_gemm_kernel[grid](
        a_data, b_data,
        a_sc, b_sc,
        C,
        m, n, k_packed,
        a_data.stride(0), a_data.stride(1),
        b_data.stride(0), b_data.stride(1),
        a_sc.stride(0), a_sc.stride(1),
        b_sc.stride(0), b_sc.stride(1),
        C.stride(0), C.stride(1),
        BLOCK_M=BLOCK_M,
        BLOCK_N=BLOCK_N,
        BLOCK_K_PACKED=BLOCK_K_PACKED,
        BLOCK_K_SCALE=BLOCK_K_SCALE,
    )
    return C


# ============================================================================
# Entry point with auto-fallback
# ============================================================================

_use_triton = None

def custom_kernel(data: input_t) -> output_t:
    global _use_triton

    if _use_triton is None:
        try:
            result = _triton_gemm(data)
            _use_triton = True
            return result
        except Exception as e:
            import sys, traceback
            print(f"Triton GEMM FAILED: {type(e).__name__}: {e}", file=sys.stderr)
            traceback.print_exc(file=sys.stderr)
            _use_triton = False
            return _aiter_gemm(data)

    if _use_triton:
        # Triton dot_scaled wins for small K, aiter wins for large K
        A, B, B_q, B_shuffle, B_scale_sh = data
        k = A.shape[1]
        if k <= 512:
            return _triton_gemm(data)
        else:
            return _aiter_gemm(data)
    else:
        return _aiter_gemm(data)
scrolls · 208 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