submission 607214
bigboiaadulla18 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 798 lines, June 9 Researcher Reciprocity License v1.0.
submission_with_fusion_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-607214?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:e054c8e406e16d8c24ead904066557b345c47fe63161f9400a546ac23c93c228
license declaredunknown
license concludedunknown
authorsbigboiaadulla18
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Fused BF16→MXFP4 quant + FP4 GEMM kernel.shared-memory
__shared__ uint8_t A_lds[4 * A_LDS_SLOT];split-k
typename OutType, bool IS_SPLITK, bool kCheckOOB, bool B_IN_LDS = false>Kernel source
submission_with_fusion_v3.py798 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Fused BF16→MXFP4 quant + FP4 GEMM kernel.
A (bf16) is quantized to MXFP4 on-the-fly inside the GEMM kernel.
B is pre-quantized and pre-shuffled. Scales in e8m0-shuffled layout.
Uses __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4 (gfx950).
"""
import os
import sys
import aiter
import torch
from aiter import dtypes, QuantType
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant # #975-patched kernel
from aiter.utility.fp4_utils import e8m0_shuffle
from torch.utils.cpp_extension import load_inline
# K must be divisible by 64 (scale group 32 and fp4 pack 2)
SCALE_GROUP_SIZE = 32
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
cuda_src = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
#include <type_traits>
// ---- vector types ----
typedef int32_t v8i32 __attribute__((ext_vector_type(8)));
typedef float v4f32 __attribute__((ext_vector_type(4)));
typedef float v16f32 __attribute__((ext_vector_type(16)));
typedef int32_t i32x4 __attribute__((ext_vector_type(4)));
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
// ---- buffer_load_lds intrinsic (direct global → LDS) ----
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
i32x4 rsrc, as3_uint32_ptr lds_ptr, int size,
int voffset, int soffset, int offset, int aux)
__asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource {
uint64_t ptr;
uint32_t range;
uint32_t config;
};
__device__ inline i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
return *reinterpret_cast<const i32x4*>(&rsrc);
}
// ---- shared constants ----
static constexpr int WAVE_SIZE = 64;
// ---- accumulator type trait ----
template <int MFMA_SIZE> struct AccType;
template <> struct AccType<16> { using type = v4f32; };
template <> struct AccType<32> { using type = v16f32; };
// ---- e8m0-shuffled scale offset ----
__device__ __forceinline__
int shuffled_scale_offset(int row, int col, int sn_pad) {
int m_block = row >> 5;
int m_half = (row >> 4) & 1;
int m_in = row & 15;
int s_block = col >> 3;
int s_half = (col >> 2) & 1;
int s_in = col & 3;
return m_block * (sn_pad << 5)
+ s_block * 256
+ s_in * 64
+ m_in * 4
+ s_half * 2
+ m_half;
}
// ---- vmcnt helper ----
// GFX940+ s_waitcnt encoding: vmcnt is 6-bit (0-63), split across bits [3:0] and [15:14]
// expcnt is 3-bit [6:4], lgkmcnt is 6-bit [13:8] on GFX940+ (set to max 0x3F = don't wait)
template <int N>
__device__ __forceinline__ void wait_vmcnt() {
static_assert(N >= 0 && N <= 63, "vmcnt out of range");
constexpr int vmcnt_lo = N & 0xF;
constexpr int vmcnt_hi = (N >> 4) & 0x3;
constexpr int encoding = (vmcnt_hi << 14) | (0x3F << 8) | (0x7 << 4) | vmcnt_lo;
__builtin_amdgcn_s_waitcnt(encoding);
}
// Per-sub-tile vmcnt dispatch: wait_vmcnt<BASE + (COUNT-1-sub) * STEP>
// BASE, STEP, COUNT are compile-time; sub is runtime but loop must be #pragma unroll
template <int BASE, int STEP, int COUNT>
__device__ __forceinline__ void wait_vmcnt_sub(int sub) {
static_assert(COUNT >= 1 && COUNT <= 4, "K_MFMAS must be 1..4");
if constexpr (COUNT == 1) {
wait_vmcnt<BASE>();
} else if constexpr (COUNT == 2) {
if (sub == 0) wait_vmcnt<BASE + STEP>();
else wait_vmcnt<BASE>();
} else if constexpr (COUNT == 3) {
if (sub == 0) wait_vmcnt<BASE + 2*STEP>();
else if (sub == 1) wait_vmcnt<BASE + STEP>();
else wait_vmcnt<BASE>();
} else {
if (sub == 0) wait_vmcnt<BASE + 3*STEP>();
else if (sub == 1) wait_vmcnt<BASE + 2*STEP>();
else if (sub == 2) wait_vmcnt<BASE + STEP>();
else wait_vmcnt<BASE>();
}
}
// ---- helpers ----
// Load B fragment (4×uint32) and B scale from global memory.
// For MFMA_SIZE=32, l can be 0..31 crossing two 16-row shuffle tiles.
template <int MFMA_SIZE>
__device__ __forceinline__
void load_b_frag(const uint8_t* __restrict__ B, const uint8_t* __restrict__ Bs,
int n_wave, int K_half, int k_byte, int k_group, int g, int l,
int sn_pad, v8i32& b_frag, int32_t& sb) {
int l_in = l;
int n_base = n_wave;
if constexpr (MFMA_SIZE == 32) {
l_in = l & 15;
n_base = n_wave + (l >> 4) * 16;
}
const int64_t b_off = (int64_t)n_base * K_half
+ ((k_byte >> 5) << 9)
+ (g << 8) + (l_in << 4);
// Single 16-byte vector load instead of 4 separate 4-byte loads
typedef uint32_t u32x4 __attribute__((ext_vector_type(4)));
const u32x4 b_vec = *reinterpret_cast<const u32x4*>(B + b_off);
b_frag = {};
b_frag[0] = b_vec[0]; b_frag[1] = b_vec[1];
b_frag[2] = b_vec[2]; b_frag[3] = b_vec[3];
sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, k_group + g, sn_pad)];
}
// Read A tile from LDS, quantize bf16 → MXFP4 (amax + scale + fp4 convert)
__device__ __forceinline__
void quantize_a_tile(const uint8_t* A_lds, uint32_t slot_off, uint32_t lds_read_base,
v8i32& a_frag, int32_t& sa) {
const uint32_t* a_pairs = reinterpret_cast<const uint32_t*>(
A_lds + slot_off + lds_read_base);
uint32_t a_data[16];
#pragma unroll
for (int i = 0; i < 16; i++)
a_data[i] = a_pairs[i];
uint32_t max_packed = 0;
#pragma unroll
for (int i = 0; i < 16; i++) {
uint32_t abs_pair = a_data[i] & 0x7FFF7FFFu;
asm volatile("v_pk_max_u16 %0, %1, %2"
: "=v"(max_packed) : "v"(max_packed), "v"(abs_pair));
}
uint32_t max_abs = max(max_packed & 0xFFFFu, max_packed >> 16);
float amax = __uint_as_float(max_abs << 16);
uint32_t amax_u = __float_as_uint(amax);
amax_u = (amax_u + 0x200000u) & 0xFF800000u;
int exp_field = (int)((amax_u >> 23) & 0xFFu);
int scale_unbiased = exp_field - 129;
scale_unbiased = max(-127, min(127, scale_unbiased));
sa = (int32_t)((uint8_t)(scale_unbiased + 127));
float hw_scale = __uint_as_float((uint32_t)sa << 23);
uint8_t fp4_bytes[16];
#pragma unroll
for (int i = 0; i < 16; i++) {
uint32_t result;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(result) : "v"(a_data[i]), "v"(hw_scale));
fp4_bytes[i] = (uint8_t)result;
}
a_frag = {};
const uint32_t* fp = reinterpret_cast<const uint32_t*>(fp4_bytes);
a_frag[0] = fp[0]; a_frag[1] = fp[1];
a_frag[2] = fp[2]; a_frag[3] = fp[3];
}
// Issue MFMA instruction (dispatches to correct intrinsic based on MFMA_SIZE)
template <int MFMA_SIZE>
__device__ __forceinline__
void do_mfma(v8i32 a_frag, v8i32 b_frag,
typename AccType<MFMA_SIZE>::type& acc, int32_t sa, int32_t sb) {
if constexpr (MFMA_SIZE == 16) {
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_frag, b_frag, acc, 4, 4, 0, sa, 0, sb);
} else {
acc = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
a_frag, b_frag, acc, 4, 4, 0, sa, 0, sb);
}
}
// Store accumulator results. OutType = __hip_bfloat16 (normal) or float (SplitK partial).
// m_tile_base = block_m + wave_m * MFMA_SIZE (base M-row for this wave's tile)
template <int MFMA_SIZE, bool kCheckOOB, typename OutType>
__device__ __forceinline__
void store_acc(const typename AccType<MFMA_SIZE>::type& acc,
OutType* __restrict__ C, int m_tile_base, int g, int l, int n_wave, int N, int M) {
if constexpr (MFMA_SIZE == 16) {
const int m_base = m_tile_base + 4 * g;
#pragma unroll
for (int i = 0; i < 4; i++) {
if (!kCheckOOB || m_base + i < M) {
if constexpr (std::is_same_v<OutType, float>)
C[(int64_t)(m_base + i) * N + n_wave + l] = acc[i];
else
C[(int64_t)(m_base + i) * N + n_wave + l] = __float2bfloat16(acc[i]);
}
}
} else {
#pragma unroll
for (int i = 0; i < 16; i++) {
const int m_row = m_tile_base + g * 4 + (i / 4) * 8 + (i % 4);
if (!kCheckOOB || m_row < M) {
if constexpr (std::is_same_v<OutType, float>)
C[(int64_t)m_row * N + n_wave + l] = acc[i];
else
C[(int64_t)m_row * N + n_wave + l] = __float2bfloat16(acc[i]);
}
}
}
}
// ---- unified templated GEMM kernel ----
//
// Template params:
// MFMA_SIZE: 16 or 32
// TILE_M_T: tile height (multiple of MFMA_SIZE)
// TILE_N_T: tile width (multiple of MFMA_SIZE)
// TILE_K_T: K elements processed per iteration (multiple of MFMA_K; enables multiple MFMAs per iter)
// OutType: __hip_bfloat16 (normal) or float (SplitK partial sums)
// IS_SPLITK: if true, each block processes a K-range subset; blockIdx.z = split index
// B_IN_LDS: if true, B data is prefetched through LDS (good for small M); if false, direct global load
template <int kM, int kN, int kK, int kNumSplits, int MFMA_SIZE, int TILE_M_T, int TILE_N_T, int TILE_K_T,
typename OutType, bool IS_SPLITK, bool kCheckOOB, bool B_IN_LDS = false>
__global__ __launch_bounds__((TILE_M_T / MFMA_SIZE) * (TILE_N_T / MFMA_SIZE) * WAVE_SIZE)
void gemm_kernel(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ Bs,
OutType* __restrict__ C,
int M, int N, int K)
{
// if kM, kN, and kK are specified, then overwrite M, N, and K. This will be known at compile time.
// Having M, N, and K known at compile time allows for better optimizations.
if constexpr (kM != 0) {
M = kM;
}
if constexpr (kN != 0) {
N = kN;
}
if constexpr (kK != 0) {
K = kK;
}
// Compile-time constants
constexpr int MFMA_K = (MFMA_SIZE == 16) ? 128 : 64;
constexpr int K_MFMAS = TILE_K_T / MFMA_K; // MFMAs per K-tile iteration
constexpr int M_TILES = TILE_M_T / MFMA_SIZE;
constexpr int N_TILES = TILE_N_T / MFMA_SIZE;
constexpr int NUM_WAVES = M_TILES * N_TILES;
constexpr int A_LDS_ROW_T = MFMA_K * 2; // 256 or 128 (per MFMA sub-tile)
constexpr int G_SHIFT = (MFMA_SIZE == 16) ? 4 : 5;
constexpr int ROW_SHIFT = (MFMA_SIZE == 16) ? 4 : 3; // for a_voffset
constexpr int COL_MASK = (MFMA_SIZE == 16) ? 15 : 7;
constexpr int ROWS_PER_WAVE = 1024 / A_LDS_ROW_T; // rows each wave loads per chunk (4 or 8)
constexpr int CHUNK_PER_WAVE = 1024; // bytes per wave per load
constexpr int TOTAL_ROWS = TILE_M_T; // rows to load
constexpr int ROWS_PER_CHUNK = NUM_WAVES * ROWS_PER_WAVE; // rows loaded per chunk by all waves
constexpr int NUM_CHUNKS = (TOTAL_ROWS + ROWS_PER_CHUNK - 1) / ROWS_PER_CHUNK; // loads per wave per sub-tile
// A sub-slot: one MFMA_K-wide sub-tile of A
constexpr int A_LDS_DATA_SUB = TILE_M_T * MFMA_K * 2; // exact A data per sub-tile
constexpr int A_LDS_LOAD_SUB = NUM_CHUNKS * NUM_WAVES * CHUNK_PER_WAVE; // load footprint per sub-tile
constexpr int A_SUB_SLOT = (A_LDS_DATA_SUB > A_LDS_LOAD_SUB) ? A_LDS_DATA_SUB : A_LDS_LOAD_SUB;
// Full A slot: K_MFMAS sub-tiles
constexpr int A_LDS_SLOT = K_MFMAS * A_SUB_SLOT;
static_assert(TILE_M_T % MFMA_SIZE == 0, "TILE_M_T must be multiple of MFMA_SIZE");
static_assert(TILE_N_T % MFMA_SIZE == 0, "TILE_N_T must be multiple of MFMA_SIZE");
static_assert(TILE_K_T % MFMA_K == 0, "TILE_K_T must be multiple of MFMA_K");
// A_LDS_SLOT >= A_LDS_LOAD is guaranteed by max() above
const int K_half = K >> 1;
const int sn_pad = ((K / 32 + 7) >> 3) << 3;
// SplitK: compute K range for this split from blockIdx.z
int k_start = 0;
int k_end = K;
if constexpr (IS_SPLITK) {
static_assert(kNumSplits != 0, "kNumSplits must be != 0 for SplitK");
int num_splits = kNumSplits;
int total_k_tiles = K / TILE_K_T;
int tiles_per_split = (total_k_tiles + num_splits - 1) / num_splits;
int my_tile_start = (int)blockIdx.z * tiles_per_split;
int my_tile_end = min(my_tile_start + tiles_per_split, total_k_tiles);
if (my_tile_start >= total_k_tiles) return;
k_start = my_tile_start * TILE_K_T;
k_end = my_tile_end * TILE_K_T;
// Advance C to this split's slice of the workspace
C = C + (int64_t)blockIdx.z * M * N;
}
const int num_k_tiles = (k_end - k_start) / TILE_K_T;
// XCD-aware block mapping: gridDim.x is padded to multiple of 8
if (blockIdx.x * TILE_N_T >= N) return;
const int block_m = blockIdx.y * TILE_M_T;
const int block_n = blockIdx.x * TILE_N_T;
const int wave_id = threadIdx.x / WAVE_SIZE;
const int lane = threadIdx.x % WAVE_SIZE;
// Wave-to-tile mapping
const int wave_m = wave_id / N_TILES;
const int wave_n = wave_id % N_TILES;
// Lane-level mapping within MFMA tile
const int g = lane >> G_SHIFT;
const int l = lane & (MFMA_SIZE - 1);
const int n_wave = block_n + wave_n * MFMA_SIZE;
typename AccType<MFMA_SIZE>::type acc = {};
// B LDS constants (only used when B_IN_LDS)
constexpr int B_SUB_SLOT = B_IN_LDS ? NUM_WAVES * 1024 : 0; // one MFMA_K sub-tile of B
constexpr int B_LDS_SLOT = K_MFMAS * B_SUB_SLOT; // full TILE_K of B
// B scale LDS constants (only used when B_IN_LDS)
// Load a full contiguous 256-byte tile via buffer_load_lds (64 lanes × 4 bytes).
// One tile covers 32 rows × 8 scale columns = 256 K-elements = one TILE_K iteration.
// Reused across all K_MFMAS sub-tiles within the iteration (no per-sub load needed).
constexpr int Bs_SCALE_SLOT = B_IN_LDS ? NUM_WAVES * 256 : 0; // one 256B tile per wave
// VMEM operation counts per TILE_K tile (per wave)
constexpr int A_TILE_VMEM = NUM_CHUNKS * K_MFMAS; // A LDS loads per tile
constexpr int B_TILE_VMEM = K_MFMAS; // B LDS loads per tile (when B_IN_LDS)
constexpr int BS_TILE_VMEM = B_IN_LDS ? 1 : 0; // one 256B tile load per iteration
constexpr int LDS_VMEM = B_IN_LDS ? (A_TILE_VMEM + B_TILE_VMEM + BS_TILE_VMEM) : A_TILE_VMEM;
// B_IN_LDS: everything through LDS (DIRECT_VMEM=0)
// !B_IN_LDS: B data+scale loaded directly, retired per-sub in inner loop
constexpr int DIRECT_VMEM = B_IN_LDS ? 0 : 2 * K_MFMAS;
constexpr int DIRECT_PER_SUB = B_IN_LDS ? 0 : 2;
// Quad-buffered LDS for A (and B + B scales when B_IN_LDS)
__shared__ uint8_t A_lds[4 * A_LDS_SLOT];
__shared__ uint8_t B_lds[B_IN_LDS ? 4 * B_LDS_SLOT : 1];
__shared__ uint8_t Bs_scale_lds[B_IN_LDS ? 4 * Bs_SCALE_SLOT : 1];
static_assert(!B_IN_LDS || TILE_K_T == 256,
"B scale tile load assumes TILE_K_T=256 (one 256B tile per iteration)");
// Each wave reads its M-tile's rows from A LDS
const uint32_t lds_read_base = (wave_m * MFMA_SIZE + l) * A_LDS_ROW_T + g * 64;
// ---- A buffer resource ----
// Clamp range to valid rows so OOB reads (when M % TILE_M != 0) return 0
const int valid_m_rows = min(TILE_M_T, M - block_m);
i32x4 a_srsrc = make_srsrc(
reinterpret_cast<const void*>(A + (int64_t)block_m * K),
(uint32_t)(valid_m_rows * K * 2));
const int a_voffset_lane = (lane >> ROW_SHIFT) * (K * 2) + (lane & COL_MASK) * 16;
// ---- B buffer resource (only for B_IN_LDS) ----
[[maybe_unused]] i32x4 b_srsrc;
[[maybe_unused]] int b_voffset = 0;
if constexpr (B_IN_LDS) {
b_srsrc = make_srsrc(
reinterpret_cast<const void*>(B),
(uint32_t)((uint64_t)N * K_half > 0xFFFFFFFFu ? 0xFFFFFFFFu : N * K_half));
int l_in = l;
int n_base = n_wave;
if constexpr (MFMA_SIZE == 32) {
l_in = l & 15;
n_base = n_wave + (l >> 4) * 16;
}
b_voffset = n_base * K_half + (g << 8) + (l_in << 4);
}
// ---- B scale buffer resource (only for B_IN_LDS) ----
// Loads a contiguous 256-byte tile via buffer_load_lds (64 lanes × 4 bytes).
// One tile covers 32 rows × 8 scale columns = 256 K-elements.
// bs_soffset_row = m_block*(sn_pad<<5), constant per wave (row-dependent base).
// bs_lds_lane_base = g*64 + m_in*4 + m_half (per-lane constant for reading).
[[maybe_unused]] i32x4 bs_srsrc;
[[maybe_unused]] int bs_soffset_row = 0;
[[maybe_unused]] int bs_lds_lane_base = 0;
if constexpr (B_IN_LDS) {
uint32_t bs_range = (uint32_t)(((N + 31u) >> 5) * ((uint32_t)sn_pad << 5));
bs_srsrc = make_srsrc(reinterpret_cast<const void*>(Bs), bs_range);
int row = n_wave + l;
int m_block = row >> 5;
int m_half = (row >> 4) & 1;
int m_in = row & 15;
bs_soffset_row = m_block * (sn_pad << 5);
bs_lds_lane_base = g * 64 + m_in * 4 + m_half;
}
// Helper: load full A tile (K_MFMAS sub-tiles) into LDS
// k_elem = starting K element offset for this tile
auto load_a_tile = [&](uint32_t slot_off, int k_elem) __attribute__((always_inline)) {
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
int soff = (k_elem + sub * MFMA_K) * 2;
#pragma unroll
for (int c = 0; c < NUM_CHUNKS; c++) {
int base_row = c * ROWS_PER_CHUNK + wave_id * ROWS_PER_WAVE;
int a_voffset = base_row * (K * 2) + a_voffset_lane;
as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
reinterpret_cast<uintptr_t>(A_lds) + slot_off
+ sub * A_SUB_SLOT
+ (uint32_t)(c * ROWS_PER_CHUNK * A_LDS_ROW_T + wave_id * CHUNK_PER_WAVE));
llvm_amdgcn_raw_buffer_load_lds(a_srsrc, lds_dst, 16, a_voffset, soff, 0, 0);
}
}
};
// B-through-LDS helpers (only when B_IN_LDS)
// k_elem = starting K element offset for this tile
auto load_b_tile = [&](uint32_t slot_off, int k_elem) __attribute__((always_inline)) {
if constexpr (B_IN_LDS) {
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
int k_sub = k_elem + sub * MFMA_K;
int b_soff = (k_sub >> 6) << 9;
as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
reinterpret_cast<uintptr_t>(B_lds) + slot_off
+ sub * B_SUB_SLOT + (uint32_t)wave_id * 1024u);
llvm_amdgcn_raw_buffer_load_lds(b_srsrc, lds_dst, 16, b_voffset, b_soff, 0, 0);
}
}
};
// Load B scales into LDS: single contiguous 256-byte tile per iteration.
// 64 lanes × 4 bytes = 256 bytes. Covers 32 rows × 8 scale columns = 256 K-elements.
// soffset = bs_soffset_row + s_block*256 where s_block = k_elem/256.
auto load_bs_tile = [&](uint32_t slot_off, int k_elem) __attribute__((always_inline)) {
if constexpr (B_IN_LDS) {
int bs_soff = bs_soffset_row + (k_elem / 256) * 256;
as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
reinterpret_cast<uintptr_t>(Bs_scale_lds) + slot_off
+ (uint32_t)wave_id * 256u);
llvm_amdgcn_raw_buffer_load_lds(bs_srsrc, lds_dst, 4, lane * 4, bs_soff, 0, 0);
}
};
// Read one MFMA_K sub-tile of B from LDS into registers
auto read_b_frag_lds = [&](uint32_t sub_off, v8i32& b_frag) __attribute__((always_inline)) {
if constexpr (B_IN_LDS) {
const uint32_t* b_ptr = reinterpret_cast<const uint32_t*>(
B_lds + sub_off + wave_id * 1024u + (uint32_t)lane * 16u);
b_frag = {};
b_frag[0] = b_ptr[0]; b_frag[1] = b_ptr[1];
b_frag[2] = b_ptr[2]; b_frag[3] = b_ptr[3];
}
};
// Read B scale byte from the contiguous 256-byte tile in LDS.
// Layout within 256B tile: offset = g*64 + m_in*4 + s_half*2 + m_half
// bs_lds_lane_base = g*64 + m_in*4 + m_half (constant per lane).
// sub selects s_half: for MFMA_K=128, sub 0 → s_half=0, sub 1 → s_half=1.
auto read_bs_scale_lds = [&](uint32_t slot_off, int sub) -> int32_t __attribute__((always_inline)) {
if constexpr (B_IN_LDS) {
return (int32_t)Bs_scale_lds[slot_off + wave_id * 256 + bs_lds_lane_base + sub * 2];
}
return 0;
};
// === Prologue: load first 2 TILE_K tiles (A + B data + B scales if B_IN_LDS) ===
load_a_tile(0, k_start);
if constexpr (B_IN_LDS) {
load_b_tile(0, k_start);
load_bs_tile(0, k_start);
}
if (num_k_tiles >= 2) {
load_a_tile(A_LDS_SLOT, k_start + TILE_K_T);
if constexpr (B_IN_LDS) {
load_b_tile(B_LDS_SLOT, k_start + TILE_K_T);
load_bs_tile(Bs_SCALE_SLOT, k_start + TILE_K_T);
}
}
// =====================================================================
// Standard K-loop with 4-stage quad-buffering, 2-ahead prefetch
// =====================================================================
//
// Each iteration processes TILE_K_T elements = K_MFMAS sub-MFMAs.
// Quad-buffered: slots 0,1,2,3. Prefetch is 2 tiles ahead.
//
// VMEM accounting (per iteration):
// LDS_VMEM = A_TILE_VMEM [+ B_TILE_VMEM + BS_TILE_VMEM] via buffer_load_lds
// B_IN_LDS: DIRECT_VMEM=0 (everything through LDS), no inner vmcnt waits
// !B_IN_LDS: DIRECT_VMEM=2*K_MFMAS (B data+scale), per-sub vmcnt waits
// === Main loop: t=0..num_k_tiles-3 (prefetch t+2) ===
for (int t = 0; t < num_k_tiles - 2; t++) {
const int k = k_start + t * TILE_K_T;
const uint32_t cur_a_slot = (t % 4) * A_LDS_SLOT;
[[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 4) * B_LDS_SLOT : 0;
[[maybe_unused]] const uint32_t cur_bs_slot = B_IN_LDS ? (t % 4) * Bs_SCALE_SLOT : 0;
// !B_IN_LDS: load B data+scale from global into VGPRs
[[maybe_unused]] v8i32 b_frags[K_MFMAS];
[[maybe_unused]] int32_t sbs[K_MFMAS];
if constexpr (!B_IN_LDS) {
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
int k_sub = k + sub * MFMA_K;
load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k_sub / 2, k_sub / 32, g, l, sn_pad, b_frags[sub], sbs[sub]);
}
}
// Prefetch tile t+2 (A + B data + B scales if B_IN_LDS)
{
const uint32_t pf_a_slot = ((t + 2) % 4) * A_LDS_SLOT;
const int k_pf = k_start + (t + 2) * TILE_K_T;
load_a_tile(pf_a_slot, k_pf);
if constexpr (B_IN_LDS) {
const uint32_t pf_b_slot = ((t + 2) % 4) * B_LDS_SLOT;
const uint32_t pf_bs_slot = ((t + 2) % 4) * Bs_SCALE_SLOT;
load_b_tile(pf_b_slot, k_pf);
load_bs_tile(pf_bs_slot, k_pf);
}
}
// Wait for current tile's LDS data (loaded 2 iters ago)
wait_vmcnt<2 * LDS_VMEM + DIRECT_VMEM>();
__syncthreads();
// Inner loop: K_MFMAS quantize+MFMA operations
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot + sub * A_SUB_SLOT, lds_read_base, a_frag, sa);
if constexpr (B_IN_LDS) {
v8i32 b_frag;
read_b_frag_lds(cur_b_slot + sub * B_SUB_SLOT, b_frag);
int32_t sb = read_bs_scale_lds(cur_bs_slot, sub);
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
} else {
wait_vmcnt_sub<LDS_VMEM, DIRECT_PER_SUB, K_MFMAS>(sub);
do_mfma<MFMA_SIZE>(a_frag, b_frags[sub], acc, sa, sbs[sub]);
}
}
}
// === Second-to-last iteration (peeled): t=num_k_tiles-2, no prefetch ===
if (num_k_tiles >= 2) {
const int t = num_k_tiles - 2;
const int k = k_start + t * TILE_K_T;
const uint32_t cur_a_slot = (t % 4) * A_LDS_SLOT;
[[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 4) * B_LDS_SLOT : 0;
[[maybe_unused]] const uint32_t cur_bs_slot = B_IN_LDS ? (t % 4) * Bs_SCALE_SLOT : 0;
[[maybe_unused]] v8i32 b_frags[K_MFMAS];
[[maybe_unused]] int32_t sbs[K_MFMAS];
if constexpr (!B_IN_LDS) {
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
int k_sub = k + sub * MFMA_K;
load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k_sub / 2, k_sub / 32, g, l, sn_pad, b_frags[sub], sbs[sub]);
}
}
wait_vmcnt<LDS_VMEM + DIRECT_VMEM>();
__syncthreads();
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot + sub * A_SUB_SLOT, lds_read_base, a_frag, sa);
if constexpr (B_IN_LDS) {
v8i32 b_frag;
read_b_frag_lds(cur_b_slot + sub * B_SUB_SLOT, b_frag);
int32_t sb = read_bs_scale_lds(cur_bs_slot, sub);
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
} else {
wait_vmcnt_sub<0, DIRECT_PER_SUB, K_MFMAS>(sub);
do_mfma<MFMA_SIZE>(a_frag, b_frags[sub], acc, sa, sbs[sub]);
}
}
}
// === Last iteration (peeled): t=num_k_tiles-1, no outstanding prefetch ===
{
const int t = num_k_tiles - 1;
const int k = k_start + t * TILE_K_T;
const uint32_t cur_a_slot = (t % 4) * A_LDS_SLOT;
[[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 4) * B_LDS_SLOT : 0;
[[maybe_unused]] const uint32_t cur_bs_slot = B_IN_LDS ? (t % 4) * Bs_SCALE_SLOT : 0;
[[maybe_unused]] v8i32 b_frags[K_MFMAS];
[[maybe_unused]] int32_t sbs[K_MFMAS];
if constexpr (!B_IN_LDS) {
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
int k_sub = k + sub * MFMA_K;
load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k_sub / 2, k_sub / 32, g, l, sn_pad, b_frags[sub], sbs[sub]);
}
}
wait_vmcnt<0>();
__syncthreads();
#pragma unroll
for (int sub = 0; sub < K_MFMAS; sub++) {
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot + sub * A_SUB_SLOT, lds_read_base, a_frag, sa);
if constexpr (B_IN_LDS) {
v8i32 b_frag;
read_b_frag_lds(cur_b_slot + sub * B_SUB_SLOT, b_frag);
int32_t sb = read_bs_scale_lds(cur_bs_slot, sub);
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
} else {
do_mfma<MFMA_SIZE>(a_frag, b_frags[sub], acc, sa, sbs[sub]);
}
}
}
store_acc<MFMA_SIZE, kCheckOOB, OutType>(acc, C, block_m + wave_m * MFMA_SIZE, g, l, n_wave, N, M);
}
// ---- SplitK reduction kernel ----
template <int kNumSplits>
__global__ void splitk_reduce(
const float* __restrict__ workspace, // [num_splits, M, N]
__hip_bfloat16* __restrict__ C, // [M, N]
int MN)
{
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= MN) return;
float sum = 0.0f;
#pragma unroll
for (int s = 0; s < kNumSplits; s++) {
sum += workspace[(int64_t)s * MN + idx];
}
C[idx] = __float2bfloat16(sum);
}
// ---- torch wrapper ----
#include <torch/extension.h>
template <int kM, int kN, int kK, int kNumSplits, int MFMA_SIZE, int TM, int TN, int TK, bool kCheckOOB, bool B_LDS>
void launch_gemm(const __hip_bfloat16* A, const uint8_t* B, const uint8_t* Bs,
__hip_bfloat16* C, int m, int n, int k) {
constexpr int NWAVES = (TM / MFMA_SIZE) * (TN / MFMA_SIZE);
// constexpr int NUM_XCDS = 8;
// int n_tiles = n / TN;
// int n_tiles_padded = ((n_tiles + NUM_XCDS - 1) / NUM_XCDS) * NUM_XCDS;
// dim3 grid(n_tiles_padded, (m + TM - 1) / TM, 1);
dim3 grid((n + TN - 1) / TN, (m + TM - 1) / TM, 1);
dim3 block(NWAVES * WAVE_SIZE);
gemm_kernel<kM, kN, kK, kNumSplits, MFMA_SIZE, TM, TN, TK, __hip_bfloat16, false, kCheckOOB, B_LDS><<<grid, block>>>(A, B, Bs, C, m, n, k);
}
template <int kM, int kN, int kK, int kNumSplits, int MFMA_SIZE, int TM, int TN, int TK, bool kCheckOOB, bool B_LDS>
void launch_gemm_splitk(const __hip_bfloat16* A, const uint8_t* B, const uint8_t* Bs,
float* workspace, __hip_bfloat16* C, int m, int n, int k) {
constexpr int NWAVES = (TM / MFMA_SIZE) * (TN / MFMA_SIZE);
// constexpr int NUM_XCDS = 8;
// int n_tiles = n / TN;
// int n_tiles_padded = ((n_tiles + NUM_XCDS - 1) / NUM_XCDS) * NUM_XCDS;
// dim3 grid(n_tiles_padded, (m + TM - 1) / TM, kNumSplits);
dim3 grid((n + TN - 1) / TN, (m + TM - 1) / TM, kNumSplits);
dim3 block(NWAVES * WAVE_SIZE);
gemm_kernel<kM, kN, kK, kNumSplits, MFMA_SIZE, TM, TN, TK, float, true, kCheckOOB, B_LDS><<<grid, block>>>(A, B, Bs, workspace, m, n, k);
// Reduce partial sums across splits
constexpr int REDUCE_THREADS = 256;
int mn = m * n;
int reduce_blocks = (mn + REDUCE_THREADS - 1) / REDUCE_THREADS;
splitk_reduce<kNumSplits><<<reduce_blocks, REDUCE_THREADS>>>(workspace, C, mn);
}
at::Tensor mxfp4_gemm(
at::Tensor A, at::Tensor B_fp4, at::Tensor B_scale,
at::Tensor C, int m, int n, int k)
{
TORCH_CHECK(k % 128 == 0, "k must be divisible by 128");
TORCH_CHECK(n % 16 == 0, "n must be divisible by 16");
const auto* A_ptr = reinterpret_cast<const __hip_bfloat16*>(A.data_ptr());
const auto* B_ptr = reinterpret_cast<const uint8_t*>(B_fp4.data_ptr());
const auto* Bs_ptr = reinterpret_cast<const uint8_t*>(B_scale.data_ptr());
auto* C_ptr = reinterpret_cast<__hip_bfloat16*>(C.data_ptr());
if (m == 4 && n == 2880 && k == 512) {
launch_gemm<4, 2880, 512, 0, 16, 16, 32, 256, true, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 16 && n == 2112 && k == 7168) {
static constexpr int num_splits = 7;
auto workspace = at::empty({num_splits, m, n}, A.options().dtype(at::kFloat));
auto* ws_ptr = reinterpret_cast<float*>(workspace.data_ptr());
launch_gemm_splitk<16, 2112, 7168, num_splits, 16, 16, 64, 256, false, true>(A_ptr, B_ptr, Bs_ptr, ws_ptr, C_ptr, m, n, k);
}
else if (m == 32 && n == 4096 && k == 512) {
launch_gemm<32, 4096, 512, 0, 16, 16, 32, 256, false, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 32 && n == 2880 && k == 512) {
launch_gemm<32, 2880, 512, 0, 16, 16, 32, 256, false, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 64 && n == 7168 && k == 2048) {
launch_gemm<64, 7168, 2048, 0, 16, 16, 64, 256, false, false>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 256 && n == 3072 && k == 1536) {
launch_gemm<256, 3072, 1536, 0, 16, 16, 64, 256, false, false>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else {
launch_gemm<0, 0, 0, 0, 16, 16, 64, 256, true, false>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
return C;
}
"""
cpp_src = r"""
at::Tensor mxfp4_gemm(at::Tensor A, at::Tensor B_fp4, at::Tensor B_scale,
at::Tensor C, int m, int n, int k);
"""
def generate_input(m: int, n: int, k: int, seed: int): # -> input_t:
"""
Generate random bf16 inputs A [m, k], B [n, k] and quantized MXFP4 B, shuffled B and B_scale.
Returns:
Tuple of (A, B), both bf16 on cuda.
"""
assert k % 64 == 0, "k must be divisible by 64 (scale group 32 and fp4 pack 2)"
gen = torch.Generator(device="cuda")
gen.manual_seed(seed)
A = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
B = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
B_q, B_scale_sh = _quant_mxfp4(B, shuffle=True)
# shuffle B(weight) to (16,16) tile coalesced
B_shuffle = shuffle_weight(B_q, layout=(16, 16))
return (A, B, B_q, B_shuffle, B_scale_sh)
def _quant_mxfp4(x, shuffle=True):
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)
if sys.stdout is None:
sys.stdout = open("/dev/stdout", "w")
if sys.stderr is None:
sys.stderr = open("/dev/stderr", "w")
module = load_inline(
name="A_bf16_B_mxfp4_C_bf16_gemm",
cpp_sources=[cpp_src],
cuda_sources=[cuda_src],
functions=["mxfp4_gemm"],
verbose=True,
extra_cuda_cflags=[
"-O3",
"--offload-arch=gfx950",
"-std=c++20",
"-ffp-contract=fast",
"-lhip_hcc",
],
)
def custom_kernel(data):
"""
data is generated by generate_input()
"""
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n, _ = B.shape
C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
module.mxfp4_gemm(
A,
B_shuffle.view(torch.uint8),
B_scale_sh.view(torch.uint8),
C,
m,
n,
k,
)
return C
scrolls · 798 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