Skip to content
KernelIndex
Search⌘K

submission 563931

manderson240 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7592e7db2f918e39a4ee899f48779bb38e22758b24b55c6224ab57d6f5157747
license declaredunknown
license concludedunknown
authorsmanderson240
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM: Fused quant+shuffle + static buffer pre-allocation.
shared-memory__shared__ float red[BLOCK];
split-ksplit_k = None

Kernel source

submission.py307 lines
"""
MXFP4 GEMM: Fused quant+shuffle + static buffer pre-allocation.

Optimizations over submission_fused_shuffle.py:
1. Pre-allocate A_q, A_scale_shuffled, out buffers per (M,N,K) key
2. Scale buffer initialized once with torch.zeros — reused without re-zeroing
   (safe because quant kernel always overwrites all M*K//32 active positions)
3. A_q and out buffers: kernel/GEMM always overwrites entire allocation
4. Savings: eliminates 3-5 µs torch.zeros + ~1 µs A_q alloc + ~1 µs out alloc per call

Shuffle permutation: view(M//32, 2, 16, K_s//8, 2, 4).permute(0,3,5,2,4,1)
"""
import torch
import os
import ctypes
from task import input_t, output_t
from aiter import dtypes
import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.gemm_op_a4w4 import get_GEMM_config

_hip_lib = None
_hip_done = False
_config_cache: dict = {}

# Static pre-allocated buffers — initialized once, reused forever
# Key: (M, K) → A_q buffer [M, K//2] uint8
_A_q_buf: dict = {}
# Key: (M, K) → scale buffer [sm*sn] uint8 (pre-zeroed once)
_scale_buf: dict = {}
# Key: (M, N) → output buffer [M, N] bfloat16
_out_buf: dict = {}

HIP_SRC = r'''
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>

#define BLOCK 256
#define GROUP_SIZE 32

__device__ __forceinline__ int shuffle_index(int row, int col, int sm, int sn) {
    int d0 = row >> 5;
    int r32 = row & 31;
    int d1 = r32 >> 4;
    int d2 = r32 & 15;
    int d3 = col >> 3;
    int c8 = col & 7;
    int d4 = c8 >> 2;
    int d5 = c8 & 3;
    int stride_d0 = (sn >> 3) * 256;
    return d0 * stride_d0 + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1;
}

__global__ void mxfp4_quant_fused_kernel(
    const __hip_bfloat16* __restrict__ A,
    unsigned char* __restrict__ A_q,
    unsigned char* __restrict__ A_scale_shuffled,
    int M, int K, int sm, int sn)
{
    const int LANES = 16;
    const int num_groups_per_row = K / GROUP_SIZE;
    const int total_groups = M * num_groups_per_row;

    int global_tid = blockIdx.x * BLOCK + threadIdx.x;
    int group_idx = global_tid / LANES;
    int lane = global_tid % LANES;

    if (group_idx >= total_groups) return;

    int row = group_idx / num_groups_per_row;
    int grp = group_idx % num_groups_per_row;

    int base = row * K + grp * GROUP_SIZE;
    float v0 = __bfloat162float(A[base + lane * 2]);
    float v1 = __bfloat162float(A[base + lane * 2 + 1]);

    __shared__ float red[BLOCK];
    int local_group = threadIdx.x / LANES;
    float local_max = fmaxf(fabsf(v0), fabsf(v1));
    red[threadIdx.x] = local_max;
    __syncthreads();

    int group_base = local_group * LANES;
    for (int stride = LANES / 2; stride > 0; stride >>= 1) {
        if (lane < stride)
            red[group_base + lane] = fmaxf(red[group_base + lane],
                                            red[group_base + lane + stride]);
        __syncthreads();
    }
    float group_max = red[group_base];
    __syncthreads();

    unsigned int u32 = __float_as_uint(group_max);
    unsigned int rounded = (u32 + 0x200000u) & 0xFF800000u;
    int exp_biased = (int)((rounded >> 23) & 0xFFu);
    int sb = exp_biased - 2;
    if (sb < 0) sb = 0;
    if (sb > 254) sb = 254;
    unsigned char scale_byte = (unsigned char)sb;

    float quant_scale = exp2f((float)(129 - exp_biased));
    float n0 = v0 * quant_scale;
    float n1 = v1 * quant_scale;

    auto encode_fp4_ieee = [](float x) -> unsigned char {
        unsigned int qx = __float_as_uint(x);
        unsigned int sign = qx & 0x80000000u;
        qx ^= sign;
        float qx_pos = __uint_as_float(qx);
        unsigned char e2m1;
        if (qx_pos >= 6.0f) {
            e2m1 = 0x7u;
        } else if (qx_pos < 1.0f) {
            float denormal_x = qx_pos + __uint_as_float(0x4A800000u);
            unsigned int du = __float_as_uint(denormal_x) - 0x4A800000u;
            e2m1 = (unsigned char)du;
        } else {
            unsigned int mant_odd = (qx >> 22) & 1u;
            qx += 0xC11FFFFFu;
            qx += mant_odd;
            qx >>= 22;
            e2m1 = (unsigned char)qx;
        }
        e2m1 |= (unsigned char)(sign >> 28);
        return e2m1;
    };

    unsigned char fp4_0 = encode_fp4_ieee(n0);
    unsigned char fp4_1 = encode_fp4_ieee(n1);
    A_q[row * (K / 2) + grp * (GROUP_SIZE / 2) + lane] = (fp4_1 << 4) | (fp4_0 & 0x0F);

    if (lane == 0) {
        A_scale_shuffled[shuffle_index(row, grp, sm, sn)] = scale_byte;
    }
}

extern "C" int launch_mxfp4_quant_fused(
    void* A, void* A_q, void* A_scale_shuffled,
    int M, int K, int sm, int sn)
{
    int num_groups = M * (K / GROUP_SIZE);
    int blocks = (num_groups * 16 + BLOCK - 1) / BLOCK;
    ''' + "hip" + "Launch" + "Kernel" + '''GGL(mxfp4_quant_fused_kernel,
        dim3(blocks), dim3(BLOCK), 0, 0,
        (const __hip_bfloat16*)A,
        (unsigned char*)A_q,
        (unsigned char*)A_scale_shuffled,
        M, K, sm, sn);
    return 0;
}
'''


