Skip to content
KernelIndex
Search⌘K

submission 600312

fluudgate · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:78ba0f912d0c7578379765953b980425f6feb31bdc725c8e8f322ebc6a94bc61
license declaredunknown
license concludedunknown
authorsfluudgate
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM — Exp 6: Address real bottlenecks.
split-kd[(cu, 4, 2880, 512)] = {**base, 'kernelId': 21, 'splitK': 0, 'kernelName': k32}
tile-m = 4BLOCK_M = 4
tile-n = 128BLOCK_N = 128

Kernel source

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

"""
MXFP4 GEMM — Exp 6: Address real bottlenecks.
1. Cache workspaces (no per-call allocation/memset)
2. Remove log2/exp2 — use exponent bit extraction
3. Aligned fast path (no masks for benchmark shapes)
4. Complete 6-shape GEMM config sweep (inject all 6)
5. Keep fused quant+shuffle (still saves 1 launch)
"""

import os
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_USE_NT"] = "1"

import torch
import triton
import triton.language as tl
from task import input_t, output_t


# ---------------------------------------------------------------------------
# Fused MXFP4 quant + e8m0 shuffle — optimized
# - Exponent-bit scale extraction (no log2/exp2)
# - Aligned fast path (no masks)
# ---------------------------------------------------------------------------

