submission 651769
kfz · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 542 lines, June 9 Researcher Reciprocity License v1.0.
sub_v31_hip.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-651769?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:159cf5fbcf42abd360737d9161260d361f7cbd9af94444bbc9ff260bb3fcd20b
license declaredunknown
license concludedunknown
authorskfz
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ uint8_t smem[];split-k
template<int WM, int WN, int GM, int GN, bool SPLIT_K, bool USE_HW_CVT, int PIPE_DEPTH>Kernel source
sub_v31_hip.py542 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# Variant 31: V29 pipeline depth + aggressive compiler flags + hardcoded best configs
# Compiler: -ffast-math -ffp-contract=fast -munsafe-fp-atomics
# Best configs from V29/V30 auto-tune (hardcoded, no warmup auto-tune overhead)
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
import sys
CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <stdint.h>
#include <math.h>
#if defined(__gfx950__)
typedef int i32x8_t __attribute__((ext_vector_type(8)));
typedef float f32x4_t __attribute__((ext_vector_type(4)));
__device__ __forceinline__ void async_global_load_lds_b16(
const uint8_t* src, uint32_t m0_val)
{
asm volatile(
"\n\ts_mov_b32 m0, %1"
"\n\tglobal_load_lds_dwordx4 %0, off"
:
: "v"(src), "s"(m0_val)
: "memory"
);
}
__device__ __forceinline__ void async_lds_fence() {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
}
__device__ __forceinline__ unsigned int compute_e8m0(float amax) {
union { float f; unsigned int u; } v;
v.f = amax;
v.u = (v.u + 0x200000u) & 0xFF800000u;
if (v.u == 0u) return 127u;
unsigned int biased_exp = (v.u >> 23) & 0xFFu;
int e8m0 = (int)biased_exp - 2;
return (unsigned int)max(0, min(254, e8m0));
}
__device__ __forceinline__ uint8_t float_to_fp4(float x) {
union { float f; unsigned int u; } v;
v.f = x;
unsigned int sign = v.u & 0x80000000u;
v.u ^= sign;
float ax = v.f;
uint8_t code;
if (ax >= 6.0f) {
code = 7u;
} else if (ax >= 1.0f) {
unsigned int mant_odd = (v.u >> 22) & 1u;
v.u += 0xC11FFFFFu;
v.u += mant_odd;
code = (uint8_t)((v.u >> 22) & 0xFu);
} else {
union { float f; unsigned int u; } d;
d.f = ax + 4194304.0f;
code = (uint8_t)((d.u - 0x4A800000u) & 0xFFu);
}
return code | (uint8_t)(sign >> 28);
}
#endif
// ═══════════════════════════════════════════════════════════════
// Kernel with PIPE_DEPTH template parameter
// PIPE_DEPTH=2: standard double buffer (V26 behavior)
// PIPE_DEPTH=3: triple buffer
// PIPE_DEPTH=4: quad buffer (full prefetch for K_steps=4)
// ═══════════════════════════════════════════════════════════════
template<int WM, int WN, int GM, int GN, bool SPLIT_K, bool USE_HW_CVT, int PIPE_DEPTH>
__global__ void fp4_gemm_fused_kernel(
const __hip_bfloat16* __restrict__ A_bf16,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
__hip_bfloat16* __restrict__ C_bf16,
float* __restrict__ C_f32,
int M, int N, int K, int chunk_Ks
) {
#if defined(__gfx950__)
constexpr int N_WARPS = GM * GN;
constexpr int BLOCK_MT = WM * GM;
constexpr int BLOCK_NT = WN * GN;
constexpr int BLOCK_M = BLOCK_MT * 16;
constexpr int BLOCK_N = BLOCK_NT * 16;
constexpr int B_ROW_BYTES = 64;
constexpr int B_TILE_BYTES = 16 * B_ROW_BYTES;
constexpr int B_BUF = BLOCK_NT * B_TILE_BYTES;
constexpr int A_ROW_BYTES = 256;
constexpr int A_TILE_BYTES = 16 * A_ROW_BYTES;
constexpr int A_BUF = BLOCK_MT * A_TILE_BYTES;
constexpr int S_CHUNK = 256;
constexpr int S_N_GROUPS = (BLOCK_N + 31) / 32;
constexpr int S_LDS_BUF = 1024;
// SMEM layout: PIPE_DEPTH buffers instead of 2
constexpr int B_BASE = 0;
constexpr int A_BASE = PIPE_DEPTH * B_BUF;
constexpr int S_BASE = A_BASE + PIPE_DEPTH * A_BUF;
extern __shared__ uint8_t smem[];
const int tid = threadIdx.x;
const int warp_id = tid / 64;
const int lane = tid % 64;
const int lane16 = lane & 15;
const int k_quarter = lane >> 4;
const int warp_row = warp_id / GN;
const int warp_col = warp_id % GN;
const int m_block = blockIdx.y * BLOCK_M;
const int n_block = blockIdx.x * BLOCK_N;
const int m_warp = m_block + warp_row * WM * 16;
const int n_warp = n_block + warp_col * WN * 16;
const int n_block_i0 = n_block >> 5;
const int half_K = K >> 1;
const int k_chunk = SPLIT_K ? (int)blockIdx.z : 0;
const int ks_base = k_chunk * chunk_Ks;
const int byte_base = ks_base * 16;
const int K_steps = chunk_Ks >> 2;
f32x4_t acc[WM][WN];
#pragma unroll
for (int wm = 0; wm < WM; wm++)
#pragma unroll
for (int wn = 0; wn < WN; wn++)
for (int i = 0; i < 4; i++)
acc[wm][wn][i] = 0.0f;
const int col_out = lane & 15;
const int row_base_out = (lane >> 4) * 4;
#define ASYNC_LOAD_B_TILES(BUF, KSTEP) do { \
int _kbyte = byte_base + (KSTEP) * 64; \
int _load_row = lane >> 2; \
int _load_quarter = lane & 3; \
for (int _task = warp_id; _task < BLOCK_NT; _task += N_WARPS) { \
uint32_t _tile_lds = (uint32_t)(B_BASE + (BUF) * B_BUF + _task * B_TILE_BYTES); \
int _g = n_block + _task * 16 + _load_row; \
int _g_safe = _g < N ? _g : 0; \
const uint8_t* _src = B_q + _g_safe * half_K + _kbyte + _load_quarter * 16; \
uint32_t _m0 = __builtin_amdgcn_readfirstlane((int)_tile_lds); \
async_global_load_lds_b16(_src, _m0); \
} \
} while(0)
#define ASYNC_LOAD_A_TILES(BUF, KSTEP) do { \
int _k_elem = ((int)ks_base + (KSTEP) * 4) * 32; \
for (int _task = warp_id; _task < BLOCK_MT * 4; _task += N_WARPS) { \
int _tile = _task >> 2; \
int _sub = _task & 3; \
uint32_t _tile_lds = (uint32_t)(A_BASE + (BUF) * A_BUF \
+ _tile * A_TILE_BYTES + _sub * 1024); \
int _chunk_base = _sub * 64 + lane; \
int _row = _chunk_base >> 4; \
int _col_chunk = _chunk_base & 15; \
int _m_g = m_block + _tile * 16 + _row; \
int _m_safe = (_m_g < M) ? _m_g : 0; \
const uint8_t* _src = (const uint8_t*)(A_bf16 + (int64_t)_m_safe * K + _k_elem) \
+ _col_chunk * 16; \
uint32_t _m0 = __builtin_amdgcn_readfirstlane((int)_tile_lds); \
async_global_load_lds_b16(_src, _m0); \
} \
} while(0)
#define ASYNC_LOAD_B_SCALE(BUF, KSTEP) do { \
if (warp_id == 0) { \
int _i3 = ((int)ks_base + (KSTEP) * 4) >> 3; \
int _group = lane >> 4; \
int _chunk_in_group = lane & 15; \
int _i0 = n_block_i0 + _group; \
int _i0_safe = (_group < S_N_GROUPS && _i0 * 32 < N) ? _i0 : 0; \
const uint8_t* _src = B_scale_sh + _i0_safe * K \
+ _i3 * 256 + _chunk_in_group * 16; \
uint32_t _s_lds = (uint32_t)(S_BASE + (BUF) * S_LDS_BUF); \
uint32_t _m0 = __builtin_amdgcn_readfirstlane((int)_s_lds); \
async_global_load_lds_b16(_src, _m0); \
} \
} while(0)
#define GET_B_SCALE_LDS(BUF, N_COL, KS) \
((unsigned int)smem[S_BASE + (BUF) * S_LDS_BUF \
+ (((N_COL) >> 5) - n_block_i0) * S_CHUNK \
+ ((KS) & 3) * 64 \
+ ((N_COL) & 15) * 4 \
+ (((KS) >> 2) & 1) * 2 \
+ (((N_COL) >> 4) & 1)])
i32x8_t a_frag[WM];
int a_scale_packed[WM];
#define QUANTIZE_A_FROM_LDS(BUF) do { \
_Pragma("unroll") \
for (int _wm = 0; _wm < WM; _wm++) { \
int _wm_tile = warp_row * WM + _wm; \
int _m_g = m_warp + _wm * 16 + lane16; \
int _lds_off = A_BASE + (BUF) * A_BUF \
+ _wm_tile * A_TILE_BYTES + lane16 * A_ROW_BYTES + k_quarter * 64; \
union { i32x8_t v; uint32_t u32[8]; uint8_t b[32]; } _abuf; \
_abuf.v = i32x8_t{0,0,0,0,0,0,0,0}; \
unsigned int _scale_byte = 127u; \
if (_m_g < M) { \
__hip_bfloat16 _av[32]; \
*(uint4*)&_av[0] = *(const uint4*)&smem[_lds_off]; \
*(uint4*)&_av[8] = *(const uint4*)&smem[_lds_off + 16]; \
*(uint4*)&_av[16] = *(const uint4*)&smem[_lds_off + 32]; \
*(uint4*)&_av[24] = *(const uint4*)&smem[_lds_off + 48]; \
float _amax = 0.0f; \
_Pragma("unroll") \
for (int _i = 0; _i < 32; _i++) { \
float _v = __bfloat162float(_av[_i]); \
_amax = fmaxf(_amax, fabsf(_v)); \
} \
_scale_byte = compute_e8m0(_amax); \
if constexpr (USE_HW_CVT) { \
union { float f; unsigned int u; } _sc; \
_sc.u = (unsigned int)_scale_byte << 23; \
float _sf = _sc.f; \
_Pragma("unroll") \
for (int _j = 0; _j < 4; _j++) { \
uint32_t _pk = 0; \
_pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
__bfloat162float(_av[8*_j+0]), __bfloat162float(_av[8*_j+1]), _sf, 0); \
_pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
__bfloat162float(_av[8*_j+2]), __bfloat162float(_av[8*_j+3]), _sf, 1); \
_pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
__bfloat162float(_av[8*_j+4]), __bfloat162float(_av[8*_j+5]), _sf, 2); \
_pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
__bfloat162float(_av[8*_j+6]), __bfloat162float(_av[8*_j+7]), _sf, 3); \
_abuf.u32[_j] = _pk; \
} \
} else { \
union { float f; unsigned int u; } _inv; \
_inv.u = ((unsigned int)(254 - (int)_scale_byte)) << 23; \
float _inv_scale = _inv.f; \
_Pragma("unroll") \
for (int _i = 0; _i < 16; _i++) { \
float _f0 = __bfloat162float(_av[2*_i]) * _inv_scale; \
float _f1 = __bfloat162float(_av[2*_i+1]) * _inv_scale; \
_abuf.b[_i] = float_to_fp4(_f0) | (float_to_fp4(_f1) << 4); \
} \
} \
} \
a_frag[_wm] = _abuf.v; \
a_scale_packed[_wm] = (int)(_scale_byte \
| (_scale_byte << 8) | (_scale_byte << 16) | (_scale_byte << 24)); \
} \
} while(0)
// ═══════════════════════════════════════════════════════════
// Prologue: issue loads for first min(PIPE_DEPTH, K_steps) ksteps
// ═══════════════════════════════════════════════════════════
#pragma unroll
for (int p = 0; p < PIPE_DEPTH; p++) {
if (p < K_steps) {
ASYNC_LOAD_B_TILES(p, p);
ASYNC_LOAD_A_TILES(p, p);
ASYNC_LOAD_B_SCALE(p, p);
}
}
async_lds_fence();
__syncthreads();
QUANTIZE_A_FROM_LDS(0);
// ═══════════════════════════════════════════════════════════
// Main loop
// ═══════════════════════════════════════════════════════════
// Pipeline: prologue loaded bufs 0..min(PIPE_DEPTH,K_steps)-1.
// At each kstep, read from cur_buf, then prefetch kstep+PIPE_DEPTH
// INTO cur_buf (pf_buf == cur_buf). Must read BEFORE prefetch!
// Fence only needed when next_buf's data came from a prefetch
// (kstep+1 >= PIPE_DEPTH), not from the already-fenced prologue.
for (int kstep = 0; kstep < K_steps; kstep++) {
int cur_buf = kstep % PIPE_DEPTH;
bool has_next = (kstep + 1 < K_steps);
// 1. Read B + B_scale from cur_buf, issue MFMAs
int ks = ks_base + kstep * 4 + k_quarter;
#pragma unroll
for (int wn = 0; wn < WN; wn++) {
int global_nt = warp_col * WN + wn;
int n_col = n_warp + wn * 16 + lane16;
union { i32x8_t v; uint8_t b[32]; } b_buf;
b_buf.v = i32x8_t{0,0,0,0,0,0,0,0};
if (n_col < N)
*(uint4*)&b_buf.b[0] = *(const uint4*)&smem[B_BASE + cur_buf * B_BUF
+ global_nt * B_TILE_BYTES + lane16 * B_ROW_BYTES + k_quarter * 16];
unsigned int b_raw = 0x7Fu;
if (n_col < N)
b_raw = GET_B_SCALE_LDS(cur_buf, n_col, ks);
int b_scale_packed = (int)(b_raw | (b_raw << 8) | (b_raw << 16) | (b_raw << 24));
#pragma unroll
for (int wm = 0; wm < WM; wm++) {
acc[wm][wn] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_frag[wm], b_buf.v,
acc[wm][wn],
4, 4,
0, a_scale_packed[wm],
0, b_scale_packed
);
}
}
// 2. Prepare for next iteration
if (has_next) {
int next_buf = (kstep + 1) % PIPE_DEPTH;
// Fence if next_buf's data came from a prefetch (not prologue)
if (kstep + 1 >= PIPE_DEPTH) {
async_lds_fence();
}
// Sync: all warps done reading cur_buf + fence propagated
__syncthreads();
// Prefetch kstep+PIPE_DEPTH into cur_buf
// (safe: all warps done reading cur_buf, cur_buf != next_buf)
int pf_step = kstep + PIPE_DEPTH;
if (pf_step < K_steps) {
ASYNC_LOAD_B_TILES(cur_buf, pf_step);
ASYNC_LOAD_A_TILES(cur_buf, pf_step);
ASYNC_LOAD_B_SCALE(cur_buf, pf_step);
}
// Quant A from next_buf (ready: prologue-fenced or just-fenced)
QUANTIZE_A_FROM_LDS(next_buf);
} else {
__syncthreads();
}
}
#undef ASYNC_LOAD_B_TILES
#undef ASYNC_LOAD_A_TILES
#undef ASYNC_LOAD_B_SCALE
#undef GET_B_SCALE_LDS
#undef QUANTIZE_A_FROM_LDS
// Store output
#pragma unroll
for (int wm = 0; wm < WM; wm++) {
#pragma unroll
for (int wn = 0; wn < WN; wn++) {
int n_col = n_warp + wn * 16 + col_out;
if (n_col < N) {
for (int i = 0; i < 4; i++) {
int m_g = m_warp + wm * 16 + row_base_out + i;
if (m_g < M) {
if (SPLIT_K) {
atomicAdd(&C_f32[m_g * N + n_col], acc[wm][wn][i]);
} else {
C_bf16[m_g * N + n_col] = __float2bfloat16(acc[wm][wn][i]);
}
}
}
}
}
}
#endif
}
__global__ void f32_to_bf16_kernel(const float* __restrict__ in,
__hip_bfloat16* __restrict__ out, int n) {
int i = blockIdx.x * 256 + threadIdx.x;
if (i < n) out[i] = __float2bfloat16(in[i]);
}
// ═══════════════════════════════════════════════════════════════
// Host dispatcher — tile configs × pipe depths
// ═══════════════════════════════════════════════════════════════
void fp4_gemm_fused(
torch::Tensor A_bf16,
int64_t B_q_ptr, int64_t B_scale_ptr,
torch::Tensor C_bf16, torch::Tensor C_f32,
int M, int N, int K,
int wm, int wn, int gm, int gn, int k_split, int pipe_depth
) {
const __hip_bfloat16* a = (const __hip_bfloat16*)A_bf16.data_ptr();
const uint8_t* b = (const uint8_t*)B_q_ptr;
const uint8_t* bs = (const uint8_t*)B_scale_ptr;
__hip_bfloat16* c_bf16 = (__hip_bfloat16*)C_bf16.data_ptr();
float* c_f32 = k_split > 1 ? (float*)C_f32.data_ptr() : nullptr;
int Ks = K >> 5;
int chunk_Ks = Ks / k_split;
int cfg = wm * 1000 + wn * 100 + gm * 10 + gn;
// Encode cfg + pipe_depth
int key = cfg * 10 + pipe_depth;
#define LAUNCH_KERNEL(WM, WN, GM, GN, SPLITK, HWCVT, PD) do { \
constexpr int BM = (WM)*(GM)*16, BN = (WN)*(GN)*16; \
constexpr int BNT = (WN)*(GN); \
constexpr int BMT = (WM)*(GM); \
constexpr int SMEM = (PD) * BNT * 16 * 64 \
+ (PD) * BMT * 16 * 256 \
+ (PD) * 1024; \
dim3 block((GM)*(GN)*64); \
if (SPLITK) { \
dim3 grid((N + BN-1) / BN, (M + BM-1) / BM, k_split); \
fp4_gemm_fused_kernel<WM, WN, GM, GN, true, HWCVT, PD><<<grid, block, SMEM>>>( \
a, b, bs, nullptr, c_f32, M, N, K, chunk_Ks); \
} else { \
dim3 grid((N + BN-1) / BN, (M + BM-1) / BM); \
fp4_gemm_fused_kernel<WM, WN, GM, GN, false, HWCVT, PD><<<grid, block, SMEM>>>( \
a, b, bs, c_bf16, nullptr, M, N, K, chunk_Ks); \
} \
} while(0)
// Only 3 configs actually used by _select_config (reduces JIT compile time)
switch (key) {
case 11114: LAUNCH_KERNEL(1, 1, 1, 1, (k_split>1), true, 4); break; // K<=512
case 11144: LAUNCH_KERNEL(1, 1, 1, 4, (k_split>1), true, 4); break; // K>512, M<=16
case 12112: LAUNCH_KERNEL(1, 2, 1, 1, (k_split>1), true, 2); break; // K>512, M>=32
}
#undef LAUNCH_KERNEL
if (k_split > 1) {
int total = M * N;
f32_to_bf16_kernel<<<(total + 255) / 256, 256>>>(c_f32, c_bf16, total);
}
}
"""
CPP_SRC = """
void fp4_gemm_fused(
torch::Tensor A_bf16,
int64_t B_q_ptr, int64_t B_scale_ptr,
torch::Tensor C_bf16, torch::Tensor C_f32,
int M, int N, int K,
int wm, int wn, int gm, int gn, int k_split, int pipe_depth
);
"""
_module = load_inline(
name="fp4_gemm_v31",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=["fp4_gemm_fused"],
verbose=True,
extra_cuda_cflags=[
"--offload-arch=gfx950", "-std=c++20", "-O3",
"-ffast-math",
"-ffp-contract=fast",
"-munsafe-fp-atomics",
],
)
def _select_config(M, N, K):
"""Hardcoded best (wm, wn, gm, gn, pipe_depth) from V29/V30 auto-tune."""
if K <= 512:
# Small K: (1,1,1,1) pipe=4 eliminates all fence stalls
wm, wn, gm, gn = (1, 1, 1, 1)
pipe_depth = 4
elif M <= 16:
# Small M, large K: (1,1,1,4) with k_split for parallelism
wm, wn, gm, gn = (1, 1, 1, 4)
pipe_depth = 4
else:
# Large shapes: (1,2,1,1) A-reuse
wm, wn, gm, gn = (1, 2, 1, 1)
pipe_depth = 2 # pipe=4 hurts on 256×3072×1536 (SMEM pressure)
block_m = wm * gm * 16
block_n = wn * gn * 16
n_m_blocks = (M + block_m - 1) // block_m
n_n_blocks = (N + block_n - 1) // block_n
base_blocks = n_m_blocks * n_n_blocks
Ks = K >> 5
k_split = 1
if K >= 1536 and base_blocks < 128:
target_blocks = 256
k_split = max(1, target_blocks // base_blocks)
MIN_CHUNK_KS = 8
k_split = min(k_split, max(1, Ks // MIN_CHUNK_KS))
while k_split > 1:
if Ks % k_split == 0 and (Ks // k_split) % 4 == 0:
break
k_split -= 1
return wm, wn, gm, gn, k_split, pipe_depth
def _run_kernel(A, b_q_ptr, b_scale_ptr, M, N, K, wm, wn, gm, gn, k_split, pipe_depth):
C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
if k_split > 1:
C_f32 = torch.zeros((M, N), dtype=torch.float32, device=A.device)
else:
C_f32 = torch.empty(1, dtype=torch.float32, device=A.device)
_module.fp4_gemm_fused(
A, b_q_ptr, b_scale_ptr,
C, C_f32, M, N, K, wm, wn, gm, gn, k_split, pipe_depth,
)
return C
# Minimal warmup — just JIT compile, no auto-tune overhead
def _warmup():
from aiter.ops.triton.quant import dynamic_mxfp4_quant
A = torch.randn((4, 512), dtype=torch.bfloat16, device="cuda")
B = torch.randn((64, 512), dtype=torch.bfloat16, device="cuda")
B_fp4, B_scale = dynamic_mxfp4_quant(B)
for _ in range(3):
_run_kernel(A, B_fp4.data_ptr(), B_scale.data_ptr(),
4, 64, 512, 1, 1, 1, 1, 1, 2)
torch.cuda.synchronize()
_warmup()
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
M, K = A.shape
N = B_q.shape[0]
wm, wn, gm, gn, k_split, pipe_depth = _select_config(M, N, K)
return _run_kernel(A, B_q.data_ptr(), B_scale_sh.data_ptr(),
M, N, K, wm, wn, gm, gn, k_split, pipe_depth)
scrolls · 542 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