submission 638406
Divyanshsingh1910 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 439 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-638406?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:6014338e610bf736704fe30a6638f80a8b572141bab449e60f1ad0001b2718c4
license declaredunknown
license concludedunknown
authorsDivyanshsingh1910
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
- A: load bf16 from global, quant to FP4 in registers, write to LDSshared-memory
__shared__ char A_lds[2][BM * A_LDS_ROW];tile-k = 256
- BM=32, BN=128, BK=256, 256 threads (4 waves)tile-m = 32
- BM=32, BN=128, BK=256, 256 threads (4 waves)tile-n = 128
- BM=32, BN=128, BK=256, 256 threads (4 waves)vector-width = uint4_t
using uint4_t = uint4;Kernel source
submission.py439 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
=== v21_fused: Single fused kernel ===
bf16 A -> quantize in registers -> MFMA 16x16x128 with B_shuffle -> bf16 C
Eliminates 3 kernel launches (quant + shuffle + gemm) into 1.
Architecture:
- BM=32, BN=128, BK=256, 256 threads (4 waves)
- A: load bf16 from global, quant to FP4 in registers, write to LDS
- B: load directly from B_shuffle to VGPRs (no LDS)
- B_scale: cooperative load to LDS with shuffled index decode
- A_scale: computed during quant, stays in registers
- Double-buffered A LDS + B_scale LDS
"""
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 <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
using fp4x2_t = unsigned char;
using fp4x64_t = fp4x2_t __attribute__((ext_vector_type(32)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
using uint4_t = uint4;
constexpr int BM = 32;
constexpr int BN = 128;
constexpr int BK = 256;
constexpr int A_LDS_ROW = BK / 2 + 4; // 132 bytes padded
// ============================================================
// Branchless FP4 E2M1 quantization (RNE, matches aiter exactly)
// ============================================================
// Quantize 32 bf16 values to FP4, return packed 16 bytes + E8M0 scale byte
// Input: 32 bf16 values in vals[0..31]
// Output: packed[0..15] (fp4x2 bytes), scale_byte
__device__ __forceinline__ void quant_group_32(
const __hip_bfloat16* vals, unsigned char* packed, unsigned char& scale_byte
) {
// Step 1: Find amax
float amax = 0.0f;
float fvals[32];
#pragma unroll
for (int i = 0; i < 32; i++) {
fvals[i] = __bfloat162float(vals[i]);
float av = fvals[i] < 0.0f ? -fvals[i] : fvals[i];
if (av > amax) amax = av;
}
// Step 2: E8M0 scale = round amax to nearest power of 2, subtract 2 from exponent
// aiter: scale_e8m0_unbiased = floor(log2(amax_rounded)) - 2
// The -2 maps FP4 E2M1 range [0,6] to normalized [0, ~4-6] for full utilization
unsigned int amax_u32 = __float_as_uint(amax);
if (amax == 0.0f) {
scale_byte = 0;
} else {
unsigned int rounded_bits = (amax_u32 + 0x200000u) & 0xFF800000u;
unsigned int exp_raw = (rounded_bits >> 23) & 0xFFu;
scale_byte = (exp_raw >= 2u) ? (unsigned char)(exp_raw - 2u) : (unsigned char)0;
}
// Reconstruct scale from adjusted E8M0 for quantization: scale = 2^(scale_byte - 127)
float scale_float = (scale_byte > 0) ? __uint_as_float(((unsigned int)scale_byte) << 23) : 0.0f;
float inv_scale = (scale_float > 0.0f) ? (1.0f / scale_float) : 0.0f;
// Step 3: Quantize each value to FP4 and pack pairs
#pragma unroll
for (int i = 0; i < 16; i++) {
unsigned char lo, hi;
// Even element (low nibble)
{
float x = fvals[2 * i];
float ax = x < 0.0f ? -x : x;
float normed = ax * inv_scale;
// Branchless 7-comparison RNE
int fp4_val = (normed > 0.25f) + (normed >= 0.75f) + (normed > 1.25f)
+ (normed >= 1.75f) + (normed > 2.5f) + (normed >= 3.5f)
+ (normed > 5.0f);
// Sign bit
int sign = (x < 0.0f) ? 8 : 0;
lo = (unsigned char)(fp4_val | sign);
}
// Odd element (high nibble)
{
float x = fvals[2 * i + 1];
float ax = x < 0.0f ? -x : x;
float normed = ax * inv_scale;
int fp4_val = (normed > 0.25f) + (normed >= 0.75f) + (normed > 1.25f)
+ (normed >= 1.75f) + (normed > 2.5f) + (normed >= 3.5f)
+ (normed > 5.0f);
int sign = (x < 0.0f) ? 8 : 0;
hi = (unsigned char)(fp4_val | sign);
}
packed[i] = (hi << 4) | (lo & 0x0F);
}
}
// ============================================================
// B_shuffle direct VGPR load (same as v20)
// ============================================================
__device__ __forceinline__ fp4x64_t load_b_direct(
const char* __restrict__ B_sh_ptr,
int global_n_tile, int k_tile_base, int total_k_tiles,
int group4, int lane16
) {
int k_tile_off = group4 >> 1;
int sub_group = group4 & 1;
const char* tile_ptr = B_sh_ptr +
(int64_t)(global_n_tile * total_k_tiles + k_tile_base + k_tile_off) * 512;
const char* src = tile_ptr + sub_group * 256 + lane16 * 16;
i32x4_t raw = *reinterpret_cast<const i32x4_t*>(src);
i32x8_t full = {raw[0], raw[1], raw[2], raw[3], 0, 0, 0, 0};
return __builtin_bit_cast(fp4x64_t, full);
}
// ============================================================
// Read A from LDS for 16x16x128 MFMA
// ============================================================
__device__ __forceinline__ fp4x64_t load_a_from_lds(
const char* __restrict__ lds,
int mt, int ki, int lane16, int group4
) {
int row = mt * 16 + lane16;
int col = ki * 64 + group4 * 16;
const char* ptr = lds + row * A_LDS_ROW + col;
i32x4_t raw = *reinterpret_cast<const i32x4_t*>(ptr);
i32x8_t full = {raw[0], raw[1], raw[2], raw[3], 0, 0, 0, 0};
return __builtin_bit_cast(fp4x64_t, full);
}
// ============================================================
// B_scale unshuffle index (inline, from shuffled layout)
// ============================================================
__device__ __forceinline__ unsigned char load_b_scale_unshuffled(
const unsigned char* __restrict__ B_scale_sh,
int n_row, int k_scale_idx, int N, int K_div_32
) {
int sm = ((N + 255) / 256) * 256;
int sn = ((K_div_32 + 7) / 8) * 8;
(void)sm;
int d0 = n_row / 32;
int d1 = (n_row & 31) / 16;
int d2 = n_row & 15;
int d3 = k_scale_idx / 8;
int d4 = (k_scale_idx & 7) / 4;
int d5 = k_scale_idx & 3;
int idx = d0 * (sn / 8 * 4 * 16 * 2 * 2)
+ d3 * (4 * 16 * 2 * 2)
+ d5 * (16 * 2 * 2)
+ d2 * (2 * 2)
+ d4 * 2
+ d1;
return B_scale_sh[idx];
}
// ============================================================
// Cooperative A quant: load bf16 from global, quant to FP4, write to LDS
// Also writes A_scale bytes to a_scale_buf for later use
//
// BM=32 rows, BK=256 cols of bf16 A
// Each row: 256 bf16 = 8 groups of 32 = 8 scale bytes, 128 packed FP4 bytes
// 256 threads total, BM*BK = 32*256 = 8192 bf16 values
// Each thread: 8192/256 = 32 bf16 values = 1 group of 32
// So each thread quants exactly 1 group and produces 16 packed bytes + 1 scale byte
// ============================================================
__device__ __forceinline__ void coop_quant_A(
const __hip_bfloat16* __restrict__ A_bf16,
char* __restrict__ a_lds,
unsigned char* __restrict__ a_scale_lds, // BM * (BK/32) = 32 * 8 = 256 bytes
int tile_m, int k, int M, int K, int tid
) {
// tid in [0,256): each handles 1 group of 32 bf16 values
// 256 groups = 32 rows * 8 groups_per_row
int row = tid / 8; // 0..31
int grp = tid % 8; // 0..7 (which group of 32 within the row)
int g_row = tile_m + row;
__hip_bfloat16 vals[32];
if (g_row < M) {
int col_start = k + grp * 32;
const __hip_bfloat16* src = A_bf16 + (int64_t)g_row * K + col_start;
// Load 32 bf16 values (64 bytes = 4 x uint4)
if (col_start + 31 < K) {
*reinterpret_cast<uint4_t*>(&vals[0]) = *reinterpret_cast<const uint4_t*>(&src[0]);
*reinterpret_cast<uint4_t*>(&vals[8]) = *reinterpret_cast<const uint4_t*>(&src[8]);
*reinterpret_cast<uint4_t*>(&vals[16]) = *reinterpret_cast<const uint4_t*>(&src[16]);
*reinterpret_cast<uint4_t*>(&vals[24]) = *reinterpret_cast<const uint4_t*>(&src[24]);
} else {
#pragma unroll
for (int i = 0; i < 32; i++)
vals[i] = (col_start + i < K) ? src[i] : __float2bfloat16(0.0f);
}
} else {
#pragma unroll
for (int i = 0; i < 32; i++)
vals[i] = __float2bfloat16(0.0f);
}
// Quantize
unsigned char packed[16];
unsigned char scale_byte;
quant_group_32(vals, packed, scale_byte);
// Write packed FP4 to A LDS
// Row layout in LDS: A_LDS_ROW bytes per row (132 bytes, 128 data + 4 pad)
// Group grp occupies bytes [grp*16 .. grp*16+15] within the row
char* dst = a_lds + row * A_LDS_ROW + grp * 16;
*reinterpret_cast<i32x4_t*>(dst) = *reinterpret_cast<i32x4_t*>(packed);
// Write scale byte
// Layout: a_scale_lds[row * 8 + grp]
a_scale_lds[row * 8 + grp] = scale_byte;
}
// ============================================================
// Main fused kernel
// ============================================================
__global__
__attribute__((amdgpu_flat_work_group_size(256, 256)))
__attribute__((amdgpu_waves_per_eu(2)))
void fused_mxfp4_gemm(
const __hip_bfloat16* __restrict__ A_bf16,
const char* __restrict__ B_shuffle,
const unsigned char* __restrict__ B_scale_sh,
__hip_bfloat16* __restrict__ C,
int M, int N, int K
) {
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 63;
const int lane16 = lane & 15;
const int group4 = lane >> 4;
const int tile_m = blockIdx.y * BM;
const int tile_n = blockIdx.x * BN;
if (tile_m >= M) return;
const int wave_n_base = tile_n + wave_id * 32;
const int K_div_32 = K / 32;
const int total_k_tiles = (K / 2) / 32;
const int gnt0 = (wave_n_base) / 16;
const int gnt1 = (wave_n_base + 16) / 16;
f32x4_t acc00 = {};
f32x4_t acc01 = {};
f32x4_t acc10 = {};
f32x4_t acc11 = {};
// LDS layout:
// A_lds: double-buffered, 2 * BM * A_LDS_ROW = 2 * 32 * 132 = 8448 bytes
// A_scale_lds: double-buffered, 2 * BM * (BK/32) = 2 * 32 * 8 = 512 bytes
__shared__ char A_lds[2][BM * A_LDS_ROW];
__shared__ unsigned char A_scale_lds[2][BM * (BK / 32)]; // 32 * 8 = 256 per buf
// Prefetch first A tile: quant bf16 -> FP4 in LDS
coop_quant_A(A_bf16, A_lds[0], A_scale_lds[0], tile_m, 0, M, K, tid);
__builtin_amdgcn_s_barrier();
int buf = 0;
for (int k = 0; k < K; k += BK) {
int next_k = k + BK;
// Prefetch next A tile (double buffer)
if (next_k < K) {
coop_quant_A(A_bf16, A_lds[1 - buf], A_scale_lds[1 - buf],
tile_m, next_k, M, K, tid);
}
int k_tile_base = (k / 2) / 32;
#pragma unroll
for (int ki = 0; ki < 2; ki++) { // BK/128 = 2
int k_abs = k + ki * 128;
int k_scale_base = k_abs / 32; // 4 scales per 128 FP4
// Load A from LDS
fp4x64_t a0 = load_a_from_lds(A_lds[buf], 0, ki, lane16, group4);
fp4x64_t a1 = load_a_from_lds(A_lds[buf], 1, ki, lane16, group4);
// Load B directly from global
fp4x64_t b0 = load_b_direct(B_shuffle, gnt0, k_tile_base + ki * 2,
total_k_tiles, group4, lane16);
fp4x64_t b1 = load_b_direct(B_shuffle, gnt1, k_tile_base + ki * 2,
total_k_tiles, group4, lane16);
// A scales from LDS (computed during quant)
// For 16x16x128 MFMA: lane16 = M-row within sub-tile, group4 = K-quarter
// A_scale_lds[row][k_scale_base - k/32 + group4] since scale indices are relative to BK chunk
// k_scale_base = (k + ki*128)/32, within BK chunk: ki*4 + group4
int a_scale_offset_0 = lane16 * 8 + ki * 4 + group4; // mt=0, row=lane16
int a_scale_offset_1 = (16 + lane16) * 8 + ki * 4 + group4; // mt=1, row=16+lane16
unsigned char sa0 = A_scale_lds[buf][a_scale_offset_0];
unsigned char sa1 = A_scale_lds[buf][a_scale_offset_1];
// Clamp sa for out-of-bounds M rows
if (tile_m + lane16 >= M) sa0 = 127;
if (tile_m + 16 + lane16 >= M) sa1 = 127;
// B scales from global (inline unshuffle)
int b_col0 = wave_n_base + lane16;
int b_col1 = wave_n_base + 16 + lane16;
unsigned char sb0 = 127, sb1 = 127;
if (b_col0 < N)
sb0 = load_b_scale_unshuffled(B_scale_sh, b_col0, k_scale_base + group4, N, K_div_32);
if (b_col1 < N)
sb1 = load_b_scale_unshuffled(B_scale_sh, b_col1, k_scale_base + group4, N, K_div_32);
// 4 MFMAs: [mt][nt]
acc00 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a0, b0, acc00, 4, 4, 0, (unsigned int)sa0, 0, (unsigned int)sb0);
acc01 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a0, b1, acc01, 4, 4, 0, (unsigned int)sa0, 0, (unsigned int)sb1);
acc10 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a1, b0, acc10, 4, 4, 0, (unsigned int)sa1, 0, (unsigned int)sb0);
acc11 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a1, b1, acc11, 4, 4, 0, (unsigned int)sa1, 0, (unsigned int)sb1);
}
if (next_k < K) {
__builtin_amdgcn_s_barrier();
buf = 1 - buf;
}
}
// Write back results
auto write_tile = [&](f32x4_t& acc_reg, int m_base, int n_col) {
if (n_col >= N) return;
#pragma unroll
for (int i = 0; i < 4; i++) {
int row = m_base + group4 * 4 + i;
if (row < M)
C[row * N + n_col] = __float2bfloat16(acc_reg[i]);
}
};
int n_col0 = wave_n_base + lane16;
int n_col1 = wave_n_base + 16 + lane16;
write_tile(acc00, tile_m, n_col0);
write_tile(acc01, tile_m, n_col1);
write_tile(acc10, tile_m + 16, n_col0);
write_tile(acc11, tile_m + 16, n_col1);
}
torch::Tensor fused_gemm(
torch::Tensor A_bf16,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh,
int M, int N, int K, int sn
) {
auto C = torch::empty({M, N},
torch::TensorOptions().dtype(torch::kBFloat16).device(A_bf16.device()));
dim3 grid((N + BN - 1) / BN, (M + BM - 1) / BM);
dim3 block(256);
auto a_ptr = reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr());
auto b_ptr = reinterpret_cast<const char*>(B_shuffle.data_ptr<uint8_t>());
auto bs_ptr = reinterpret_cast<const unsigned char*>(B_scale_sh.data_ptr<uint8_t>());
auto c_ptr = reinterpret_cast<__hip_bfloat16*>(C.data_ptr());
fused_mxfp4_gemm<<<grid, block>>>(a_ptr, b_ptr, bs_ptr, c_ptr, M, N, K);
return C;
}
"""
CPP_SRC = """
torch::Tensor fused_gemm(
torch::Tensor A_bf16,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh,
int M, int N, int K, int sn
);
"""
module = load_inline(
name='mxfp4_fused_v21b',
cpp_sources=[CPP_SRC],
cuda_sources=[HIP_SRC],
functions=['fused_gemm'],
verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3",
"-mllvm", "-amdgpu-early-inline-all=true"],
)
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B_q.shape[0]
# Dispatch: fused kernel for small K, aiter for large K
# Fused kernel wins on K=512 (11.6 us vs 18.8 us)
# aiter wins on K>=1536 (22.9-33 us vs 26-97 us)
if k <= 1024:
B_sh_u8 = B_shuffle.view(torch.uint8)
B_sc_u8 = B_scale_sh.view(torch.uint8)
sn = ((k // 32 + 7) // 8) * 8
return module.fused_gemm(A, B_sh_u8, B_sc_u8, m, n, k, sn)
else:
A_q, A_scale = dynamic_mxfp4_quant(A)
A_scale_sh = e8m0_shuffle(A_scale)
return aiter.gemm_a4w4(
A_q.view(dtypes.fp4x2), B_shuffle,
A_scale_sh.view(dtypes.fp8_e8m0), B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
scrolls · 439 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