submission 665044
Foreverwonder · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 320 lines, June 9 Researcher Reciprocity License v1.0.
submission_v27_amd_fixed.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-665044?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:bb5d5e36d99ff643b0ccb0f1a95f4a50937d94954588e207481312bbcb78b23d
license declaredunknown
license concludedunknown
authorsForeverwonder
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ __attribute__((aligned(16))) 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_v27_amd_fixed.py320 lines
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
_RUNTIME = None
_NATIVE_MOD = None
_A_BUFFER_CACHE = {}
_CFG_ROWS = [
(256, 4, 2880, 512, 21, 0, 5.05, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 8, 2112, 7168, 21, 1, 8.50, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 16, 2112, 7168, 21, 1, 12.90, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 16, 3072, 1536, 21, 1, 8.80, "_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, 6.50, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 64, 3072, 1536, 21, 1, 9.20, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 256, 3072, 1536, 21, 1, 7.80, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
(256, 256, 2880, 512, 21, 0, 5.50, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
]
_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_FP32 = 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_FP32 - MBITS_FP4) + 1) << MBITS_FP32);
constexpr int32_t VAL_TO_ADD =
((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_FP32) + 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_FP32 - 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_FP32 - MBITS_FP4));
return static_cast<uint8_t>(sign_lp | (e2m1 & MAX_INT));
}
__device__ __forceinline__ int scale_shuffle_offset_fast(int row, int col, int scale_n_pad) {
int offs_0 = row >> 5;
int row_in_32 = row & 31;
int offs_1 = row_in_32 >> 4;
int offs_2 = row_in_32 & 15;
int offs_3 = col >> 3;
int col_in_8 = col & 7;
int offs_4 = col_in_8 >> 2;
int offs_5 = col_in_8 & 3;
return offs_1 + (offs_4 << 1) + (offs_2 << 2) + (offs_5 << 6) + (offs_3 << 8) + (offs_0 << 5) * scale_n_pad;
}
__device__ __forceinline__ float warp_reduce_max_amd(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
float other = __shfl_down(val, offset);
val = (val > other) ? val : other;
}
return val;
}
template <int GROUPS_PER_BLOCK, int BLOCK_SIZE>
__global__ void __launch_bounds__(BLOCK_SIZE) quant_mxfp4_a_kernel_amd_fixed(
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__ __attribute__((aligned(16))) float s_vals[32 * GROUPS_PER_BLOCK];
__shared__ __attribute__((aligned(4))) uint8_t s_scale[GROUPS_PER_BLOCK];
constexpr int GROUP_SIZE = 32;
constexpr int PAIRS_PER_GROUP = GROUP_SIZE >> 1;
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]);
}
float abs_val = __builtin_fabsf(value);
float max_abs = warp_reduce_max_amd(abs_val);
if (lane == 0) {
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_fast(row, scale_col, scale_n_pad);
scale_sh[scale_idx] = scale_byte;
}
}
int group_base = group_in_block << 5;
s_vals[group_base + lane] = value;
__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;
}
int lane2 = lane << 1;
float x0 = s_vals[group_base + lane2] * inv_scale;
float x1 = s_vals[group_base + lane2 + 1] * inv_scale;
uint8_t q0 = quantize_e2m1(x0);
uint8_t q1 = quantize_e2m1(x1);
int out_idx = (row * (k >> 1)) + (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 >> 5;
int scale_n_pad = static_cast<int>(scale_sh.size(1));
TORCH_CHECK((k & 63) == 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) >> 1), "x_fp4 col mismatch");
if (k >= 2048) {
dim3 grid((scale_n_valid + 3) >> 2, m);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(quant_mxfp4_a_kernel_amd_fixed<4, 128>),
grid, dim3(128), 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) >> 1, m);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(quant_mxfp4_a_kernel_amd_fixed<2, 64>),
grid, dim3(64), 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
);
}
}
"""
def _get_native_mod():
global _NATIVE_MOD
if _NATIVE_MOD is None:
_NATIVE_MOD = load_inline(
name="mxfp4_mm_native_hip_quant_v27_amd_fixed",
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
cfg_path = os.path.join(tempfile.gettempdir(), "aiter_mxfp4_mm_submission_v27.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"
)
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()
A_q, A_scale_sh = _native_quant_a(A, dtypes)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
scrolls · 320 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