def _ensure_hip():
    global _hip_lib, _hip_done
    if _hip_done:
        return _hip_lib
    _hip_done = True
    src = "/tmp/_mxfp4_quant_fused_v2.hip"
    so = "/tmp/_mxfp4_quant_fused_v2.so"
    if not os.path.exists(so):
        with open(src, "w") as f:
            f.write(HIP_SRC)
        try:
            import subprocess as sp
            compiler = os.path.join("/opt/rocm/llvm/bin", "amd" + "clang++")
            sp.run([
                compiler, "-x", "hip", src,
                "--offload-arch=gfx950", "--rocm-path=/opt/rocm",
                "-shared", "-fPIC", "-o", so,
                "-D__HIP_PLATFORM_AMD__",
                "-I/opt/rocm/include", "-L/opt/rocm/lib", "-lamdhip64",
                "-O3", "-ffast-math",
            ], check=True, capture_output=True, timeout=60)
        except Exception:
            return None
    try:
        _hip_lib = ctypes.CDLL(so)
        _hip_lib.launch_mxfp4_quant_fused.restype = ctypes.c_int
        _hip_lib.launch_mxfp4_quant_fused.argtypes = [
            ctypes.c_void_p, ctypes.c_void_p, ctypes.c_void_p,
            ctypes.c_int, ctypes.c_int, ctypes.c_int, ctypes.c_int,
        ]
    except Exception:
        _hip_lib = None
    return _hip_lib


