submission 662817
darkness098600 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 347 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-662817?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:cdc7609b04ff5577801e7f410ea05b6384cf28c775a73ecc6c39be530f318858
license declaredunknown
license concludedunknown
authorsdarkness098600
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM - Inject tuned configs + custom quant kernel.shared-memory
__shared__ float s_vals[32 * GROUPS_PER_BLOCK];split-k
f.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")Kernel source
submission.py347 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM - Inject tuned configs + custom quant kernel.
Based on v20 optimization:
1. Inject tuned config CSV so aiter uses optimal kernels
2. Custom quant kernel for A (faster than aiter's dynamic_mxfp4_quant)
"""
import os
import tempfile
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# Tuned configs for benchmark shapes
_CFG_ROWS = [
(256, 4, 2880, 512, 21, 0, 5.05, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 16, 2112, 7168, 21, 1, 12.90, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 32, 4096, 512, 21, 0, 5.40, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 32, 2880, 512, 21, 0, 5.10, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 64, 7168, 2048, 21, 0, 8.50, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E"),
(256, 256, 3072, 1536, 21, 0, 10.0, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E"),
]
_INLINE_CPP = r"""
#include <torch/extension.h>
void quant_mxfp4_a(torch::Tensor x, torch::Tensor x_fp4, torch::Tensor scale_sh);
"""
_INLINE_HIP = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <cstdint>
namespace {
__device__ __forceinline__ uint32_t f32_as_u32(float x) {
union { float f; uint32_t u; } v;
v.f = x;
return v.u;
}
__device__ __forceinline__ float u32_as_f32(uint32_t x) {
union { float f; uint32_t u; } v;
v.u = x;
return v.f;
}
__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
__hip_bfloat16 v;
*reinterpret_cast<uint16_t*>(&v) = x;
return static_cast<float>(v);
}
__device__ __forceinline__ uint8_t quantize_e2m1(float x) {
constexpr int EXP_BIAS_FP32 = 127;
constexpr int EXP_BIAS_FP4 = 1;
constexpr int MBITS_F32 = 23;
constexpr int MBITS_FP4 = 1;
constexpr uint8_t MAX_INT = 0x7;
constexpr uint8_t SIGN_MASK = 0x8;
constexpr uint32_t MAGIC_ADDER = (1u << 21) - 1u;
constexpr float MAX_NORMAL = 6.0f;
constexpr float MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT =
static_cast<uint32_t>(((EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1) << MBITS_F32);
constexpr int32_t VAL_TO_ADD =
((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + static_cast<int32_t>(MAGIC_ADDER);
uint32_t ux = f32_as_u32(x);
uint8_t sign_lp = static_cast<uint8_t>((ux >> 28) & SIGN_MASK);
ux &= 0x7FFFFFFFu;
float ax = u32_as_f32(ux);
if (ax >= MAX_NORMAL) {
return static_cast<uint8_t>(sign_lp | MAX_INT);
}
if (ax < MIN_NORMAL) {
float denormal_f = ax + u32_as_f32(DENORM_MASK_INT);
uint32_t denormal_u = f32_as_u32(denormal_f);
uint8_t denormal_x = static_cast<uint8_t>(denormal_u - DENORM_MASK_INT);
return static_cast<uint8_t>(sign_lp | (denormal_x & MAX_INT));
}
uint32_t mant_odd = (ux >> (MBITS_F32 - MBITS_FP4)) & 1u;
uint32_t normal_x = ux + static_cast<uint32_t>(VAL_TO_ADD) + mant_odd;
uint8_t e2m1 = static_cast<uint8_t>(normal_x >> (MBITS_F32 - MBITS_FP4));
return static_cast<uint8_t>(sign_lp | (e2m1 & MAX_INT));
}
__device__ __forceinline__ int scale_shuffle_offset(int row, int col, int scale_n_pad) {
int offs_0 = row / 32;
int row_in_32 = row % 32;
int offs_1 = row_in_32 / 16;
int offs_2 = row_in_32 % 16;
int offs_3 = col / 8;
int col_in_8 = col % 8;
int offs_4 = col_in_8 / 4;
int offs_5 = col_in_8 % 4;
return offs_1 + offs_4 * 2 + offs_2 * 4 + offs_5 * 64 + offs_3 * 256 + offs_0 * 32 * scale_n_pad;
}
template <int GROUPS_PER_BLOCK>
__global__ void quant_mxfp4_a_kernel(
const uint16_t* __restrict__ x,
uint8_t* __restrict__ x_fp4,
uint8_t* __restrict__ scale_sh,
int m,
int k,
int scale_n_valid,
int scale_n_pad
) {
__shared__ float s_vals[32 * GROUPS_PER_BLOCK];
__shared__ float s_abs[32 * GROUPS_PER_BLOCK];
__shared__ uint8_t s_scale[GROUPS_PER_BLOCK];
constexpr int GROUP_SIZE = 32;
constexpr int PAIRS_PER_GROUP = GROUP_SIZE / 2;
int tid = static_cast<int>(threadIdx.x);
int lane = tid & 31;
int group_in_block = tid >> 5;
int row = static_cast<int>(blockIdx.y);
int scale_col = static_cast<int>(blockIdx.x) * GROUPS_PER_BLOCK + group_in_block;
bool active_group = row < m && scale_col < scale_n_valid;
float value = 0.0f;
if (active_group) {
int x_idx = row * k + scale_col * GROUP_SIZE + lane;
value = bf16_to_f32(x[x_idx]);
}
int group_base = group_in_block * GROUP_SIZE;
float stored = active_group ? value : 0.0f;
s_vals[group_base + lane] = stored;
s_abs[group_base + lane] = fabsf(stored);
__syncthreads();
if (lane < 16) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 16];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane < 8) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 8];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane < 4) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 4];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane < 2) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 2];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane == 0) {
float a = s_abs[group_base];
float b = s_abs[group_base + 1];
float max_abs = a > b ? a : b;
uint32_t max_bits = f32_as_u32(max_abs);
max_bits = (max_bits + 0x00200000u) & 0xFF800000u;
uint8_t scale_byte = 0;
if (max_bits != 0) {
scale_byte = static_cast<uint8_t>(((max_bits >> 23) & 0xFFu) - 2u);
}
s_scale[group_in_block] = scale_byte;
if (active_group) {
int scale_idx = scale_shuffle_offset(row, scale_col, scale_n_pad);
scale_sh[scale_idx] = scale_byte;
}
}
__syncthreads();
if (lane < PAIRS_PER_GROUP && active_group) {
uint8_t scale_byte = s_scale[group_in_block];
float inv_scale = 0.0f;
if (scale_byte != 0) {
float scale_f = u32_as_f32(static_cast<uint32_t>(scale_byte) << 23);
inv_scale = 1.0f / scale_f;
}
float x0 = s_vals[group_base + lane * 2] * inv_scale;
float x1 = s_vals[group_base + lane * 2 + 1] * inv_scale;
uint8_t q0 = quantize_e2m1(x0);
uint8_t q1 = quantize_e2m1(x1);
int out_idx = row * (k / 2) + scale_col * PAIRS_PER_GROUP + lane;
x_fp4[out_idx] = static_cast<uint8_t>(q0 | (q1 << 4));
}
}
} // namespace
void quant_mxfp4_a(torch::Tensor x, torch::Tensor x_fp4, torch::Tensor scale_sh) {
TORCH_CHECK(x.is_cuda(), "x must be a CUDA tensor");
TORCH_CHECK(x_fp4.is_cuda(), "x_fp4 must be a CUDA tensor");
TORCH_CHECK(scale_sh.is_cuda(), "scale_sh must be a CUDA tensor");
TORCH_CHECK(x.scalar_type() == at::ScalarType::BFloat16, "x must be bf16");
TORCH_CHECK(x_fp4.scalar_type() == at::ScalarType::Byte, "x_fp4 must be uint8");
TORCH_CHECK(scale_sh.scalar_type() == at::ScalarType::Byte, "scale_sh must be uint8");
TORCH_CHECK(x.dim() == 2, "x must be 2D");
TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
TORCH_CHECK(x_fp4.is_contiguous(), "x_fp4 must be contiguous");
TORCH_CHECK(scale_sh.is_contiguous(), "scale_sh must be contiguous");
int m = static_cast<int>(x.size(0));
int k = static_cast<int>(x.size(1));
int scale_n_valid = k / 32;
int scale_n_pad = static_cast<int>(scale_sh.size(1));
TORCH_CHECK(k % 64 == 0, "k must be divisible by 64");
TORCH_CHECK(x_fp4.size(0) == x.size(0), "x_fp4 row mismatch");
TORCH_CHECK(x_fp4.size(1) == x.size(1) / 2, "x_fp4 col mismatch");
if (k >= 2048) {
dim3 grid((scale_n_valid + 3) / 4, m);
dim3 block(128);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(quant_mxfp4_a_kernel<4>),
grid, block, 0, 0,
reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>()),
x_fp4.data_ptr<uint8_t>(),
scale_sh.data_ptr<uint8_t>(),
m, k, scale_n_valid, scale_n_pad
);
} else {
dim3 grid((scale_n_valid + 1) / 2, m);
dim3 block(64);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(quant_mxfp4_a_kernel<2>),
grid, block, 0, 0,
reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>()),
x_fp4.data_ptr<uint8_t>(),
scale_sh.data_ptr<uint8_t>(),
m, k, scale_n_valid, scale_n_pad
);
}
}
"""
_RUNTIME = None
_NATIVE_MOD = None
_A_BUFFER_CACHE = {}
def _get_native_mod():
global _NATIVE_MOD
if _NATIVE_MOD is None:
_NATIVE_MOD = load_inline(
name="mxfp4_mm_native_hip_quant_v20",
cpp_sources=[_INLINE_CPP],
cuda_sources=[_INLINE_HIP],
functions=["quant_mxfp4_a"],
extra_cuda_cflags=["--offload-arch=gfx950", "-O3", "-std=c++20"],
verbose=False,
)
return _NATIVE_MOD
def _ensure_runtime():
global _RUNTIME
if _RUNTIME is not None:
return _RUNTIME
# Create tuned config CSV
cfg_path = os.path.join(tempfile.gettempdir(), "aiter_mxfp4_mm_submission.csv")
if not os.path.exists(cfg_path):
with open(cfg_path, "w", encoding="ascii", newline="") as f:
f.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")
for cu_num, m, n, k, kernel_id, split_k, us, kernel_name in _CFG_ROWS:
f.write(
f"{cu_num},{m},{n},{k},{kernel_id},{split_k},{us},{kernel_name},0.0,0.0,0.0\n"
)
# Override aiter config
default_cfg = "/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
os.environ["AITER_CONFIG_GEMM_A4W4"] = cfg_path + os.pathsep + default_cfg
import aiter
from aiter import dtypes
_RUNTIME = (aiter, dtypes)
return _RUNTIME
def _native_quant_a(A: torch.Tensor, dtypes):
mod = _get_native_mod()
m, k = A.shape
scale_m_pad = ((m + 255) // 256) * 256
scale_n_valid = k // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
cache_key = (A.device.index, m, k)
cached = _A_BUFFER_CACHE.get(cache_key)
if cached is None:
A_q_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=A.device)
A_scale_u8 = torch.zeros((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=A.device)
cached = (
A_q_u8,
A_scale_u8,
A_q_u8.view(dtypes.fp4x2),
A_scale_u8.view(dtypes.fp8_e8m0),
)
_A_BUFFER_CACHE[cache_key] = cached
A_q_u8, A_scale_u8, A_q_view, A_scale_view = cached
mod.quant_mxfp4_a(A, A_q_u8, A_scale_u8)
return A_q_view, A_scale_view
def custom_kernel(data: input_t) -> output_t:
aiter, dtypes = _ensure_runtime()
A, _, _, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
# Use custom quant kernel (faster than aiter's dynamic_mxfp4_quant)
A_q, A_scale_sh = _native_quant_a(A, dtypes)
# Use aiter GEMM (now with tuned config)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
scrolls · 347 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