Skip to content
KernelIndex
Search⌘K

submission 595946

Borui Xu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:046eea6b711d42fe1c07c7f986c9733dbe4361367ea847ce31192ad9bc9ddcc9
license declaredunknown
license concludedunknown
authorsBorui Xu
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM — Virtual M-padding for better GEMM occupancy.
split-kreturn {'kernelId': 21, 'splitK': 0, 'us': 10.0,

Kernel source

submission.py205 lines
"""
MXFP4 GEMM — Virtual M-padding for better GEMM occupancy.
For small M with large K (low block count), pad M to increase blocks.
The quant kernel handles virtual padding (zero rows beyond actual M).
GEMM gets more blocks → better memory latency hiding.
"""
from task import input_t, output_t
import torch
import aiter
from aiter import dtypes

_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_e8m0 = dtypes.fp8_e8m0
_K32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

# Config patch
try:
    from aiter.ops import gemm_op_a4w4 as _gmod
    _orig_cfg = _gmod.get_GEMM_config
    def _p(M, N, K):
        try:
            c = _orig_cfg(M, N, K)
            if c is not None: return c
        except: pass
        return {'kernelId': 21, 'splitK': 0, 'us': 10.0,
                'kernelName': _K32, 'tflops': 0, 'bw': 0, 'errRatio': 0}
    _gmod.get_GEMM_config = _p
except: pass

_hip = None
_asm_fn = None
_init_done = False
_cache = {}


def _build_hip():
    global _hip
    if _hip is not None: return _hip
    from torch.utils.cpp_extension import load_inline
    cuda_src = r'''
#include <torch/extension.h>
#include <math.h>
__device__ __forceinline__ unsigned f2u(float f) {
    union { float ff; unsigned uu; } c; c.ff = f; return c.uu;
}
__device__ __forceinline__ float u2f(unsigned u) {
    union { unsigned uu; float ff; } c; c.uu = u; return c.ff;
}
__device__ __forceinline__ uint8_t quant_fp4(float v) {
    uint8_t s = (v < 0.f) ? 8u : 0u;
    float a = fabsf(v);
    uint8_t c;
    if      (a > 5.0f)  c = 7;
    else if (a == 5.0f)  c = 6;
    else if (a > 3.5f)  c = 6;
    else if (a == 3.5f)  c = 6;
    else if (a > 2.5f)  c = 5;
    else if (a == 2.5f)  c = 4;
    else if (a > 1.75f) c = 4;
    else if (a == 1.75f) c = 4;
    else if (a > 1.25f) c = 3;
    else if (a == 1.25f) c = 2;
    else if (a > 0.75f) c = 2;
    else if (a == 0.75f) c = 2;
    else if (a > 0.25f) c = 1;
    else if (a == 0.25f) c = 0;
    else                 c = 0;
    return s | c;
}

// Quant kernel with virtual M padding
// Processes virtual_M rows, but only reads A for rows < actual_M (rest = 0)
__global__ void mxfp4_quant_shuffle_kernel(
    const uint16_t* __restrict__ A,
    uint8_t* __restrict__ A_fp4,
    uint8_t* __restrict__ A_scale_sh,
    int actual_M, int virtual_M, int K, int ngroups, int sn_pad) {
    int gtid = blockIdx.x * blockDim.x + threadIdx.x;
    int hw = gtid >> 5, lane = gtid & 31;
    if (hw >= virtual_M * ngroups) return;
    int row = hw / ngroups, gcol = hw % ngroups;

    // Virtual padding: zero for rows beyond actual_M
    float val = 0.f;
    if (row < actual_M)
        val = u2f((unsigned)A[row * K + gcol * 32 + lane] << 16);

    float mx = fabsf(val);
    mx = fmaxf(mx, __shfl_xor(mx, 16));
    mx = fmaxf(mx, __shfl_xor(mx, 8));
    mx = fmaxf(mx, __shfl_xor(mx, 4));
    mx = fmaxf(mx, __shfl_xor(mx, 2));
    mx = fmaxf(mx, __shfl_xor(mx, 1));
    uint8_t e8, fp4;
    if (mx == 0.f) { e8 = 0; fp4 = 0; }
    else {
        unsigned rounded = (f2u(mx) + 0x200000u) & 0xFF800000u;
        float su = floorf(log2f(u2f(rounded))) - 2.0f;
        su = fminf(fmaxf(su, -127.f), 127.f);
        e8 = (uint8_t)((int)su + 127);
        fp4 = quant_fp4(val * exp2f(-su));
    }
    uint8_t p = (uint8_t)__shfl_xor((int)fp4, 1);
    if ((lane & 1) == 0)
        A_fp4[row * (K >> 1) + (gcol << 4) + (lane >> 1)] =
            (fp4 & 0xFu) | ((p & 0xFu) << 4);
    if (lane == 0) {
        int d0=row>>5, d1=(row>>4)&1, d2=row&15, d3=gcol>>3, d4=(gcol>>2)&1, d5=gcol&3;
        A_scale_sh[d0*((sn_pad>>3)*256) + d3*256 + d5*64 + d2*4 + d4*2 + d1] = e8;
    }
}

void mxfp4_quant_vpad(torch::Tensor A, torch::Tensor fp4, torch::Tensor scale_sh,
                       int actual_M, int virtual_M) {
    int K = A.size(1), ngrp = K / 32;
    int sn_pad = scale_sh.size(1);
    int tot = virtual_M * ngrp * 32, BS = 256;
    mxfp4_quant_shuffle_kernel<<<(tot+BS-1)/BS, BS>>>(
        (const uint16_t*)A.data_ptr(),
        fp4.data_ptr<uint8_t>(), scale_sh.data_ptr<uint8_t>(),
        actual_M, virtual_M, K, ngrp, sn_pad);
}
'''
    cpp_src = "void mxfp4_quant_vpad(torch::Tensor, torch::Tensor, torch::Tensor, int, int);\n"
    _hip = load_inline(name="hip_vpad", cpp_sources=cpp_src,
        cuda_sources=cuda_src, functions=["mxfp4_quant_vpad"],
        extra_cuda_cflags=["-O3"], verbose=False)
    return _hip


def _init():
    global _asm_fn, _init_done
    if _init_done: return
    _init_done = True
    try:
        from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
        _asm_fn = gemm_a4w4_asm
    except: pass


def _compute_virtual_m(m, n, k):
    """Compute virtual M that gives good GEMM occupancy."""
    n_tiles = (n + 127) // 128
    # Current blocks = ceil(m/32) * n_tiles
    cur_blocks = ((m + 31) // 32) * n_tiles
    # Target: at least 64 blocks for decent occupancy
    if cur_blocks >= 64:
        return m
    # Pad M to get more M-tiles (each adds n_tiles blocks)
    target_m_tiles = max(2, (64 + n_tiles - 1) // n_tiles)
    virtual_m = target_m_tiles * 32
    return virtual_m


def _get_cache(m, n, k, virtual_m, device):
    key = (m, n, k, virtual_m)
    if key not in _cache:
        ngrp = k // 32
        sm_pad = ((virtual_m + 31) // 32) * 32
        sn_pad = ((ngrp + 7) // 8) * 8
        _cache[key] = {
            'fp4': torch.empty((virtual_m, k // 2), dtype=torch.uint8, device=device),
            'scale': torch.zeros((sm_pad, sn_pad), dtype=torch.uint8, device=device),
            'out': torch.empty((virtual_m, n), dtype=torch.bfloat16, device=device),
        }
    return _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]

    if not _init_done:
        from aiter.ops.triton.quant import dynamic_mxfp4_quant
        from aiter.utility.fp4_utils import e8m0_shuffle
        aq, asc = dynamic_mxfp4_quant(A)
        ash = e8m0_shuffle(asc)
        out = aiter.gemm_a4w4(aq.view(_fp4x2), B_shuffle,
                               ash.view(_e8m0), B_scale_sh,
                               dtype=_bf16, bpreshuffle=True)
        _init()
        return out

    hip = _build_hip()
    virtual_m = _compute_virtual_m(m, n, k)
    c = _get_cache(m, n, k, virtual_m, A.device)

    # Quant with virtual padding (no extra copies!)
    hip.mxfp4_quant_vpad(A, c['fp4'], c['scale'], m, virtual_m)

    # GEMM on virtual_m rows (more blocks → better occupancy)
    if _asm_fn is not None:
        _asm_fn(c['fp4'].view(_fp4x2), B_shuffle,
                c['scale'].view(_e8m0), B_scale_sh,
                c['out'], _K32, bpreshuffle=True)
    else:
        c['out'][:] = aiter.gemm_a4w4(
            c['fp4'].view(_fp4x2), B_shuffle,
            c['scale'].view(_e8m0), B_scale_sh,
            dtype=_bf16, bpreshuffle=True)

    return c['out'][:m, :n]
scrolls · 205 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