submission 712263
sharkconi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 703 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-712263?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:ee9664847b421b083a9ac0a931bef52d9078a074daec64bfc48242688eeb3c8f
license declaredunknown
license concludedunknown
authorssharkconi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
__global__ void fused_mxfp4_gemm_splitk_kernel(vector-width = uint4
const uint4* row_ptr128 = reinterpret_cast<const uint4*>(row_ptr);Kernel source
submission.py703 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Combined module: old mxfp4_quant_fused HIP kernel + fused MFMA GEMM kernel,
both in a single load_inline module for server compatibility.
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <torch/extension.h>
#include <cstdint>
// ======================================================================
// KERNEL 1: Fused MXFP4 quantization + E8M0 scale shuffle kernel.
// Each thread processes one 32-element block of the input.
// Produces packed FP4x2 output and directly writes the shuffled E8M0 scale.
// ======================================================================
__global__ void mxfp4_quant_shuffle_kernel(
const __hip_bfloat16* __restrict__ x_in,
uint8_t* __restrict__ x_fp4_out,
uint8_t* __restrict__ scale_out,
const int M,
const int K,
const int scaleN_valid,
const int scaleN_pad,
const int padM256
) {
const int row = blockIdx.x * blockDim.y + threadIdx.y;
const int scale_col = blockIdx.y * blockDim.x + threadIdx.x;
if (row >= padM256 || scale_col >= scaleN_pad) return;
// Compute shuffled scale index for ALL positions (valid + padding)
int i0 = row / 32;
int rem32 = row % 32;
int i1 = rem32 / 16;
int i2 = rem32 % 16;
int i3 = scale_col / 8;
int sc_rem8 = scale_col % 8;
int i4 = sc_rem8 / 4;
int i5 = sc_rem8 % 4;
int sn8 = scaleN_pad / 8;
int out_idx = i0 * (sn8 * 256)
+ i3 * 256
+ i5 * 64
+ i2 * 4
+ i4 * 2
+ i1;
// Padding position: write neutral E8M0 scale (127) and skip fp4
if (row >= M || scale_col >= scaleN_valid) {
scale_out[out_idx] = 127;
return;
}
const int k_start = scale_col * 32;
// Load 32 bf16 values and compute amax simultaneously
float vals[32];
float amax_val = 0.0f;
const __hip_bfloat16* row_ptr = x_in + (int64_t)row * K + k_start;
// Load 32 bf16 values using 128-bit vector loads (8 bf16 per load = 4 loads)
const uint4* row_ptr128 = reinterpret_cast<const uint4*>(row_ptr);
#pragma unroll
for (int j = 0; j < 4; j++) {
uint4 data = row_ptr128[j];
const __hip_bfloat16* bvals = reinterpret_cast<const __hip_bfloat16*>(&data);
#pragma unroll
for (int i = 0; i < 8; i++) {
float v = __bfloat162float(bvals[i]);
vals[j * 8 + i] = v;
amax_val = fmaxf(amax_val, fabsf(v));
}
}
// Compute E8M0 scale using bitwise operations (no transcendentals)
uint32_t amax_bits = __float_as_uint(amax_val);
amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
uint32_t exponent = (amax_bits >> 23) & 0xFFu;
uint32_t e8m0_val = (exponent >= 2u) ? (exponent - 2u) : 0u;
uint8_t bs_e8m0 = (uint8_t)e8m0_val;
float inverted_scale;
if (e8m0_val == 0u) {
inverted_scale = __uint_as_float(0x7F800000u); // +inf
bs_e8m0 = 0;
} else {
inverted_scale = __uint_as_float(e8m0_val << 23);
}
// Quantize 32 values using hardware fp4 conversion
uint32_t packed32[4];
#pragma unroll
for (int g = 0; g < 4; g++) {
uint32_t w = 0;
int base = g * 8;
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[base+0], vals[base+1], inverted_scale, 0);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[base+2], vals[base+3], inverted_scale, 1);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[base+4], vals[base+5], inverted_scale, 2);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[base+6], vals[base+7], inverted_scale, 3);
packed32[g] = w;
}
// Write packed fp4 output (16 bytes = 32 fp4 values) as a single 128-bit store
uint4* out_ptr128 = reinterpret_cast<uint4*>(
x_fp4_out + (int64_t)row * (K / 2) + k_start / 2);
uint4 out_data;
out_data.x = packed32[0];
out_data.y = packed32[1];
out_data.z = packed32[2];
out_data.w = packed32[3];
out_ptr128[0] = out_data;
// Write scale in shuffled order
scale_out[out_idx] = bs_e8m0;
}
// C++ wrapper for quant kernel
std::vector<torch::Tensor> mxfp4_quant_fused(torch::Tensor x_in) {
const int M = x_in.size(0);
const int K = x_in.size(1);
const int scaleN_valid = (K + 31) / 32;
const int scaleN_pad = ((scaleN_valid + 7) / 8) * 8;
const int padM256 = ((M + 31) / 32) * 32;
auto x_fp4 = torch::empty({M, K / 2},
torch::TensorOptions().dtype(torch::kUInt8).device(x_in.device()));
auto scale = torch::empty({padM256 * scaleN_pad},
torch::TensorOptions().dtype(torch::kUInt8).device(x_in.device()));
int block_x = 32;
if (scaleN_valid <= 16) block_x = 16;
if (scaleN_valid <= 8) block_x = 8;
int block_y = 256 / block_x;
if (block_y > 32) block_y = 32;
if (block_y < 1) block_y = 1;
dim3 block(block_x, block_y);
dim3 grid(
(padM256 + block_y - 1) / block_y,
(scaleN_pad + block_x - 1) / block_x
);
mxfp4_quant_shuffle_kernel<<<grid, block>>>(
reinterpret_cast<const __hip_bfloat16*>(x_in.data_ptr<at::BFloat16>()),
x_fp4.data_ptr<uint8_t>(),
scale.data_ptr<uint8_t>(),
M, K, scaleN_valid, scaleN_pad, padM256
);
scale = scale.view({padM256, scaleN_pad});
return {x_fp4, scale};
}
// In-place quant: writes into pre-allocated output tensors (no allocation)
void mxfp4_quant_inplace(torch::Tensor x_in, torch::Tensor x_fp4_out, torch::Tensor scale_out) {
const int M = x_in.size(0);
const int K = x_in.size(1);
const int scaleN_valid = (K + 31) / 32;
const int scaleN_pad = ((scaleN_valid + 7) / 8) * 8;
const int padM256 = ((M + 31) / 32) * 32;
int block_x = 32;
if (scaleN_valid <= 16) block_x = 16;
if (scaleN_valid <= 8) block_x = 8;
int block_y = 256 / block_x;
if (block_y > 32) block_y = 32;
if (block_y < 1) block_y = 1;
dim3 block(block_x, block_y);
dim3 grid(
(padM256 + block_y - 1) / block_y,
(scaleN_pad + block_x - 1) / block_x
);
mxfp4_quant_shuffle_kernel<<<grid, block>>>(
reinterpret_cast<const __hip_bfloat16*>(x_in.data_ptr<at::BFloat16>()),
x_fp4_out.data_ptr<uint8_t>(),
scale_out.data_ptr<uint8_t>(),
M, K, scaleN_valid, scaleN_pad, padM256
);
}
// ======================================================================
// KERNEL 2: Fused MFMA GEMM kernel (16x16x128 variant for K<=512).
// bf16 A -> fp4 quant + GEMM with pre-quantized fp4 B in a single kernel.
// Each wavefront (64 threads) computes one 16x16 output tile.
// Uses __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4 on gfx950.
// K per MFMA = 128, so K=512 needs only 4 iterations (vs 8 with 32x32x64).
// ======================================================================
#if defined(__gfx950__)
typedef int __attribute__((ext_vector_type(8))) i32x8_t;
typedef float __attribute__((ext_vector_type(16))) fp32x16_t;
typedef float __attribute__((ext_vector_type(4))) fp32x4_t;
#endif
template<int CONST_K>
__launch_bounds__(64, 1)
__global__ void fused_mxfp4_gemm_kernel(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
__hip_bfloat16* __restrict__ C,
const int M, const int N, const int K,
const int scaleN_pad
) {
#if defined(__gfx950__)
const int lane = threadIdx.x; // 0-63
const int lane16 = lane & 15; // row within 16-row tile (A), col within 16-col tile (B)
const int group4 = lane >> 4; // K-quarter: 0,1,2,3
const int tile_m = blockIdx.x;
const int tile_n = blockIdx.y;
const int out_row_base = tile_m * 16;
const int out_col = tile_n * 16 + lane16;
const int a_row = out_row_base + lane16;
const int K_half = K / 2; // bytes per row of B_q
const int sn8 = scaleN_pad / 8; // for un-shuffle index computation
// Pre-compute B scale shuffle index components constant across K loop
// B is indexed by (out_col, scale_col). out_col is constant per lane.
const int b_i0 = out_col / 32;
const int b_rem32 = out_col & 31;
const int b_i1 = b_rem32 / 16;
const int b_i2 = b_rem32 & 15;
const int b_i0_stride = b_i0 * (sn8 * 256);
const int b_i2_i1_part = b_i2 * 4 + b_i1;
// Initialize accumulator — 4 fp32 values for 16x16x128 MFMA
fp32x4_t c_reg = {};
// Declare registers outside K loop
i32x8_t a_reg = {};
i32x8_t b_reg = {};
int scale_a_val = 0;
int scale_b_val = 127;
// K loop: fully unrolled for compile-time K (4 iters for K=512)
#pragma unroll
for (int k_iter = 0; k_iter < CONST_K; k_iter += 128) {
// Each group4 handles 32 fp4 elements (16 bytes) within the 128-element step
const int k_quarter_start = k_iter + group4 * 32;
// ============ Load and quantize A (branch-free, single-pass) ============
const int safe_row = (a_row < M) ? a_row : 0;
const __hip_bfloat16* a_ptr = A + (int64_t)safe_row * K + k_quarter_start;
// Load 32 bf16 values, compute amax, quantize to 16 bytes fp4
float vals[32];
float amax_val = 0.0f;
const uint4* a_ptr128 = reinterpret_cast<const uint4*>(a_ptr);
#pragma unroll
for (int j = 0; j < 4; j++) {
uint4 data = a_ptr128[j];
const __hip_bfloat16* bv = reinterpret_cast<const __hip_bfloat16*>(&data);
#pragma unroll
for (int i = 0; i < 8; i++) {
float v = __bfloat162float(bv[i]);
vals[j * 8 + i] = v;
amax_val = fmaxf(amax_val, fabsf(v));
}
}
// Compute E8M0 scale (branchless)
uint32_t amax_bits = __float_as_uint(amax_val);
amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
uint32_t exponent = (amax_bits >> 23) & 0xFFu;
uint32_t a_e8m0 = (exponent >= 2u) ? (exponent - 2u) : 0u;
float inv_scale = (a_e8m0 > 0u) ? __uint_as_float(a_e8m0 << 23) : __uint_as_float(0x7F800000u);
scale_a_val = (a_row < M) ? (int)a_e8m0 : 0;
// Quantize from vals[] in registers — produces 4 x uint32 = 16 bytes = 32 fp4
#pragma unroll
for (int g = 0; g < 4; g++) {
uint32_t w = 0;
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+0], vals[g*8+1], inv_scale, 0);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+2], vals[g*8+3], inv_scale, 1);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+4], vals[g*8+5], inv_scale, 2);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+6], vals[g*8+7], inv_scale, 3);
a_reg[g] = __builtin_bit_cast(int, w);
}
// Zero a_reg for out-of-bounds rows
if (a_row >= M) { a_reg = {}; scale_a_val = 0; }
// ============ Load B (vectorized 128-bit load) ============
b_reg = {};
scale_b_val = 127;
if (out_col < N) {
const uint4* b_src128 = reinterpret_cast<const uint4*>(
B_q + (int64_t)out_col * K_half + k_quarter_start / 2);
uint4 b_data = b_src128[0];
b_reg[0] = __builtin_bit_cast(int, b_data.x);
b_reg[1] = __builtin_bit_cast(int, b_data.y);
b_reg[2] = __builtin_bit_cast(int, b_data.z);
b_reg[3] = __builtin_bit_cast(int, b_data.w);
// Compute un-shuffle index to read from shuffled B_scale_sh
// scale_col = k_iter/32 + group4 (each group covers one 32-element scale block)
int scale_col = k_iter / 32 + group4;
int i3 = scale_col / 8;
int sc_rem8 = scale_col & 7;
int i4 = sc_rem8 / 4;
int i5 = sc_rem8 & 3;
int shuffled_idx = b_i0_stride + i3 * 256 + i5 * 64 + b_i2_i1_part + i4 * 2;
scale_b_val = (int)B_scale_sh[shuffled_idx];
}
// ============ Execute MFMA (16x16x128) ============
c_reg = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_reg, b_reg, c_reg,
4, 4, 0, scale_a_val, 0, scale_b_val
);
}
// ============ Store output ============
// 16x16x128 output mapping: 4 groups of 16 lanes, each group writes 4 rows
// Output row = tile_m*16 + group4*4 + i (i=0..3)
// Output col = tile_n*16 + lane16
if (out_col < N) {
#pragma unroll
for (int i = 0; i < 4; i++) {
const int64_t r = out_row_base + group4 * 4 + i;
C[r * N + out_col] = __float2bfloat16(c_reg[i]);
}
}
#endif // __gfx950__
}
// C++ wrapper for fused GEMM kernel (16x16 tiles)
void fused_mxfp4_gemm(
torch::Tensor A, // [M, K] bf16
torch::Tensor B_q, // [N, K/2] uint8 fp4x2
torch::Tensor B_scale_sh, // [padN256, scaleN_pad] uint8 E8M0 shuffled
torch::Tensor C, // [M_padded, N] bf16 output (pre-allocated)
int N_dim // actual N dimension
) {
const int M = A.size(0);
const int K = A.size(1);
const int N = N_dim;
const int K_scale = K / 32;
const int scaleN_pad = ((K_scale + 7) / 8) * 8;
const int M_padded = C.size(0);
dim3 grid(
(M_padded + 15) / 16,
(N + 15) / 16
);
dim3 block(64);
// Dispatch to template-instantiated kernel based on K
auto launch = [&](auto k_tag) {
fused_mxfp4_gemm_kernel<decltype(k_tag)::value><<<grid, block>>>(
reinterpret_cast<const __hip_bfloat16*>(A.data_ptr<at::BFloat16>()),
B_q.data_ptr<uint8_t>(),
B_scale_sh.data_ptr<uint8_t>(),
reinterpret_cast<__hip_bfloat16*>(C.data_ptr<at::BFloat16>()),
M, N, K, scaleN_pad
);
};
if (K == 512) launch(std::integral_constant<int, 512>{});
else launch(std::integral_constant<int, 512>{}); // fallback, shouldnt happen for K<=512
}
// ======================================================================
// KERNEL 3: Fused MFMA GEMM with splitK — splits K across blockIdx.z
// Each block computes a partial 16x16 tile for its K-split range,
// and stores fp32 partial results to partial_out[k_split, M_padded, N].
// No atomicAdd needed — each K-split writes to its own slice.
// ======================================================================
__launch_bounds__(64, 1)
__global__ void fused_mxfp4_gemm_splitk_kernel(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
float* __restrict__ partial_out, // [num_k_splits, M_padded, N] fp32
const int M, const int N, const int K,
const int scaleN_pad, const int num_k_splits, const int M_padded
) {
#if defined(__gfx950__)
const int lane = threadIdx.x; // 0-63
const int lane16 = lane & 15;
const int group4 = lane >> 4; // 0,1,2,3
const int tile_m = blockIdx.x;
const int tile_n = blockIdx.y;
const int k_split = blockIdx.z; // which K-split
const int out_row_base = tile_m * 16;
const int out_col = tile_n * 16 + lane16;
const int a_row = out_row_base + lane16;
const int K_half = K / 2;
const int sn8 = scaleN_pad / 8;
// Compute K range for this split
// K is always multiple of 128; divide K/128 steps among splits
const int total_k_steps = K / 128;
const int steps_per_split = (total_k_steps + num_k_splits - 1) / num_k_splits;
const int k_step_start = k_split * steps_per_split;
const int k_step_end_raw = k_step_start + steps_per_split;
const int k_step_end = (k_step_end_raw < total_k_steps) ? k_step_end_raw : total_k_steps;
const int k_start = k_step_start * 128;
const int k_end = k_step_end * 128;
// If this split has no work, bail out
if (k_start >= k_end) return;
// Pre-compute B scale shuffle index components
const int b_i0 = out_col / 32;
const int b_rem32 = out_col & 31;
const int b_i1 = b_rem32 / 16;
const int b_i2 = b_rem32 & 15;
const int b_i0_stride = b_i0 * (sn8 * 256);
const int b_i2_i1_part = b_i2 * 4 + b_i1;
// Initialize accumulator
fp32x4_t c_reg = {};
i32x8_t a_reg = {};
i32x8_t b_reg = {};
int scale_a_val = 0;
int scale_b_val = 127;
// K loop over this split's range
for (int k_iter = k_start; k_iter < k_end; k_iter += 128) {
const int k_quarter_start = k_iter + group4 * 32;
// ============ Load and quantize A ============
const int safe_row = (a_row < M) ? a_row : 0;
const __hip_bfloat16* a_ptr = A + (int64_t)safe_row * K + k_quarter_start;
float vals[32];
float amax_val = 0.0f;
const uint4* a_ptr128 = reinterpret_cast<const uint4*>(a_ptr);
#pragma unroll
for (int j = 0; j < 4; j++) {
uint4 data = a_ptr128[j];
const __hip_bfloat16* bv = reinterpret_cast<const __hip_bfloat16*>(&data);
#pragma unroll
for (int i = 0; i < 8; i++) {
float v = __bfloat162float(bv[i]);
vals[j * 8 + i] = v;
amax_val = fmaxf(amax_val, fabsf(v));
}
}
uint32_t amax_bits = __float_as_uint(amax_val);
amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
uint32_t exponent = (amax_bits >> 23) & 0xFFu;
uint32_t a_e8m0 = (exponent >= 2u) ? (exponent - 2u) : 0u;
float inv_scale = (a_e8m0 > 0u) ? __uint_as_float(a_e8m0 << 23) : __uint_as_float(0x7F800000u);
scale_a_val = (a_row < M) ? (int)a_e8m0 : 0;
#pragma unroll
for (int g = 0; g < 4; g++) {
uint32_t w = 0;
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+0], vals[g*8+1], inv_scale, 0);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+2], vals[g*8+3], inv_scale, 1);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+4], vals[g*8+5], inv_scale, 2);
w = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w, vals[g*8+6], vals[g*8+7], inv_scale, 3);
a_reg[g] = __builtin_bit_cast(int, w);
}
if (a_row >= M) { a_reg = {}; scale_a_val = 0; }
// ============ Load B ============
b_reg = {};
scale_b_val = 127;
if (out_col < N) {
const uint4* b_src128 = reinterpret_cast<const uint4*>(
B_q + (int64_t)out_col * K_half + k_quarter_start / 2);
uint4 b_data = b_src128[0];
b_reg[0] = __builtin_bit_cast(int, b_data.x);
b_reg[1] = __builtin_bit_cast(int, b_data.y);
b_reg[2] = __builtin_bit_cast(int, b_data.z);
b_reg[3] = __builtin_bit_cast(int, b_data.w);
int scale_col = k_iter / 32 + group4;
int i3 = scale_col / 8;
int sc_rem8 = scale_col & 7;
int i4 = sc_rem8 / 4;
int i5 = sc_rem8 & 3;
int shuffled_idx = b_i0_stride + i3 * 256 + i5 * 64 + b_i2_i1_part + i4 * 2;
scale_b_val = (int)B_scale_sh[shuffled_idx];
}
// ============ Execute MFMA ============
c_reg = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_reg, b_reg, c_reg,
4, 4, 0, scale_a_val, 0, scale_b_val
);
}
// ============ Store: direct write to partial_out[k_split, :, :] ============
if (out_col < N) {
const int slice_offset = k_split * M_padded * N;
#pragma unroll
for (int i = 0; i < 4; i++) {
const int r = out_row_base + group4 * 4 + i;
partial_out[slice_offset + r * N + out_col] = (r < M) ? c_reg[i] : 0.0f;
}
}
#endif // __gfx950__
}
// ======================================================================
// KERNEL 4: Reduce partial K-splits and convert fp32 -> bf16
// ======================================================================
__global__ void reduce_and_convert_kernel(
const float* __restrict__ partial, // [num_k_splits, M_padded, N]
__hip_bfloat16* __restrict__ out, // [M_padded, N]
const int M_padded, const int N, const int num_k_splits
) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
int total = M_padded * N;
if (idx >= total) return;
float sum = 0.0f;
for (int k = 0; k < num_k_splits; k++) {
sum += partial[k * total + idx];
}
out[idx] = __float2bfloat16(sum);
}
// C++ wrapper for fused splitK GEMM
void fused_mxfp4_gemm_splitk(
torch::Tensor A, // [M, K] bf16
torch::Tensor B_q, // [N, K/2] uint8 fp4x2
torch::Tensor B_scale_sh, // shuffled E8M0 scales
torch::Tensor partial, // [num_k_splits, M_padded, N] fp32 partial buffer
torch::Tensor C, // [M_padded, N] bf16 output
int N_dim,
int num_k_splits
) {
const int M = A.size(0);
const int K = A.size(1);
const int N = N_dim;
const int K_scale = K / 32;
const int scaleN_pad = ((K_scale + 7) / 8) * 8;
const int M_padded = C.size(0);
// Launch splitK kernel: 3D grid (M_tiles, N_tiles, num_k_splits)
dim3 grid(
(M_padded + 15) / 16,
(N + 15) / 16,
num_k_splits
);
dim3 block(64);
fused_mxfp4_gemm_splitk_kernel<<<grid, block>>>(
reinterpret_cast<const __hip_bfloat16*>(A.data_ptr<at::BFloat16>()),
B_q.data_ptr<uint8_t>(),
B_scale_sh.data_ptr<uint8_t>(),
partial.data_ptr<float>(),
M, N, K, scaleN_pad, num_k_splits, M_padded
);
// Launch reduce + convert kernel: sum over K-splits and convert to bf16
int total = M_padded * N;
int conv_threads = 256;
int conv_blocks = (total + conv_threads - 1) / conv_threads;
reduce_and_convert_kernel<<<conv_blocks, conv_threads>>>(
partial.data_ptr<float>(),
reinterpret_cast<__hip_bfloat16*>(C.data_ptr<at::BFloat16>()),
M_padded, N, num_k_splits
);
}
"""
CPP_SRC = """
#include <torch/extension.h>
#include <vector>
std::vector<torch::Tensor> mxfp4_quant_fused(torch::Tensor x_in);
void mxfp4_quant_inplace(torch::Tensor x_in, torch::Tensor x_fp4_out, torch::Tensor scale_out);
void fused_mxfp4_gemm(torch::Tensor A, torch::Tensor B_q, torch::Tensor B_scale_sh, torch::Tensor C, int N_dim);
void fused_mxfp4_gemm_splitk(torch::Tensor A, torch::Tensor B_q, torch::Tensor B_scale_sh, torch::Tensor partial, torch::Tensor C, int N_dim, int num_k_splits);
"""
_module = load_inline(
name='mxfp4_quant_hip',
cpp_sources=[CPP_SRC],
cuda_sources=[HIP_SRC],
functions=['mxfp4_quant_fused', 'mxfp4_quant_inplace', 'fused_mxfp4_gemm', 'fused_mxfp4_gemm_splitk'],
verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
)
# Keep aiter imports for server compatibility
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
_fused_gemm = _module.fused_mxfp4_gemm
_fused_gemm_splitk = _module.fused_mxfp4_gemm_splitk
_quant = _module.mxfp4_quant_fused
_quant_ip = _module.mxfp4_quant_inplace
_gemm = gemm_a4w4_asm
_fp4x2 = dtypes.fp4x2
_e8m0 = dtypes.fp8_e8m0
_bf16 = torch.bfloat16
_f32 = torch.float32
_empty = torch.empty
_zeros = torch.zeros
_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_shape_cache = {} # (m, k, n) -> dict with all pre-computed values for this shape
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B_shuffle.shape[0]
# Cache only shape-dependent buffers. Ranked mode can reuse a shape with new B tensors.
cache_key = (m, k, n)
sc = _shape_cache.get(cache_key)
if sc is None:
sc = _build_shape_cache(m, k, n, A.device)
_shape_cache[cache_key] = sc
return sc[1](A, B_q, B_scale_sh, B_shuffle)
def _build_shape_cache(m, k, n, device):
"""Build all pre-computed data for a given (m, k, n) shape. Called once."""
if k <= 512:
# Fused MFMA path
m_padded = (m + 15) & ~15 # bitwise round up to 16
out = _empty((m_padded, n), dtype=_bf16, device=device)
out_slice = out[:m]
_fg = _fused_gemm
_o, _os = out, out_slice
def _run_fused(A, B_q, B_scale_sh, B_shuffle, _fg=_fg, _o=_o, _os=_os, _n=n):
_fg(A, B_q.view(torch.uint8), B_scale_sh.view(torch.uint8), _o, _n)
return _os
return (0, _run_fused)
if m <= 16:
# SplitK path
m_padded = (m + 15) & ~15
out = _empty((m_padded, n), dtype=_bf16, device=device)
out_slice = out[:m]
k_steps = k >> 7 # k // 128
num_k_splits = min(7, k_steps)
while num_k_splits > 1 and k_steps % num_k_splits != 0:
num_k_splits -= 1
partial = _empty((num_k_splits, m_padded, n), dtype=_f32, device=device)
_fsk = _fused_gemm_splitk
def _run_splitk(A, B_q, B_scale_sh, B_shuffle, _fsk=_fsk, _p=partial, _o=out, _os=out_slice, _n=n, _nks=num_k_splits):
_fsk(A, B_q.view(torch.uint8), B_scale_sh.view(torch.uint8), _p, _o, _n, _nks)
return _os
return (1, _run_splitk)
# ASM GEMM path
m_padded = (m + 31) & ~31 # bitwise round up to 32
out = _empty((m_padded, n), dtype=_bf16, device=device)
out_slice = out[:m]
# Pre-allocate quant buffers
scaleN_valid = (k + 31) >> 5 # k // 32 rounded up
scaleN_pad = (scaleN_valid + 7) & ~7 # round up to 8
padM256 = (m + 31) & ~31 # round up to 32 (matches MFMA tile alignment)
A_q_buf = _empty((m, k >> 1), dtype=torch.uint8, device=device)
A_scale_buf = _empty((padM256 * scaleN_pad,), dtype=torch.uint8, device=device)
A_q_view = A_q_buf.view(_fp4x2).view(m, k >> 1)
A_scale_view = A_scale_buf.view(padM256, scaleN_pad).view(_e8m0)
# Pre-bound closure: eliminates tuple unpacking, attribute lookups, and
# argument construction from the hot path
_qi = _quant_ip
_gm = _gemm
_kn = _KERNEL
_out = out
# Ranked-valid probe: only the visible large-M (64) shape tries K-split in ASM.
_log2_k_split = 1 if (m, k, n) == (64, 2048, 7168) else 0
def _run_asm(A, B_q, B_scale_sh, B_shuffle):
_qi(A, A_q_buf, A_scale_buf)
_gm(A_q_view, B_shuffle, A_scale_view, B_scale_sh, _out,
_kn, None, 1.0, 0.0, True, log2_k_split=_log2_k_split)
return out_slice
return (2, _run_asm)
scrolls · 703 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