Skip to content
KernelIndex
Search⌘K

submission 747026

SuminBae · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-747026?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
14.1µs
#488 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5048ca3c3ec6d1eb3a6e8436a0bd3312d95d9e7917e691d234f970b7f17f62d5
license declaredunknown
license concludedunknown
authorsSuminBae
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM v11c — Fused quant+shuffle + ASM GEMM via hipModuleLoad.

Kernel source

submission.py314 lines
"""
MXFP4 GEMM v11c — Fused quant+shuffle + ASM GEMM via hipModuleLoad.
Single C++ call: fused HIP quant/e8m0_shuffle + ASM GEMM launch.
Eliminates Python overhead and minimizes kernel launches.
"""
from task import input_t, output_t
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
import torch
from torch.utils.cpp_extension import load_inline

# Find aiter .co file directory
import aiter
_aiter_root = os.path.dirname(os.path.dirname(aiter.__file__))
_co_dir = os.path.join(_aiter_root, 'hsa', 'gfx950', 'f4gemm')

_cuda_src = r"""
#include <torch/extension.h>
#include <cstdint>
#include <vector>
#include <cmath>
#include <string>
#include <unordered_map>
#include <hip/hip_runtime.h>

#define HIP_CHECK(call) do { \
    hipError_t err = call; \
    if (err != hipSuccess) { \
        TORCH_CHECK(false, "HIP error in ", #call, ": ", hipGetErrorString(err)); \
    } \
} while(0)

__host__ __device__ __forceinline__ int ceildiv(int a, int b) {
    return (a + b - 1) / b;
}

__device__ __forceinline__ float bf16_to_float(const void* ptr, int idx) {
    uint16_t raw = reinterpret_cast<const uint16_t*>(ptr)[idx];
    uint32_t f32_bits = ((uint32_t)raw) << 16;
    return __uint_as_float(f32_bits);
}

// Compute flat index into e8m0-shuffled scale tensor
__device__ __forceinline__ int shuffled_scale_idx(int row, int col, int Nb) {
    int d0 = row >> 5;
    int d1 = (row >> 4) & 1;
    int d2 = row & 15;
    int d3 = col >> 3;
    int d4 = (col >> 2) & 1;
    int d5 = col & 3;
    return d0 * (Nb << 8) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1;
}


// ════════════════════════════════════════════════════════════════════
// Fused Quantization + e8m0 Shuffle kernel
// Grid: (M, ceil(K/64)),  Block: 64
// Writes FP4 data to A_fp4 and shuffled scales directly to A_scale_sh
// ════════════════════════════════════════════════════════════════════

__global__ void mxfp4_quant_shuffle_kernel(
    const void* __restrict__ A,
    uint8_t* __restrict__ A_fp4,
    uint8_t* __restrict__ A_scale_sh,
    int M, int K,
    int stride_a, int stride_fp4,
    int scale_Nb
) {
    int row  = blockIdx.x;
    int half = threadIdx.x >> 5;
    int lane = threadIdx.x & 31;
    int blk  = blockIdx.y * 2 + half;
    int col  = blk * 32 + lane;

    float val = 0.0f, abs_val = 0.0f;
    if (row < M && col < K) {
        val = bf16_to_float(A, row * stride_a + col);
        abs_val = fabsf(val);
    }

    float amax = abs_val;
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1)
        amax = fmaxf(amax, __shfl_xor(amax, offset));

    uint32_t amax_u = __float_as_uint(amax);
    amax_u = (amax_u + 0x200000u) & 0xFF800000u;
    float log2_amax = floorf(log2f(__uint_as_float(amax_u))) - 2.0f;
    log2_amax = fminf(fmaxf(log2_amax, -127.0f), 127.0f);
    uint8_t scale_e8m0 = (uint8_t)((int)log2_amax + 127);

    // Write scale directly to shuffled position
    if (lane == 0 && row < M && blk < ceildiv(K, 32)) {
        int sh_idx = shuffled_scale_idx(row, blk, scale_Nb);
        A_scale_sh[sh_idx] = scale_e8m0;
    }

    float qx = val * exp2f(-log2_amax);
    uint32_t qx_u = __float_as_uint(qx);
    uint32_t sign = qx_u & 0x80000000u;
    qx_u ^= sign;
    float qx_pos = __uint_as_float(qx_u);

    uint8_t e2m1;
    if (qx_pos >= 6.0f) {
        e2m1 = 0x7;
    } else if (qx_pos < 1.0f) {
        const uint32_t DENORM_MASK_INT = 149u << 23;
        uint32_t du = __float_as_uint(qx_pos + __uint_as_float(DENORM_MASK_INT));
        e2m1 = (uint8_t)(du - DENORM_MASK_INT);
    } else {
        uint32_t nu = qx_u;
        uint32_t mant_odd = (nu >> 22) & 1;
        nu += ((1 - 127) << 23) + ((1 << 21) - 1);
        nu += mant_odd;
        e2m1 = (uint8_t)(nu >> 22);
    }
    e2m1 |= (uint8_t)(sign >> 28);
    if (row >= M || col >= K) e2m1 = 0;

    uint8_t partner = (uint8_t)__shfl_down((int)e2m1, 1);
    if ((lane & 1) == 0) {
        uint8_t packed = e2m1 | (partner << 4);
        int out_col = blk * 16 + (lane >> 1);
        if (row < M && out_col < K / 2)
            A_fp4[row * stride_fp4 + out_col] = packed;
    }
}


// ════════════════════════════════════════════════════════════════════
// ASM GEMM KernelArgs struct (matches aiter layout exactly)
// ════════════════════════════════════════════════════════════════════

struct p3 { unsigned int _p0, _p1, _p2; };
struct p2 { unsigned int _p0, _p1; };

struct __attribute__((packed)) KernelArgs {
    void* ptr_D;    p2 _p0;
    void* ptr_C;    p2 _p1;
    void* ptr_A;    p2 _p2;
    void* ptr_B;    p2 _p3;
    float alpha;    p3 _p4;
    float beta;     p3 _p5;
    unsigned int stride_D0; p3 _p6;
    unsigned int stride_D1; p3 _p7;
    unsigned int stride_C0; p3 _p8;
    unsigned int stride_C1; p3 _p9;
    unsigned int stride_A0; p3 _p10;
    unsigned int stride_A1; p3 _p11;
    unsigned int stride_B0; p3 _p12;
    unsigned int stride_B1; p3 _p13;
    unsigned int M; p3 _p14;
    unsigned int N; p3 _p15;
    unsigned int K; p3 _p16;
    void* ptr_ScaleA; p2 _p17;
    void* ptr_ScaleB; p2 _p18;
    unsigned int stride_ScaleA0; p3 _p19;
    unsigned int stride_ScaleA1; p3 _p20;
    unsigned int stride_ScaleB0; p3 _p21;
    unsigned int stride_ScaleB1; p3 _p22;
    int log2_k_split;
};


// ════════════════════════════════════════════════════════════════════
// ASM kernel cache
// ════════════════════════════════════════════════════════════════════

struct AsmKernel {
    hipModule_t module = nullptr;
    hipFunction_t func = nullptr;
};

static std::string g_co_dir;
static std::unordered_map<std::string, AsmKernel> g_kernel_cache;

void set_co_dir(const std::string& dir) {
    g_co_dir = dir;
}

hipFunction_t get_asm_func(const char* co_name, const char* func_name) {
    auto it = g_kernel_cache.find(func_name);
    if (it != g_kernel_cache.end()) return it->second.func;

    AsmKernel k;
    std::string path = g_co_dir + "/" + co_name;
    HIP_CHECK(hipModuleLoad(&k.module, path.c_str()));
    HIP_CHECK(hipModuleGetFunction(&k.func, k.module, func_name));
    g_kernel_cache[func_name] = k;
    return k.func;
}


// ════════════════════════════════════════════════════════════════════
// All-in-one: fused quant+shuffle + ASM GEMM
// ════════════════════════════════════════════════════════════════════

torch::Tensor mxfp4_gemm_full(
    torch::Tensor A,           // bf16 [M, K]
    torch::Tensor B_shuffle,   // [N, K/2] (pre-shuffled B, any dtype)
    torch::Tensor B_scale_sh,  // [padded, padded] (shuffled B_scale, any dtype)
    int64_t m, int64_t n, int64_t k
) {
    int M = (int)m, N = (int)n, K = (int)k;
    auto device = A.device();
    auto u8_opts = torch::TensorOptions().dtype(torch::kUInt8).device(device);
    auto bf16_opts = torch::TensorOptions().dtype(torch::kBFloat16).device(device);

    // --- Fused quantize + shuffle ---
    int scale_n = ceildiv(K, 32);
    int padded_m = ceildiv(M, 256) * 256;
    int padded_sn = ceildiv(scale_n, 8) * 8;
    int scale_Nb = padded_sn >> 3;

    auto A_fp4      = torch::empty({m, k / 2}, u8_opts);
    auto A_scale_sh = torch::empty({(int64_t)padded_m, (int64_t)padded_sn}, u8_opts);

    mxfp4_quant_shuffle_kernel<<<dim3(M, ceildiv(K, 64)), 64>>>(
        A.data_ptr(), A_fp4.data_ptr<uint8_t>(), A_scale_sh.data_ptr<uint8_t>(),
        M, K, (int)A.stride(0), (int)A_fp4.stride(0), scale_Nb);

    // --- Select ASM kernel ---
    const char* co_name;
    const char* func_name;
    int tile_M, tile_N;

    if (M <= 64) {
        co_name = "f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co";
        func_name = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E";
        tile_M = 32; tile_N = 128;
    } else if (M <= 128) {
        co_name = "f4gemm_bf16_per1x32Fp4_BpreShuffle_128x128.co";
        func_name = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_128x128E";
        tile_M = 128; tile_N = 128;
    } else {
        co_name = "f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128.co";
        func_name = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E";
        tile_M = 192; tile_N = 128;
    }

    // --- Allocate output ---
    int padded_M32 = ceildiv(M, 32) * 32;
    auto out = torch::empty({(int64_t)padded_M32, n}, bf16_opts);

    // --- Compute grid ---
    int gdx = ceildiv(N, tile_N);
    int gdy = ceildiv(M, tile_M);

    // --- Pack KernelArgs ---
    KernelArgs args;
    memset(&args, 0, sizeof(args));
    args.ptr_D      = out.data_ptr();
    args.ptr_C      = nullptr;
    args.ptr_A      = A_fp4.data_ptr();
    args.ptr_B      = B_shuffle.data_ptr();
    args.alpha      = 1.0f;
    args.beta       = 0.0f;
    args.stride_D0  = (unsigned int)N;
    args.stride_C0  = (unsigned int)N;
    args.stride_A0  = (unsigned int)K;
    args.stride_B0  = (unsigned int)K;
    args.M          = (unsigned int)M;
    args.N          = (unsigned int)N;
    args.K          = (unsigned int)K;
    args.ptr_ScaleA = A_scale_sh.data_ptr();
    args.ptr_ScaleB = B_scale_sh.data_ptr();
    args.stride_ScaleA0 = (unsigned int)padded_sn;
    args.stride_ScaleB0 = (unsigned int)B_scale_sh.stride(0);
    args.log2_k_split   = 0;

    // --- Launch ASM GEMM ---
    hipFunction_t func = get_asm_func(co_name, func_name);
    size_t arg_size = sizeof(args);
    void* config[] = {
        HIP_LAUNCH_PARAM_BUFFER_POINTER, &args,
        HIP_LAUNCH_PARAM_BUFFER_SIZE, &arg_size,
        HIP_LAUNCH_PARAM_END
    };
    HIP_CHECK(hipModuleLaunchKernel(func, gdx, gdy, 1, 256, 1, 1,
                                     0, 0, nullptr, (void**)config));

    return out.slice(0, 0, m);
}
"""

_cpp_src = """
#include <torch/extension.h>
#include <string>
void set_co_dir(const std::string& dir);
torch::Tensor mxfp4_gemm_full(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
    int64_t m, int64_t n, int64_t k);
"""

_ext = load_inline(
    name="mxfp4_v11c",
    cpp_sources=[_cpp_src],
    cuda_sources=[_cuda_src],
    functions=["set_co_dir", "mxfp4_gemm_full"],
    verbose=False,
    extra_cuda_cflags=["-O3"],
)

_ext.set_co_dir(_co_dir)


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n = B.shape[0]
    return _ext.mxfp4_gemm_full(A, B_shuffle, B_scale_sh, m, n, k)
scrolls · 314 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