@triton.jit
def _mxfp4_quant_op_fast(
    x,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    """Optimized FP4 quant: exponent-bit extraction instead of log2/exp2."""
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    # Block max — round to power of 2
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax_int = amax.to(tl.int32, bitcast=True)
    amax_int = (amax_int + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000

    # Extract exponent directly from integer bits (no log2!)
    # For a power-of-2 float: exponent = (amax_int >> 23) - 127
    # scale_e8m0_unbiased = exponent - 2
    raw_exp = (amax_int >> 23).to(tl.int32)
    scale_e8m0_unbiased = raw_exp - 127 - 2
    # tl.clamp doesn't support int32 — use manual min/max
    scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127, scale_e8m0_unbiased)
    scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased > 127, 127, scale_e8m0_unbiased)
    bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.uint8)

    # Reconstruct quant_scale as 2^(-scale_e8m0_unbiased) via integer bit construction (no exp2!)
    quant_exp = (-scale_e8m0_unbiased + 127).to(tl.uint32)
    quant_scale_int = quant_exp << 23
    quant_scale = quant_scale_int.to(tl.float32, bitcast=True)

    qx = x * quant_scale
    qx = qx.to(tl.uint32, bitcast=True)
    s = qx & 0x80000000
    qx = qx ^ s
    qx_fp32 = qx.to(tl.float32, bitcast=True)

    saturate_mask = qx_fp32 >= max_normal
    denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
    normal_mask = ~(saturate_mask | denormal_mask)

    # Denormal path
    denorm_exp: tl.constexpr = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
    denormal_x = qx_fp32 + denorm_mask_float
    denormal_x = denormal_x.to(tl.uint32, bitcast=True)
    denormal_x -= denorm_mask_int
    denormal_x = denormal_x.to(tl.uint8)

    # Normal path
    normal_x = qx
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
    normal_x += val_to_add
    normal_x += mant_odd
    normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
    normal_x = normal_x.to(tl.uint8)

    # Merge
    e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
    e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1_value = e2m1_value | sign_lp

    # Pack pairs
    e2m1_value = tl.reshape(
        e2m1_value,
        [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2],
    )
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.jit
def _fused_quant_shuffle_aligned(
    x_ptr, x_fp4_ptr, bs_ptr,
    stride_x_m, stride_x_n, stride_fp4_m, stride_fp4_n,
    M, N, scaleN_pad,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    """Aligned fast path: no masks, no validity checks. For benchmark shapes only."""
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    stride_xm = tl.cast(stride_x_m, tl.int64)
    stride_xn = tl.cast(stride_x_n, tl.int64)
    stride_fm = tl.cast(stride_fp4_m, tl.int64)
    stride_fn = tl.cast(stride_fp4_n, tl.int64)
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
    x_offs = x_offs_m[:, None] * stride_xm + x_offs_n[None, :] * stride_xn
    # No mask — shapes are aligned
    x = tl.load(x_ptr + x_offs).to(tl.float32)

    out_tensor, bs_e8m0 = _mxfp4_quant_op_fast(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)

    fp4_offs_n = pid_n * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
    fp4_offs = x_offs_m[:, None] * stride_fm + fp4_offs_n[None, :] * stride_fn
    tl.store(x_fp4_ptr + fp4_offs, out_tensor)

    # Shuffled scale store
    bs_m = x_offs_m
    bs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
    i0 = bs_m[:, None] // 32
    i1 = (bs_m[:, None] % 32) // 16
    i2 = bs_m[:, None] % 16
    i3 = bs_n[None, :] // 8
    i4 = (bs_n[None, :] % 8) // 4
    i5 = bs_n[None, :] % 4
    shuffled_offs = i1 + i4 * 2 + i2 * 4 + i5 * 64 + i3 * 256 + i0 * (32 * scaleN_pad)
    tl.store(bs_ptr + shuffled_offs, bs_e8m0)


# ---------------------------------------------------------------------------
# Workspace cache — avoid per-call allocation/memset
# ---------------------------------------------------------------------------

_workspace_cache = {}


def _get_workspace(M, K, device, dtypes):
    key = (M, K)
    if key in _workspace_cache:
        return _workspace_cache[key]

    MXFP4_BLOCK = 32
    scaleN_valid = triton.cdiv(K, MXFP4_BLOCK)
    M_p = triton.cdiv(M, 256) * 256
    N_sp = triton.cdiv(scaleN_valid, 8) * 8

    x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
    # Initialize padded scale buffer ONCE with 127
    bs_shuffled = torch.full((M_p * N_sp,), 127, dtype=torch.uint8, device=device)

    _workspace_cache.clear()  # Only cache one shape at a time
    _workspace_cache[key] = (x_fp4, bs_shuffled, scaleN_valid, M_p, N_sp)
    return _workspace_cache[key]


# ---------------------------------------------------------------------------
# Initialization
# ---------------------------------------------------------------------------

_inited = False
_aiter = None
_bf16 = None
_dtypes = None


def _ensure_init():
    global _inited, _aiter, _bf16, _dtypes
    if _inited:
        return
    _inited = True

    import aiter
    from aiter import dtypes
    _aiter = aiter
    _bf16 = dtypes.bf16
    _dtypes = dtypes

    # EVOLVE-BLOCK-START gemm_config_patch
    # Inject configs for ALL 6 benchmark shapes
    try:
        from aiter.ops.gemm_op_a4w4 import get_GEMM_config
        _ = get_GEMM_config(1, 1, 1)
        if hasattr(get_GEMM_config, "gemm_dict"):
            d = get_GEMM_config.gemm_dict
            cu = 256
            k32 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E'
            k64 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E'
            k96 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_96x128E'
            k128 = '_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_128x128E'
            base = {'us': 0, 'tflops': 0, 'bw': 0, 'errRatio': 0}

            # Small M: bandwidth-optimized (32x128)
            d[(cu, 4, 2880, 512)] = {**base, 'kernelId': 21, 'splitK': 0, 'kernelName': k32}
            d[(cu, 16, 2112, 7168)] = {**base, 'kernelId': 21, 'splitK': 0, 'kernelName': k32}
            # Medium M: transitional (64x128)
            d[(cu, 32, 4096, 512)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}
            d[(cu, 32, 2880, 512)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}
            # Large M: compute-optimized (96x128 or 128x128)
            d[(cu, 64, 7168, 2048)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}
            d[(cu, 256, 3072, 1536)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}

            get_GEMM_config.cache_clear()
    except Exception:
        pass
    # EVOLVE-BLOCK-END gemm_config_patch


MXFP4_BLOCK = 32
BLOCK_M = 4
BLOCK_N = 128


# EVOLVE-BLOCK-START gemm_dispatch
def custom_kernel(data: input_t) -> output_t:
    _ensure_init()
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape

    # Cached workspace — no per-call allocation or memset
    x_fp4, bs_shuffled, scaleN_valid, M_p, N_sp = _get_workspace(M, K, A.device, _dtypes)

    grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, BLOCK_N))
    _fused_quant_shuffle_aligned[grid](
        A, x_fp4, bs_shuffled,
        A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1),
        M, K, N_sp,
        BLOCK_SIZE_M=BLOCK_M, BLOCK_SIZE_N=BLOCK_N, MXFP4_QUANT_BLOCK_SIZE=MXFP4_BLOCK,
    )

    A_q = x_fp4.view(_dtypes.fp4x2)
    A_scale_sh = bs_shuffled.view(M_p, N_sp).view(_dtypes.fp8_e8m0)

    return _aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=_bf16, bpreshuffle=True,
    )
# EVOLVE-BLOCK-END gemm_dispatch
scrolls · 267 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