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
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.
fp4
MXFP4 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