submission 628737
hongquant.17 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 365 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-628737?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:bb2993c634cdad2260a0d7ddf2367c4144d718a5f9d45e0b21a4a7497d161dae
license declaredunknown
license concludedunknown
authorshongquant.17
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
constexpr int kBlockElems = 32; // MXFP4 block sizeKernel source
submission.py365 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os
from functools import lru_cache
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
if "PYTORCH_ROCM_ARCH" not in os.environ:
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"
quant_src = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <stdint.h>
#include <vector>
#include <math.h>
constexpr int kWaveSize = 64;
constexpr int kGroupLanes = 16; // one 16-lane subgroup per 1x32 block
constexpr int kBlocksPerWave = 4; // 64 lanes / 16
constexpr int kBlockElems = 32; // MXFP4 block size
constexpr int kBytesPerBlock = 16; // 32 fp4 elems / 2 per byte
constexpr int kScalePadRows = 256;
constexpr int kScalePadCols = 8;
using native_bf16 = __bf16;
using native_bf16x2 = __attribute__((ext_vector_type(2))) native_bf16;
template <typename T>
__host__ __device__ __forceinline__ T round_up(T x, T a) {
return ((x + a - 1) / a) * a;
}
template <typename T>
__host__ __device__ __forceinline__ T ceil_div(T x, T y) {
return (x + y - 1) / y;
}
__host__ __device__ __forceinline__ uint32_t bitcast_f32_to_u32(float x) {
union { float f; uint32_t u; } v;
v.f = x;
return v.u;
}
__host__ __device__ __forceinline__ float bitcast_u32_to_f32(uint32_t x) {
union { uint32_t u; float f; } v;
v.u = x;
return v.f;
}
__device__ __forceinline__ native_bf16 bitcast_u16_to_native_bf16(uint16_t x) {
union { uint16_t u; native_bf16 b; } v;
v.u = x;
return v.b;
}
__host__ __device__ __forceinline__ float bf16_bits_to_f32(uint16_t x) {
return bitcast_u32_to_f32(uint32_t(x) << 16);
}
// Same E8M0 reconstruction AITER uses.
__host__ __device__ __forceinline__ float e8m0_to_f32_quant(uint8_t e8m0) {
uint32_t bits;
if (e8m0 == 0x00u) {
bits = 0x00400000u;
} else if (e8m0 == 0xFFu) {
bits = 0x7F800001u;
} else {
bits = uint32_t(e8m0) << 23;
}
return bitcast_u32_to_f32(bits);
}
// Match aiter.utility.fp4_utils.dynamic_mxfp4_quant scale selection:
// amax_bits = (amax_bits + 0x200000) & 0xFF800000
// scale_e8m0 = biased_exp(amax_rounded) - 2
__host__ __device__ __forceinline__ uint8_t choose_scale_e8m0_aiter(float amax) {
if (amax == 0.0f) return 0u;
uint32_t u = bitcast_f32_to_u32(amax);
u = (u + 0x00200000u) & 0xFF800000u;
int scale_biased = int((u >> 23) & 0xFFu) - 2;
if (scale_biased < 0) scale_biased = 0;
if (scale_biased > 0xFF) scale_biased = 0xFF;
return static_cast<uint8_t>(scale_biased);
}
// Exact mapping for:
// scale = scale.view(sm // 32, 2, 16, sn // 8, 2, 4)
// scale = scale.permute(0, 3, 5, 2, 4, 1).contiguous()
// scale = scale.view(sm, sn)
__host__ __device__ __forceinline__ uint64_t e8m0_shuffle_flat_index(
int row, int col, int sn_padded)
{
const int a = row / 32;
const int b = (row % 32) / 16;
const int c = row % 16;
const int d = col / 8;
const int e = (col % 8) / 4;
const int f = col % 4;
const int d_extent = sn_padded / 8;
return ((((uint64_t(a) * uint64_t(d_extent) + uint64_t(d)) * 4ull
+ uint64_t(f)) * 16ull + uint64_t(c)) * 2ull
+ uint64_t(e)) * 2ull + uint64_t(b);
}
__device__ __forceinline__ float subgroup16_max(float v, int lane) {
#pragma unroll
for (int mask = kGroupLanes >> 1; mask > 0; mask >>= 1) {
const int peer = (lane & ~(kGroupLanes - 1)) | ((lane & (kGroupLanes - 1)) ^ mask);
const float other = __shfl(v, peer, kWaveSize);
v = fmaxf(v, other);
}
return v;
}
__device__ __forceinline__ uint8_t pack_two_bf16_to_fp4x2(
uint16_t x0_bits, uint16_t x1_bits, float scale_f32)
{
native_bf16x2 src;
src[0] = bitcast_u16_to_native_bf16(x0_bits);
src[1] = bitcast_u16_to_native_bf16(x1_bits);
// opsel=0 -> write packed fp4x2 into byte 0 of returned u32
const unsigned packed =
__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0u, src, scale_f32, 0);
return static_cast<uint8_t>(packed & 0xFFu);
}
template <bool kShuffleScale>
__global__ __launch_bounds__(256)
void quant_bf16_to_mxfp4_kernel(
const uint16_t* __restrict__ x_bf16,
uint8_t* __restrict__ q, // [M, K/2]
uint8_t* __restrict__ s, // [M, K/32] or shuffled [Mp, Sp]
int M,
int K,
int64_t stride_xm,
int64_t stride_xn,
int64_t stride_qm,
int64_t stride_qn,
int64_t stride_sm,
int64_t stride_sn)
{
const int row = int(blockIdx.y);
if (row >= M) return;
const int tid = int(threadIdx.x);
const int lane = tid & (kWaveSize - 1);
const int wave_in_block = tid >> 6;
const int waves_per_block = int(blockDim.x) >> 6;
const int subgroup = lane / kGroupLanes;
const int lane16 = lane & (kGroupLanes - 1);
const int blocks_per_row = K / kBlockElems;
const int sn_padded = kShuffleScale ? round_up(blocks_per_row, kScalePadCols) : 0;
const int block_id =
((int(blockIdx.x) * waves_per_block + wave_in_block) * kBlocksPerWave) + subgroup;
if (block_id >= blocks_per_row) return;
const int x_col0 = block_id * kBlockElems + lane16 * 2 + 0;
const int x_col1 = block_id * kBlockElems + lane16 * 2 + 1;
const int64_t x_row_offset = int64_t(row) * stride_xm;
const int64_t q_row_offset = int64_t(row) * stride_qm;
const uint16_t x0_bits =
x_bf16[x_row_offset + int64_t(x_col0) * stride_xn];
const uint16_t x1_bits =
x_bf16[x_row_offset + int64_t(x_col1) * stride_xn];
const float x0 = bf16_bits_to_f32(x0_bits);
const float x1 = bf16_bits_to_f32(x1_bits);
const float local_amax = fmaxf(fabsf(x0), fabsf(x1));
const float block_amax = subgroup16_max(local_amax, lane);
const uint8_t scale_e8m0 = choose_scale_e8m0_aiter(block_amax);
uint8_t q_byte = 0;
if (block_amax != 0.0f) {
const float scale_f32 = e8m0_to_f32_quant(scale_e8m0);
q_byte = pack_two_bf16_to_fp4x2(x0_bits, x1_bits, scale_f32);
}
const int q_col = block_id * kBytesPerBlock + lane16;
q[q_row_offset + int64_t(q_col) * stride_qn] = q_byte;
if (lane16 == 0) {
if constexpr (!kShuffleScale) {
s[int64_t(row) * stride_sm + int64_t(block_id) * stride_sn] = scale_e8m0;
} else {
const uint64_t flat = e8m0_shuffle_flat_index(row, block_id, sn_padded);
const int srow = int(flat / uint64_t(sn_padded));
const int scol = int(flat % uint64_t(sn_padded));
s[int64_t(srow) * stride_sm + int64_t(scol) * stride_sn] = scale_e8m0;
}
}
}
template <bool kShuffleScale>
void launch_quantize(
const uint16_t* x_ptr,
uint8_t* q_ptr,
uint8_t* s_ptr,
int M,
int K,
int64_t stride_xm,
int64_t stride_xn,
int64_t stride_qm,
int64_t stride_qn,
int64_t stride_sm,
int64_t stride_sn)
{
constexpr int kThreads = 256;
constexpr int kBlocksPerCta = (kThreads / kWaveSize) * kBlocksPerWave;
dim3 threads(kThreads);
dim3 blocks(ceil_div(K / kBlockElems, kBlocksPerCta), M);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(quant_bf16_to_mxfp4_kernel<kShuffleScale>),
blocks,
threads,
0,
0,
x_ptr,
q_ptr,
s_ptr,
M,
K,
stride_xm,
stride_xn,
stride_qm,
stride_qn,
stride_sm,
stride_sn);
}
std::vector<torch::Tensor> quantize(torch::Tensor x, bool shuffle) {
TORCH_CHECK(x.is_cuda(), "x must be a CUDA/HIP tensor");
TORCH_CHECK(x.scalar_type() == at::kBFloat16, "x must be bf16");
TORCH_CHECK(x.dim() == 2, "x must be 2D [M, K]");
x = x.contiguous();
const auto M = x.size(0);
const auto K = x.size(1);
TORCH_CHECK(K % 32 == 0, "K must be a multiple of 32");
TORCH_CHECK(M <= INT_MAX, "M exceeds kernel int indexing range");
TORCH_CHECK(K <= INT_MAX, "K exceeds kernel int indexing range");
const c10::cuda::CUDAGuard device_guard(x.device());
auto u8_opts = x.options().dtype(at::kByte);
auto q = torch::empty({M, K / 2}, u8_opts);
torch::Tensor s;
if (shuffle) {
const auto Mp = round_up<int64_t>(M, kScalePadRows);
const auto Sp = round_up<int64_t>(K / 32, kScalePadCols);
s = torch::full({Mp, Sp}, 127, u8_opts); // pad value must be 127
} else {
s = torch::empty({M, K / 32}, u8_opts);
}
const auto* x_ptr = reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>());
auto* q_ptr = q.data_ptr<uint8_t>();
auto* s_ptr = s.data_ptr<uint8_t>();
const int M_i = static_cast<int>(M);
const int K_i = static_cast<int>(K);
if (shuffle) {
launch_quantize<true>(
x_ptr,
q_ptr,
s_ptr,
M_i,
K_i,
x.stride(0),
x.stride(1),
q.stride(0),
q.stride(1),
s.stride(0),
s.stride(1));
} else {
launch_quantize<false>(
x_ptr,
q_ptr,
s_ptr,
M_i,
K_i,
x.stride(0),
x.stride(1),
q.stride(0),
q.stride(1),
s.stride(0),
s.stride(1));
}
auto err = hipGetLastError();
TORCH_CHECK(
err == hipSuccess,
"quant_bf16_to_mxfp4_kernel launch failed: ",
hipGetErrorString(err));
return {q, s};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("quantize", &quantize, "HIP MXFP4 quant kernel");
}
"""
@lru_cache(maxsize=1)
def _get_quant_module():
return load_inline(
name="mxfp4_quant_ext_v2", # bump this name if you want to force a rebuild
cpp_sources="",
cuda_sources=quant_src,
with_cuda=True,
verbose=False,
extra_cuda_cflags=["-std=c++20", "-O3"],
no_implicit_headers=True,
)
def _quant_mxfp4_custom(x, shuffle=True):
from aiter import dtypes
q_u8, s_u8 = _get_quant_module().quantize(x, shuffle)
return q_u8.view(dtypes.fp4x2), s_u8.view(dtypes.fp8_e8m0)
def custom_kernel(data: input_t) -> output_t:
import aiter
from aiter import dtypes
A, _, _, B_shuffle, B_scale_sh = data
A = A.contiguous()
A_q, A_scale_sh = _quant_mxfp4_custom(A, shuffle=True)
out_gemm = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out_gemm
scrolls · 365 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