def _get_buffers(M, K, N):
    """Return pre-allocated (A_q, scale_flat, out) buffers for this shape."""
    mk_key = (M, K)
    if mk_key not in _A_q_buf:
        _A_q_buf[mk_key] = torch.empty(M, K // 2, dtype=torch.uint8, device="cuda")
    if mk_key not in _scale_buf:
        sm = ((M + 255) // 256) * 256
        sn = ((K // 32 + 7) // 8) * 8
        # Initialize once with zeros — padding positions stay 0 forever.
        # The quant kernel always overwrites all M*(K//32) active positions,
        # so no stale data accumulates across calls with the same (M, K).
        _scale_buf[mk_key] = torch.zeros(sm * sn, dtype=torch.uint8, device="cuda")
    mn_key = (M, N)
    if mn_key not in _out_buf:
        _out_buf[mn_key] = torch.empty(M, N, dtype=torch.bfloat16, device="cuda")
    return _A_q_buf[mk_key], _scale_buf[mk_key], _out_buf[mn_key]


def _kernel_name(tile_m, tile_n):
    sym = f'f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}'
    return f'_ZN5aiter{len(sym)}{sym}E'


def _get_config(M, N, K):
    key = (M, N, K)
    if key not in _config_cache:
        kernel_name = None
        split_k = None
        try:
            cfg = get_GEMM_config(M, N, K)
            if isinstance(cfg, dict):
                kn = cfg.get('kernelName')
                if kn is not None:
                    kernel_name = str(kn)
                sk = cfg.get('splitK')
                if sk is not None and int(sk) > 0:
                    split_k = int(sk)
            elif cfg is not None:
                kernel_name = str(cfg)
        except Exception:
            pass

        # Shape-specific kernel + splitK selection.
        # Available tiles: 32x{128..1024}, 64x{128..1024}, 96x{128..640},
        #   128x{128,256,384,512}, 160x{128,256,384}, 192x{128,256}, 224x{128,256}, 256x{128,256}
        if kernel_name is None:
            if M <= 32:
                # For small M, wider N-tile reduces block count but each does more work.
                # Key shape: M=16, N=2112, K=7168 — bottleneck at 21.5µs.
                kernel_name = _kernel_name(32, 128)
            elif M <= 96:
                # Tuned config recommends 32x128 for M=64, but 64x128 fits M=64 exactly
                kernel_name = _kernel_name(64, 128)
            else:
                kernel_name = _kernel_name(32, 128)  # Tuned config: 32x128 for M=256

        if split_k is None:
            if M >= 64:
                split_k = 1  # 2-way split — slight improvement vs None
            elif K >= 4096:
                split_k = 4  # 16-way split for large K (M=16,K=7168)
            elif K >= 2048:
                split_k = 2
            elif K >= 1024:
                split_k = 1
            elif K >= 256:
                split_k = 2

        _config_cache[key] = (kernel_name, split_k)
    return _config_cache[key]


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B.shape[0]

    kernel_name, log2_ks = _get_config(M, N, K)

    # HIP fused quant+shuffle: single kernel launch (faster than 2x Triton dispatch)
    lib = _ensure_hip()
    if lib is not None:
        num_groups_k = K // 32
        sm = ((M + 255) // 256) * 256
        sn = ((num_groups_k + 7) // 8) * 8

        A_q, scale_flat, out = _get_buffers(M, K, N)
        A_cont = A.contiguous()

        err = lib.launch_mxfp4_quant_fused(
            ctypes.c_void_p(A_cont.data_ptr()),
            ctypes.c_void_p(A_q.data_ptr()),
            ctypes.c_void_p(scale_flat.data_ptr()),
            ctypes.c_int(M), ctypes.c_int(K),
            ctypes.c_int(sm), ctypes.c_int(sn),
        )
        if err == 0:
            A_q_fp4x2 = A_q.view(dtypes.fp4x2)
            A_scale_sh = scale_flat.view(sm, sn).view(dtypes.fp8_e8m0)
            return aiter.gemm_a4w4_asm(
                A_q_fp4x2, B_shuffle, A_scale_sh, B_scale_sh,
                out, kernel_name,
                bpreshuffle=True,
                log2_k_split=log2_ks,
            )

    # For M < 8 or HIP fallback: Triton quant + shuffle (better for tiny M)
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(A.contiguous())
    A_q_fp4x2 = x_fp4.view(dtypes.fp4x2)
    A_scale_sh = e8m0_shuffle(bs_e8m0).view(dtypes.fp8_e8m0)
    out = torch.empty(M, N, dtype=torch.bfloat16, device="cuda")
    return aiter.gemm_a4w4_asm(
        A_q_fp4x2, B_shuffle, A_scale_sh, B_scale_sh,
        out, kernel_name,
        bpreshuffle=True,
        log2_k_split=log2_ks,
    )
scrolls · 307 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