submission 596924
Purple rain · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1084 lines, June 9 Researcher Reciprocity License v1.0.
submission_amd_mxfp4_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-596924?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:4e49dbfab68f4675da0fa9e6cde3642b97848823ae46d60f7545a9d24985e8a5
license declaredunknown
license concludedunknown
authorsPurple rain
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");shared-memory
__shared__ __align__(16) uint8_t smem_A[GEMM_BLOCK_M * FP4_BYTES_PER_KBLOCK];split-k
_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"tile-k = 32
constexpr int BLOCK_K = 32;tile-m = 64
constexpr int GEMM_BLOCK_M = 64;tile-n = 64
constexpr int GEMM_BLOCK_N = 64;vector-width = uint4
uint4 v4 = *reinterpret_cast<const uint4*>(row_ptr + vec_idx * 16);Kernel source
submission_amd_mxfp4_mm.py1084 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os
import threading
from typing import Dict, Tuple
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
_QUANT_BACKEND_ENV = "MXFP4_MM_QUANT_BACKEND" # auto | triton | hip
_EXEC_BACKEND_ENV = "MXFP4_MM_EXEC_BACKEND" # auto | aiter | hip
_B_LAYOUT_ENV = "MXFP4_MM_B_LAYOUT" # raw | shuffle | auto
_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"
_DEFAULT_QUANT_BACKEND = "hip"
_DEFAULT_EXEC_BACKEND = "auto"
_DEFAULT_B_LAYOUT = "shuffle"
_HIP_LOCK = threading.Lock()
_HIP_MODULE = None
_HIP_BUILD_ERROR = None
_B_SCALE_RAW_LOCK = threading.Lock()
_B_SCALE_RAW_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_B_SCALE_RAW_CACHE_MAX = 8
_WORKSPACE_LOCK = threading.Lock()
_WORKSPACE_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_WORKSPACE_CACHE_MAX = 16
_B_SHUFFLE_INNER_LOCK = threading.Lock()
_B_SHUFFLE_INNER_CACHE: Dict[Tuple[int, ...], int] = {}
_STATIC_SPLITK_LOG2: Dict[Tuple[int, int, int], int] = {
(4, 2880, 512): 2,
(16, 2112, 7168): 3,
(32, 4096, 512): 2,
(32, 2880, 512): 2,
(64, 7168, 2048): 1,
(256, 3072, 1536): 0,
}
CPP_WRAPPER = r"""
#include <cstdint>
#include <vector>
std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x);
torch::Tensor hip_gemm_mxfp4(
torch::Tensor a_fp4_u8,
torch::Tensor b_u8,
torch::Tensor a_scale_u8,
torch::Tensor b_scale_u8,
int64_t layout_mode,
int64_t log2_k_split,
torch::Tensor workspace,
int64_t workspace_stride,
int64_t b_shuffle_inner_mode);
"""
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <vector>
namespace {
constexpr int BLOCK_K = 32;
constexpr int PAD_M_KERNEL = 64;
constexpr int PAD_M_SCALE = 256;
constexpr int PAD_SCALE_N = 8;
constexpr int FP4_BYTES_PER_KBLOCK = BLOCK_K / 2;
constexpr int GEMM_BLOCK_M = 64;
constexpr int GEMM_BLOCK_N = 64;
constexpr int GEMM_THREADS = 256;
constexpr int WAVE_SIZE = 64;
constexpr int WAVE_TILE_M = 32;
constexpr int WAVE_TILE_N = 32;
constexpr int MFMA_TILE_M = 16;
constexpr int MFMA_TILE_N = 16;
__device__ __constant__ uint16_t FP4_TO_BF16_LUT[16] = {
0x0000, 0x3f00, 0x3f80, 0x3fc0, 0x4000, 0x4040, 0x4080, 0x40c0,
0x8000, 0xbf00, 0xbf80, 0xbfc0, 0xc000, 0xc040, 0xc080, 0xc0c0
};
__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
return __uint_as_float(static_cast<uint32_t>(x) << 16);
}
__device__ __forceinline__ uint8_t quantize_e2m1(float x_scaled) {
constexpr uint32_t T0 = 0x3E800000u; // 0.25
constexpr uint32_t T1 = 0x3F400000u; // 0.75
constexpr uint32_t T2 = 0x3FA00000u; // 1.25
constexpr uint32_t T3 = 0x3FE00000u; // 1.75
constexpr uint32_t T4 = 0x40200000u; // 2.5
constexpr uint32_t T5 = 0x40600000u; // 3.5
constexpr uint32_t T6 = 0x40A00000u; // 5.0
uint32_t bits = __float_as_uint(x_scaled);
uint32_t sign = (bits >> 31) & 0x1u;
uint32_t abs_bits = bits & 0x7FFFFFFFu;
uint32_t mag = 0;
mag += static_cast<uint32_t>(abs_bits >= T0);
mag += static_cast<uint32_t>(abs_bits >= T1);
mag += static_cast<uint32_t>(abs_bits >= T2);
mag += static_cast<uint32_t>(abs_bits >= T3);
mag += static_cast<uint32_t>(abs_bits >= T4);
mag += static_cast<uint32_t>(abs_bits >= T5);
mag += static_cast<uint32_t>(abs_bits >= T6);
return static_cast<uint8_t>(mag | (sign << 3));
}
__device__ __forceinline__ int64_t shuffled_scale_offset(
int64_t row,
int64_t col,
int64_t scale_n_pad) {
int64_t bs_offs_0 = row / 32;
int64_t bs_offs_1 = row % 32;
int64_t bs_offs_2 = bs_offs_1 % 16;
bs_offs_1 = bs_offs_1 / 16;
int64_t bs_offs_3 = col / 8;
int64_t bs_offs_4 = col % 8;
int64_t bs_offs_5 = bs_offs_4 % 4;
bs_offs_4 = bs_offs_4 / 4;
return bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4 + bs_offs_5 * 64 +
bs_offs_3 * 256 + bs_offs_0 * 32 * scale_n_pad;
}
__device__ __forceinline__ int64_t get_b_scale_shuffled_offset(
int64_t n,
int64_t k_blk,
int64_t scale_k_pad) {
int64_t n_outer = n / 32;
int64_t n_inner = n % 32;
int64_t k_outer = k_blk / 8;
int64_t k_inner = k_blk % 8;
int64_t n_16 = n_inner % 16;
int64_t n_2 = n_inner / 16;
int64_t k_4 = k_inner % 4;
int64_t k_2 = k_inner / 4;
return n_2 + (k_2 * 2) + (n_16 * 4) + (k_4 * 64) + (k_outer * 256) +
(n_outer * 32 * scale_k_pad);
}
__device__ __forceinline__ int64_t get_b_fp4_shuffled_offset(
int64_t n,
int64_t k_fp4,
int64_t k_fp4_pad,
int64_t inner_mode) {
int64_t n_blk = n / 16;
int64_t k_blk = k_fp4 / 16;
int64_t n_in = n % 16;
int64_t k_in = k_fp4 % 16;
int64_t blk_stride = k_fp4_pad / 16;
int64_t blk_offset = (n_blk * blk_stride + k_blk) * 256;
int64_t inner_offset = (inner_mode == 0) ? (k_in * 16 + n_in) : (n_in * 16 + k_in);
return blk_offset + inner_offset;
}
__device__ __forceinline__ float e8m0_to_f32_fast(uint8_t e) {
if (e == 0) {
return __uint_as_float(0x00400000u);
}
if (e == 0xFF) {
return __uint_as_float(0x7F800001u);
}
return __uint_as_float(static_cast<uint32_t>(e) << 23);
}
__device__ __forceinline__ uint16_t fp4_to_bf16_bits(uint8_t v) {
return FP4_TO_BF16_LUT[v & 0x0Fu];
}
using floatx4 = __attribute__((__vector_size__(4 * sizeof(float)))) float;
using bit16x4 = __attribute__((__vector_size__(4 * sizeof(uint16_t)))) uint16_t;
using bit16x8 = __attribute__((__vector_size__(8 * sizeof(uint16_t)))) uint16_t;
struct B16x8 {
bit16x4 xy[2];
};
__device__ __forceinline__ B16x8 unpack_fp4x8_to_b16x8(uint32_t pack) {
B16x8 reg;
const uint8_t* bytes = reinterpret_cast<const uint8_t*>(&pack);
#pragma unroll
for (int i = 0; i < 8; ++i) {
uint8_t packed = bytes[i >> 1];
uint8_t nib = (i & 1) ? static_cast<uint8_t>(packed >> 4)
: static_cast<uint8_t>(packed & 0x0F);
uint16_t bits = fp4_to_bf16_bits(nib);
if (i < 4) {
reg.xy[0][i] = bits;
} else {
reg.xy[1][i - 4] = bits;
}
}
return reg;
}
__device__ __forceinline__ uint32_t load_packed_fp4_word(
const uint8_t* tile,
int tile_idx,
int word_idx) {
const uint32_t* words = reinterpret_cast<const uint32_t*>(
tile + tile_idx * FP4_BYTES_PER_KBLOCK);
return words[word_idx];
}
__device__ __forceinline__ void accum_scaled_output_fragment(
floatx4& dst,
const floatx4& src,
const floatx4& a_scale,
float b_scale) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
dst[i] += src[i] * a_scale[i] * b_scale;
}
}
__device__ __forceinline__ floatx4 decode_scale4(const uint8_t* scale_ptr) {
floatx4 out;
#pragma unroll
for (int i = 0; i < 4; ++i) {
out[i] = e8m0_to_f32_fast(scale_ptr[i]);
}
return out;
}
__device__ __forceinline__ void store_accum_tile(
float* workspace,
int64_t workspace_stride,
int64_t split_idx,
int64_t m,
int64_t n,
int64_t row_base,
int64_t col,
int row_group,
const floatx4& acc) {
if (col >= n) {
return;
}
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = row_base + row_group * 4 + i;
if (row < m) {
int64_t out_off = split_idx * workspace_stride + row * n + col;
workspace[out_off] = acc[i];
}
}
}
__device__ __forceinline__ floatx4 gcn_mfma16x16x32_bf16(
const B16x8& a,
const B16x8& b,
const floatx4& c) {
#if defined(__gfx950__)
bit16x8 ta = __builtin_shufflevector(a.xy[0], a.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
bit16x8 tb = __builtin_shufflevector(b.xy[0], b.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
return __builtin_amdgcn_mfma_f32_16x16x32_bf16(ta, tb, c, 0, 0, 0);
#else
return c;
#endif
}
__global__ void quant_mxfp4_kernel(
const __hip_bfloat16* x,
uint8_t* out_fp4,
uint8_t* out_scale,
int64_t m,
int64_t m_pad,
int64_t k,
int64_t k_blocks_valid,
int64_t k_blocks_pad) {
int64_t linear = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t total = m_pad * k_blocks_pad;
if (linear >= total) {
return;
}
int64_t row = linear / k_blocks_pad;
int64_t kb = linear % k_blocks_pad;
int64_t scale_off = shuffled_scale_offset(row, kb, k_blocks_pad);
if (row >= m || kb >= k_blocks_valid) {
out_scale[scale_off] = 127;
if (row < m_pad && kb < k_blocks_valid) {
int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);
#pragma unroll
for (int i = 0; i < BLOCK_K / 2; ++i) {
out_fp4[out_base + i] = 0;
}
}
return;
}
int64_t in_base = row * k + kb * BLOCK_K;
float vals[BLOCK_K];
float amax = 0.0f;
const uint8_t* x_bytes = reinterpret_cast<const uint8_t*>(x);
const uint8_t* row_ptr = x_bytes + in_base * static_cast<int64_t>(sizeof(__hip_bfloat16));
#pragma unroll
for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
uint4 v4 = *reinterpret_cast<const uint4*>(row_ptr + vec_idx * 16);
uint32_t w0 = v4.x;
uint32_t w1 = v4.y;
uint32_t w2 = v4.z;
uint32_t w3 = v4.w;
uint16_t b0 = static_cast<uint16_t>(w0 & 0xFFFFu);
uint16_t b1 = static_cast<uint16_t>((w0 >> 16) & 0xFFFFu);
uint16_t b2 = static_cast<uint16_t>(w1 & 0xFFFFu);
uint16_t b3 = static_cast<uint16_t>((w1 >> 16) & 0xFFFFu);
uint16_t b4 = static_cast<uint16_t>(w2 & 0xFFFFu);
uint16_t b5 = static_cast<uint16_t>((w2 >> 16) & 0xFFFFu);
uint16_t b6 = static_cast<uint16_t>(w3 & 0xFFFFu);
uint16_t b7 = static_cast<uint16_t>((w3 >> 16) & 0xFFFFu);
int base = vec_idx * 8;
float f0 = bf16_to_f32(b0);
float f1 = bf16_to_f32(b1);
float f2 = bf16_to_f32(b2);
float f3 = bf16_to_f32(b3);
float f4 = bf16_to_f32(b4);
float f5 = bf16_to_f32(b5);
float f6 = bf16_to_f32(b6);
float f7 = bf16_to_f32(b7);
vals[base + 0] = f0;
vals[base + 1] = f1;
vals[base + 2] = f2;
vals[base + 3] = f3;
vals[base + 4] = f4;
vals[base + 5] = f5;
vals[base + 6] = f6;
vals[base + 7] = f7;
amax = fmaxf(amax, fabsf(f0));
amax = fmaxf(amax, fabsf(f1));
amax = fmaxf(amax, fabsf(f2));
amax = fmaxf(amax, fabsf(f3));
amax = fmaxf(amax, fabsf(f4));
amax = fmaxf(amax, fabsf(f5));
amax = fmaxf(amax, fabsf(f6));
amax = fmaxf(amax, fabsf(f7));
}
int exp_unbiased = 0;
float scale = 1.0f;
if (amax > 0.0f) {
float target = amax * (1.0f / 6.0f);
exp_unbiased = static_cast<int>(ceilf(log2f(target)));
scale = exp2f(static_cast<float>(exp_unbiased));
}
int exp_biased = exp_unbiased + 127;
exp_biased = exp_biased < 0 ? 0 : (exp_biased > 255 ? 255 : exp_biased);
out_scale[scale_off] = static_cast<uint8_t>(exp_biased);
float inv_scale = 1.0f / scale;
int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);
#pragma unroll
for (int i = 0; i < BLOCK_K / 2; ++i) {
uint8_t lo = quantize_e2m1(vals[2 * i] * inv_scale);
uint8_t hi = quantize_e2m1(vals[2 * i + 1] * inv_scale);
out_fp4[out_base + i] = static_cast<uint8_t>((hi << 4) | lo);
}
}
__global__ void gemm_mxfp4_splitk_kernel(
const uint8_t* a_fp4,
const uint8_t* b_u8,
const uint8_t* a_scale_u8,
const uint8_t* b_scale_u8,
float* workspace,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t layout_mode,
int64_t b_shuffle_inner_mode,
int64_t split_k) {
int tid = static_cast<int>(threadIdx.x);
int wave_id = tid / WAVE_SIZE;
int lane = tid & (WAVE_SIZE - 1);
int lane16 = lane & 15;
int row_group = lane >> 4; // 0..3
int wave_row = wave_id >> 1;
int wave_col = wave_id & 1;
int quad_row_base = wave_row * WAVE_TILE_M;
int quad_col_base = wave_col * WAVE_TILE_N;
int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
int local_col0 = quad_col_base + lane16;
int local_col1 = local_col0 + MFMA_TILE_N;
int64_t row_base0 = tile_m + quad_row_base;
int64_t row_base1 = row_base0 + MFMA_TILE_M;
int64_t col0 = tile_n + local_col0;
int64_t col1 = tile_n + local_col1;
floatx4 c_acc00 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc01 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc10 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc11 = {0.0f, 0.0f, 0.0f, 0.0f};
__shared__ __align__(16) uint8_t smem_A[GEMM_BLOCK_M * FP4_BYTES_PER_KBLOCK];
__shared__ __align__(16) uint8_t smem_B[GEMM_BLOCK_N * FP4_BYTES_PER_KBLOCK];
__shared__ uint8_t smem_A_scale[GEMM_BLOCK_M];
__shared__ uint8_t smem_B_scale[GEMM_BLOCK_N];
auto* smem_a_vec = reinterpret_cast<uint4*>(smem_A);
auto* smem_b_vec = reinterpret_cast<uint4*>(smem_B);
int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
int64_t kb_end = kb_start + kb_per_split;
if (kb_end > k_blocks_valid) {
kb_end = k_blocks_valid;
}
for (int64_t kb = kb_start; kb < kb_end; ++kb) {
if (tid < GEMM_BLOCK_M) {
int local_row = tid;
int64_t a_row = tile_m + local_row;
uint4 a_vec = make_uint4(0u, 0u, 0u, 0u);
uint8_t a_scale = 127;
if (a_row < m) {
a_vec = *reinterpret_cast<const uint4*>(
a_fp4 + a_row * k2 + kb * FP4_BYTES_PER_KBLOCK);
a_scale = a_scale_u8[a_row * k_blocks_pad + kb];
}
smem_a_vec[local_row] = a_vec;
smem_A_scale[local_row] = a_scale;
}
if (tid >= GEMM_BLOCK_M && tid < GEMM_BLOCK_M + GEMM_BLOCK_N) {
int local_n = tid - GEMM_BLOCK_M;
int64_t b_row = tile_n + local_n;
uint4 b_vec = make_uint4(0u, 0u, 0u, 0u);
uint8_t b_scale = 127;
if (b_row < n) {
if (layout_mode == 0) {
b_vec = *reinterpret_cast<const uint4*>(
b_u8 + b_row * k2 + kb * FP4_BYTES_PER_KBLOCK);
b_scale = b_scale_u8[b_row * k_blocks_pad + kb];
} else {
if (b_shuffle_inner_mode == 1) {
int64_t boff = get_b_fp4_shuffled_offset(
b_row, kb * FP4_BYTES_PER_KBLOCK, k2, b_shuffle_inner_mode);
b_vec = *reinterpret_cast<const uint4*>(b_u8 + boff);
} else {
uint8_t* dst = reinterpret_cast<uint8_t*>(&b_vec);
#pragma unroll
for (int bi = 0; bi < FP4_BYTES_PER_KBLOCK; ++bi) {
int64_t k_fp4 = kb * FP4_BYTES_PER_KBLOCK + bi;
int64_t boff = get_b_fp4_shuffled_offset(
b_row, k_fp4, k2, b_shuffle_inner_mode);
dst[bi] = b_u8[boff];
}
}
int64_t soff = get_b_scale_shuffled_offset(b_row, kb, k_blocks_pad);
b_scale = b_scale_u8[soff];
}
}
smem_b_vec[local_n] = b_vec;
smem_B_scale[local_n] = b_scale;
}
__syncthreads();
uint32_t a_pack0 = load_packed_fp4_word(
smem_A, quad_row_base + lane16, row_group);
uint32_t a_pack1 = load_packed_fp4_word(
smem_A, quad_row_base + MFMA_TILE_M + lane16, row_group);
uint32_t b_pack0 = load_packed_fp4_word(
smem_B, quad_col_base + lane16, row_group);
uint32_t b_pack1 = load_packed_fp4_word(
smem_B, quad_col_base + MFMA_TILE_N + lane16, row_group);
B16x8 a_reg0 = unpack_fp4x8_to_b16x8(a_pack0);
B16x8 a_reg1 = unpack_fp4x8_to_b16x8(a_pack1);
B16x8 b_reg0 = unpack_fp4x8_to_b16x8(b_pack0);
B16x8 b_reg1 = unpack_fp4x8_to_b16x8(b_pack1);
floatx4 t_acc00 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 t_acc01 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 t_acc10 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 t_acc11 = {0.0f, 0.0f, 0.0f, 0.0f};
t_acc00 = gcn_mfma16x16x32_bf16(a_reg0, b_reg0, t_acc00);
t_acc01 = gcn_mfma16x16x32_bf16(a_reg0, b_reg1, t_acc01);
t_acc10 = gcn_mfma16x16x32_bf16(a_reg1, b_reg0, t_acc10);
t_acc11 = gcn_mfma16x16x32_bf16(a_reg1, b_reg1, t_acc11);
int row_scale_base0 = quad_row_base + row_group * 4;
int row_scale_base1 = row_scale_base0 + MFMA_TILE_M;
const uint8_t* a_scale_ptr0 = smem_A_scale + row_scale_base0;
const uint8_t* a_scale_ptr1 = smem_A_scale + row_scale_base1;
floatx4 a_scale0 = decode_scale4(a_scale_ptr0);
floatx4 a_scale1 = decode_scale4(a_scale_ptr1);
float b_scale0 = e8m0_to_f32_fast(smem_B_scale[local_col0]);
float b_scale1 = e8m0_to_f32_fast(smem_B_scale[local_col1]);
accum_scaled_output_fragment(c_acc00, t_acc00, a_scale0, b_scale0);
accum_scaled_output_fragment(c_acc01, t_acc01, a_scale0, b_scale1);
accum_scaled_output_fragment(c_acc10, t_acc10, a_scale1, b_scale0);
accum_scaled_output_fragment(c_acc11, t_acc11, a_scale1, b_scale1);
__syncthreads();
}
store_accum_tile(
workspace,
workspace_stride,
static_cast<int64_t>(blockIdx.z),
m,
n,
row_base0,
col0,
row_group,
c_acc00);
store_accum_tile(
workspace,
workspace_stride,
static_cast<int64_t>(blockIdx.z),
m,
n,
row_base0,
col1,
row_group,
c_acc01);
store_accum_tile(
workspace,
workspace_stride,
static_cast<int64_t>(blockIdx.z),
m,
n,
row_base1,
col0,
row_group,
c_acc10);
store_accum_tile(
workspace,
workspace_stride,
static_cast<int64_t>(blockIdx.z),
m,
n,
row_base1,
col1,
row_group,
c_acc11);
}
__global__ void reduce_splitk_kernel(
const float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t split_k) {
int64_t row = static_cast<int64_t>(blockIdx.y) * blockDim.y + threadIdx.y;
int64_t col = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
if (row >= m || col >= n) {
return;
}
float sum = 0.0f;
int64_t base = row * n + col;
for (int64_t s = 0; s < split_k; ++s) {
sum += workspace[s * workspace_stride + base];
}
out[base] = static_cast<__hip_bfloat16>(sum);
}
inline void check_hip_error(const char* where) {
hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}
} // namespace
std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x) {
TORCH_CHECK(x.is_cuda(), "x must be CUDA tensor");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bfloat16");
TORCH_CHECK(x.dim() == 2, "x must be 2D [M, K]");
auto x_contig = x.contiguous();
int64_t m = x_contig.size(0);
int64_t k = x_contig.size(1);
TORCH_CHECK(k % 64 == 0, "k must be divisible by 64");
int64_t k_blocks_valid = k / BLOCK_K;
int64_t k_blocks_pad = ((k_blocks_valid + (PAD_SCALE_N - 1)) / PAD_SCALE_N) * PAD_SCALE_N;
int64_t m_pad_kernel = ((m + (PAD_M_KERNEL - 1)) / PAD_M_KERNEL) * PAD_M_KERNEL;
int64_t m_pad_scale = ((m + (PAD_M_SCALE - 1)) / PAD_M_SCALE) * PAD_M_SCALE;
auto u8_opts = x_contig.options().dtype(torch::kUInt8);
auto out_fp4 = torch::empty({m_pad_kernel, k / 2}, u8_opts);
auto out_scale = torch::full({m_pad_scale, k_blocks_pad}, 127, u8_opts);
int64_t total = m_pad_kernel * k_blocks_pad;
int threads = 256;
int blocks = static_cast<int>((total + threads - 1) / threads);
if (blocks > 0) {
hipLaunchKernelGGL(
quant_mxfp4_kernel,
dim3(blocks),
dim3(threads),
0,
0,
reinterpret_cast<const __hip_bfloat16*>(x_contig.data_ptr()),
reinterpret_cast<uint8_t*>(out_fp4.data_ptr()),
reinterpret_cast<uint8_t*>(out_scale.data_ptr()),
m,
m_pad_kernel,
k,
k_blocks_valid,
k_blocks_pad);
check_hip_error("quant_mxfp4_kernel");
}
return {out_fp4, out_scale};
}
torch::Tensor hip_gemm_mxfp4(
torch::Tensor a_fp4_u8,
torch::Tensor b_u8,
torch::Tensor a_scale_u8,
torch::Tensor b_scale_u8,
int64_t layout_mode,
int64_t log2_k_split,
torch::Tensor workspace,
int64_t workspace_stride,
int64_t b_shuffle_inner_mode) {
TORCH_CHECK(a_fp4_u8.is_cuda(), "a_fp4_u8 must be CUDA");
TORCH_CHECK(b_u8.is_cuda(), "b_u8 must be CUDA");
TORCH_CHECK(a_scale_u8.is_cuda(), "a_scale_u8 must be CUDA");
TORCH_CHECK(b_scale_u8.is_cuda(), "b_scale_u8 must be CUDA");
TORCH_CHECK(a_fp4_u8.scalar_type() == torch::kUInt8, "a_fp4_u8 must be uint8");
TORCH_CHECK(b_u8.scalar_type() == torch::kUInt8, "b_u8 must be uint8");
TORCH_CHECK(a_scale_u8.scalar_type() == torch::kUInt8, "a_scale_u8 must be uint8");
TORCH_CHECK(b_scale_u8.scalar_type() == torch::kUInt8, "b_scale_u8 must be uint8");
TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");
TORCH_CHECK(a_scale_u8.dim() == 2 && b_scale_u8.dim() == 2, "A/B scale must be 2D");
TORCH_CHECK(layout_mode == 0 || layout_mode == 1, "layout_mode must be 0 or 1");
auto a_fp4 = a_fp4_u8.contiguous();
auto b_fp4 = b_u8.contiguous();
auto a_scale = a_scale_u8.contiguous();
auto b_scale = b_scale_u8.contiguous();
int64_t m = a_fp4.size(0);
int64_t n = b_fp4.size(0);
int64_t k2 = a_fp4.size(1);
TORCH_CHECK(b_fp4.size(1) == k2, "A/B K/2 mismatch");
int64_t k = k2 * 2;
TORCH_CHECK(k % BLOCK_K == 0, "K must be divisible by 32");
int64_t k_blocks_valid = k / BLOCK_K;
int64_t k_blocks_pad = a_scale.size(1);
TORCH_CHECK(b_scale.size(1) == k_blocks_pad, "A/B scale padded K-block mismatch");
TORCH_CHECK(a_scale.size(0) >= m, "a_scale rows must cover m");
TORCH_CHECK(b_scale.size(0) >= n, "b_scale rows must cover n");
int64_t split_k = 1;
if (log2_k_split > 0) {
split_k = static_cast<int64_t>(1) << log2_k_split;
}
if (split_k < 1) {
split_k = 1;
}
int64_t mn = m * n;
if (workspace_stride <= 0) {
workspace_stride = mn;
}
TORCH_CHECK(workspace_stride >= mn, "workspace_stride must be >= m*n");
auto ws = workspace;
auto ws_opts = a_fp4.options().dtype(torch::kFloat);
int64_t need = split_k * workspace_stride;
if (!ws.defined() || !ws.is_cuda() || ws.scalar_type() != torch::kFloat || ws.numel() < need) {
ws = torch::empty({split_k, workspace_stride}, ws_opts);
} else {
ws = ws.contiguous().view({split_k, workspace_stride});
}
auto out = torch::empty({m, n}, a_fp4.options().dtype(torch::kBFloat16));
dim3 block(GEMM_THREADS);
dim3 grid(
static_cast<unsigned int>((n + GEMM_BLOCK_N - 1) / GEMM_BLOCK_N),
static_cast<unsigned int>((m + GEMM_BLOCK_M - 1) / GEMM_BLOCK_M),
static_cast<unsigned int>(split_k));
hipLaunchKernelGGL(
gemm_mxfp4_splitk_kernel,
grid,
block,
0,
0,
reinterpret_cast<const uint8_t*>(a_fp4.data_ptr()),
reinterpret_cast<const uint8_t*>(b_fp4.data_ptr()),
reinterpret_cast<const uint8_t*>(a_scale.data_ptr()),
reinterpret_cast<const uint8_t*>(b_scale.data_ptr()),
reinterpret_cast<float*>(ws.data_ptr()),
workspace_stride,
m,
n,
k2,
k_blocks_valid,
k_blocks_pad,
layout_mode,
b_shuffle_inner_mode,
split_k);
check_hip_error("gemm_mxfp4_splitk_kernel");
dim3 rblock(16, 16);
dim3 rgrid(
static_cast<unsigned int>((n + 15) / 16),
static_cast<unsigned int>((m + 15) / 16));
hipLaunchKernelGGL(
reduce_splitk_kernel,
rgrid,
rblock,
0,
0,
reinterpret_cast<const float*>(ws.data_ptr()),
reinterpret_cast<__hip_bfloat16*>(out.data_ptr()),
workspace_stride,
m,
n,
split_k);
check_hip_error("reduce_splitk_kernel");
return out;
}
"""
def _sanitize_quant_backend(mode: str) -> str:
mode = (mode or _DEFAULT_QUANT_BACKEND).strip().lower()
if mode in {"auto", "triton", "hip"}:
return mode
return _DEFAULT_QUANT_BACKEND
def _get_quant_backend() -> str:
return _sanitize_quant_backend(os.getenv(_QUANT_BACKEND_ENV, _DEFAULT_QUANT_BACKEND))
def _sanitize_exec_backend(mode: str) -> str:
mode = (mode or _DEFAULT_EXEC_BACKEND).strip().lower()
if mode in {"auto", "aiter", "hip"}:
return mode
return _DEFAULT_EXEC_BACKEND
def _get_exec_backend() -> str:
mode = _sanitize_exec_backend(os.getenv(_EXEC_BACKEND_ENV, _DEFAULT_EXEC_BACKEND))
if mode == "auto":
return "aiter"
return mode
def _sanitize_b_layout(mode: str) -> str:
mode = (mode or _DEFAULT_B_LAYOUT).strip().lower()
if mode in {"raw", "shuffle", "auto"}:
return mode
return _DEFAULT_B_LAYOUT
def _get_b_layout() -> str:
mode = _sanitize_b_layout(os.getenv(_B_LAYOUT_ENV, _DEFAULT_B_LAYOUT))
if mode == "auto":
return "shuffle"
return mode
def _get_splitk_override() -> int | None:
raw = os.getenv(_SPLITK_ENV)
if raw is None:
return None
try:
return max(0, int(raw))
except ValueError:
return None
def _e8m0_unshuffle(scale_sh: torch.Tensor) -> torch.Tensor:
if scale_sh.ndim != 2:
raise RuntimeError(f"scale_sh must be 2D, got {tuple(scale_sh.shape)}")
sm, sn = scale_sh.shape
if sm % 32 != 0 or sn % 8 != 0:
raise RuntimeError(f"scale_sh shape must be divisible by (32,8), got {(sm, sn)}")
s = scale_sh.view(torch.uint8)
s = s.view(sm // 32, sn // 8, 4, 16, 2, 2)
s = s.permute(0, 5, 3, 1, 4, 2).contiguous()
s = s.view(sm, sn)
return s.view(scale_sh.dtype)
def _get_b_scale_raw_cached(b_scale_sh: torch.Tensor) -> torch.Tensor:
dev = int(b_scale_sh.device.index) if b_scale_sh.device.index is not None else -1
key = (
int(b_scale_sh.data_ptr()),
int(b_scale_sh.shape[0]),
int(b_scale_sh.shape[1]),
dev,
)
cached = _B_SCALE_RAW_CACHE.get(key)
if cached is not None:
return cached
with _B_SCALE_RAW_LOCK:
cached = _B_SCALE_RAW_CACHE.get(key)
if cached is not None:
return cached
raw = _e8m0_unshuffle(b_scale_sh).contiguous()
if len(_B_SCALE_RAW_CACHE) >= _B_SCALE_RAW_CACHE_MAX:
_B_SCALE_RAW_CACHE.pop(next(iter(_B_SCALE_RAW_CACHE)))
_B_SCALE_RAW_CACHE[key] = raw
return raw
def _get_workspace(device: torch.device, m: int, n: int, split_k: int) -> torch.Tensor:
dev = int(device.index) if device.index is not None else -1
key = (dev, int(m), int(n), int(split_k))
cached = _WORKSPACE_CACHE.get(key)
need = split_k * m * n
if cached is not None and cached.numel() >= need:
return cached
with _WORKSPACE_LOCK:
cached = _WORKSPACE_CACHE.get(key)
if cached is not None and cached.numel() >= need:
return cached
ws = torch.empty((split_k, m * n), dtype=torch.float32, device=device)
if len(_WORKSPACE_CACHE) >= _WORKSPACE_CACHE_MAX:
_WORKSPACE_CACHE.pop(next(iter(_WORKSPACE_CACHE)))
_WORKSPACE_CACHE[key] = ws
return ws
def _pick_splitk_log2(m: int, n: int, k: int) -> int:
override = _get_splitk_override()
if override is not None:
return override
key = (int(m), int(n), int(k))
if key in _STATIC_SPLITK_LOG2:
return _STATIC_SPLITK_LOG2[key]
if m <= 16 and k >= 4096:
return 3
if m <= 32:
return 2
if m <= 64 and k >= 2048:
return 1
return 0
def _get_hip_module():
global _HIP_MODULE, _HIP_BUILD_ERROR
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
with _HIP_LOCK:
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
try:
os.environ.setdefault("CXX", "clang++")
_HIP_MODULE = load_inline(
name="mxfp4_mm_inline_quant_gemm_v4",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[HIP_SRC],
functions=["hip_quant_mxfp4", "hip_gemm_mxfp4"],
verbose=False,
extra_cuda_cflags=["-O3", "-std=c++20"],
)
except Exception as e: # pragma: no cover - runtime dependent
_HIP_BUILD_ERROR = e
raise RuntimeError(f"HIP inline build failed: {e}") from e
return _HIP_MODULE
def _quant_triton_mxfp4(x: torch.Tensor, shuffle: bool = True):
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
if shuffle:
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
def _run_aiter_pipeline(
a_bf16: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
import aiter
from aiter import dtypes
a_q_sh, a_scale_sh = _quant_triton_mxfp4(a_bf16, shuffle=True)
return aiter.gemm_a4w4(
a_q_sh,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
def _detect_b_shuffle_inner_mode(
module,
a_bf16: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> int:
dev = int(a_bf16.device.index) if a_bf16.device.index is not None else -1
key = (int(b_shuffle.shape[0]), int(b_shuffle.shape[1]), dev)
cached = _B_SHUFFLE_INNER_CACHE.get(key)
if cached is not None:
return cached
with _B_SHUFFLE_INNER_LOCK:
cached = _B_SHUFFLE_INNER_CACHE.get(key)
if cached is not None:
return cached
best_mode = 0
try:
import aiter
from aiter import dtypes
m_probe = min(8, int(a_bf16.shape[0]))
a_probe = a_bf16[:m_probe, :].contiguous()
a_q_sh, a_scale_sh = _quant_triton_mxfp4(a_probe, shuffle=True)
ref = aiter.gemm_a4w4(
a_q_sh,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
a_q_raw, a_scale_raw = _quant_triton_mxfp4(a_probe, shuffle=False)
a_fp4_u8 = a_q_raw.view(torch.uint8).contiguous()
a_scale_u8 = a_scale_raw.view(torch.uint8).contiguous()
b_u8 = b_shuffle.view(torch.uint8).contiguous()
b_scale_u8 = b_scale_sh.view(torch.uint8).contiguous()
ws = _get_workspace(a_probe.device, m_probe, int(b_shuffle.shape[0]), 1)
errs = []
for mode in (0, 1):
out = module.hip_gemm_mxfp4(
a_fp4_u8,
b_u8,
a_scale_u8,
b_scale_u8,
1,
0,
ws,
int(m_probe * int(b_shuffle.shape[0])),
mode,
)
err = (out.float() - ref.float()).abs().max().item()
errs.append(err)
best_mode = 0 if errs[0] <= errs[1] else 1
except Exception:
best_mode = 0
_B_SHUFFLE_INNER_CACHE[key] = best_mode
return best_mode
def _run_hip_full_pipeline(
a_bf16: torch.Tensor,
b_q: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
module = _get_hip_module()
quant_backend = _get_quant_backend()
b_layout = _get_b_layout()
if quant_backend == "hip":
a_fp4_u8, a_scale_sh_u8 = module.hip_quant_mxfp4(a_bf16)
a_fp4_u8 = a_fp4_u8[: a_bf16.shape[0], :].contiguous()
a_scale_u8 = _e8m0_unshuffle(a_scale_sh_u8.view(torch.uint8)).contiguous()
else:
a_q, a_scale_raw = _quant_triton_mxfp4(a_bf16, shuffle=False)
a_fp4_u8 = a_q.view(torch.uint8).contiguous()
a_scale_u8 = a_scale_raw.view(torch.uint8).contiguous()
if b_layout == "shuffle":
b_u8 = b_shuffle.view(torch.uint8).contiguous()
b_scale_u8 = b_scale_sh.view(torch.uint8).contiguous()
layout_mode = 1
b_inner_mode = _detect_b_shuffle_inner_mode(module, a_bf16, b_shuffle, b_scale_sh)
else:
b_u8 = b_q.view(torch.uint8).contiguous()
b_scale_u8 = _get_b_scale_raw_cached(b_scale_sh).view(torch.uint8).contiguous()
layout_mode = 0
b_inner_mode = 0
m = int(a_bf16.shape[0])
k = int(a_bf16.shape[1])
n = int(b_q.shape[0])
log2_k_split = _pick_splitk_log2(m, n, k)
split_k = 1 << log2_k_split
ws = _get_workspace(a_bf16.device, m, n, split_k)
return module.hip_gemm_mxfp4(
a_fp4_u8,
b_u8,
a_scale_u8,
b_scale_u8,
int(layout_mode),
int(log2_k_split),
ws,
int(m * n),
int(b_inner_mode),
)
def custom_kernel(data: input_t) -> output_t:
"""
Default path uses aiter's preshuffled GEMM for leaderboard throughput.
Set MXFP4_MM_EXEC_BACKEND=hip to force the custom HIP pipeline.
"""
A, B, B_q, B_shuffle, B_scale_sh = data
del B
A = A.contiguous()
B_q = B_q.contiguous()
B_shuffle = B_shuffle.contiguous()
B_scale_sh = B_scale_sh.contiguous()
if _get_exec_backend() == "aiter":
return _run_aiter_pipeline(
a_bf16=A,
b_shuffle=B_shuffle,
b_scale_sh=B_scale_sh,
)
return _run_hip_full_pipeline(
a_bf16=A,
b_q=B_q,
b_shuffle=B_shuffle,
b_scale_sh=B_scale_sh,
)
scrolls · 1084 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