submission 597059
Ashwin Adulla · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 750 lines, June 9 Researcher Reciprocity License v1.0.
submission_with_fusion.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-597059?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:0340615111c479089c97614e37adeecc4b2b464bf26664214b1b4a57507d9fac
license declaredunknown
license concludedunknown
authorsAshwin Adulla
imported2026-08-26
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[3 * A_LDS_SLOT];split-k
template <int MFMA_SIZE, int TILE_M_T, int TILE_N_T, typename OutType, bool IS_SPLITK,tile-m = 16
TILE_M = 16Kernel source
submission_with_fusion.py750 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 ----
template <int N>
__device__ __forceinline__ void wait_vmcnt() {
static_assert(N >= 0 && N <= 15, "vmcnt out of range");
if constexpr (N == 0) asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
else if constexpr (N == 1) asm volatile("s_waitcnt vmcnt(1)" ::: "memory");
else if constexpr (N == 2) asm volatile("s_waitcnt vmcnt(2)" ::: "memory");
else if constexpr (N == 3) asm volatile("s_waitcnt vmcnt(3)" ::: "memory");
else if constexpr (N == 4) asm volatile("s_waitcnt vmcnt(4)" ::: "memory");
else if constexpr (N == 5) asm volatile("s_waitcnt vmcnt(5)" ::: "memory");
else if constexpr (N == 6) asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
else if constexpr (N == 7) asm volatile("s_waitcnt vmcnt(7)" ::: "memory");
else if constexpr (N == 8) asm volatile("s_waitcnt vmcnt(8)" ::: "memory");
else if constexpr (N == 9) asm volatile("s_waitcnt vmcnt(9)" ::: "memory");
else if constexpr (N == 10) asm volatile("s_waitcnt vmcnt(10)" ::: "memory");
else if constexpr (N == 11) asm volatile("s_waitcnt vmcnt(11)" ::: "memory");
else if constexpr (N == 12) asm volatile("s_waitcnt vmcnt(12)" ::: "memory");
else if constexpr (N == 13) asm volatile("s_waitcnt vmcnt(13)" ::: "memory");
else if constexpr (N == 14) asm volatile("s_waitcnt vmcnt(14)" ::: "memory");
else if constexpr (N == 15) asm volatile("s_waitcnt vmcnt(15)" ::: "memory");
}
// ---- 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, 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) {
if constexpr (MFMA_SIZE == 16) {
const int m_base = m_tile_base + 4 * g;
#pragma unroll
for (int i = 0; i < 4; i++) {
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 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)
// 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
// PINGPONG: if true, use 8-wave ping-pong scheduling (requires WAVES_PER_WG == 2, B_IN_LDS == true)
template <int MFMA_SIZE, int TILE_M_T, int TILE_N_T, typename OutType, bool IS_SPLITK,
bool B_IN_LDS = false, bool PINGPONG = 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)
{
// Compile-time constants
constexpr int MFMA_K = (MFMA_SIZE == 16) ? 128 : 64;
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 WAVES_PER_WG = (NUM_WAVES <= 4) ? 1 : 2;
constexpr int A_LDS_ROW_T = MFMA_K * 2; // 256 or 128
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
constexpr int A_LDS_DATA = TILE_M_T * MFMA_K * 2; // exact A data per slot
constexpr int A_LDS_LOAD = NUM_CHUNKS * NUM_WAVES * CHUNK_PER_WAVE; // load footprint
constexpr int A_LDS_SLOT = (A_LDS_DATA > A_LDS_LOAD) ? A_LDS_DATA : A_LDS_LOAD;
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");
// A_LDS_SLOT >= A_LDS_LOAD is guaranteed by max() above
static_assert(!PINGPONG || WAVES_PER_WG == 2,
"PINGPONG requires 8 waves (WAVES_PER_WG == 2)");
static_assert(!PINGPONG || B_IN_LDS,
"PINGPONG requires B_IN_LDS");
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) {
int total_k_tiles = K / MFMA_K;
int tiles_per_split = (total_k_tiles + (int)gridDim.z - 1) / (int)gridDim.z;
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 * MFMA_K;
k_end = my_tile_end * MFMA_K;
// 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) / MFMA_K;
// 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;
[[maybe_unused]] const int wavegroup = wave_id / WAVES_PER_WG;
[[maybe_unused]] const int wave_in_wg = wave_id % WAVES_PER_WG;
// 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_LDS_SLOT = B_IN_LDS ? NUM_WAVES * 1024 : 0;
// VMEM operation counts per K tile (per wave)
constexpr int A_VMEM_OPS = NUM_CHUNKS; // buffer_load_lds calls for A
// Triple-buffered LDS for A (and B when B_IN_LDS)
__shared__ uint8_t A_lds[3 * A_LDS_SLOT];
__shared__ uint8_t B_lds[B_IN_LDS ? 3 * B_LDS_SLOT : 1];
// 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 ----
i32x4 a_srsrc = make_srsrc(
reinterpret_cast<const void*>(A + (int64_t)block_m * K),
(uint32_t)(TILE_M_T * 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);
}
// Helper: load A tile into LDS (issues NUM_CHUNKS buffer_load_lds)
auto load_a_tile = [&](uint32_t slot_off, int soff) __attribute__((always_inline)) {
#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
+ (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)
auto load_b_tile = [&](uint32_t slot_off, int b_soff) __attribute__((always_inline)) {
if constexpr (B_IN_LDS) {
as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
reinterpret_cast<uintptr_t>(B_lds) + slot_off + (uint32_t)wave_id * 1024u);
llvm_amdgcn_raw_buffer_load_lds(b_srsrc, lds_dst, 16, b_voffset, b_soff, 0, 0);
}
};
auto read_b_frag_lds = [&](uint32_t slot_off, v8i32& b_frag) __attribute__((always_inline)) {
if constexpr (B_IN_LDS) {
const uint32_t* b_ptr = reinterpret_cast<const uint32_t*>(
B_lds + slot_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];
}
};
// === Prologue: load first A tile (and B if B_IN_LDS) into LDS slot 0 ===
load_a_tile(0, k_start * 2);
if constexpr (B_IN_LDS) {
load_b_tile(0, (k_start >> 6) << 9);
}
if constexpr (!PINGPONG) {
// =====================================================================
// Standard K-loop (no ping-pong)
// =====================================================================
//
// B_IN_LDS=true: B prefetched to LDS one K-tile ahead, read from LDS
// Issue: B_scale(1), A_next(A_VMEM_OPS), B_next(1)
// Outstanding after prev: prev_A + prev_B + new = 2*A_VMEM_OPS + 3
// Keep = A_VMEM_OPS + 2, second wait = A_VMEM_OPS + 1
//
// B_IN_LDS=false: B loaded directly from global in current iteration
// Issue: B_data(1) + B_scale(1), A_next(A_VMEM_OPS)
// Outstanding after prev: prev_A + new = 2*A_VMEM_OPS + 2
// Keep = A_VMEM_OPS + 2, second wait = A_VMEM_OPS
// === Main loop: all iterations except the last ===
for (int t = 0; t < num_k_tiles - 1; t++) {
const int k = k_start + t * MFMA_K;
const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
[[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 3) * B_LDS_SLOT : 0;
v8i32 b_frag; int32_t sb;
if constexpr (B_IN_LDS) {
// B scale from global (1 VMEM)
sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
} else {
// B data + B scale from global (2 VMEM)
load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k / 2, k / 32, g, l, sn_pad, b_frag, sb);
}
// Load A_next (and B_next if B_IN_LDS) to LDS
{
const uint32_t next_a_slot = ((t + 1) % 3) * A_LDS_SLOT;
const int k_next = k + MFMA_K;
load_a_tile(next_a_slot, k_next * 2);
if constexpr (B_IN_LDS) {
const uint32_t next_b_slot = ((t + 1) % 3) * B_LDS_SLOT;
load_b_tile(next_b_slot, (k_next >> 6) << 9);
}
}
// Wait for prev A (and prev B if B_IN_LDS) LDS loads to complete.
// Both paths: keep A_VMEM_OPS + 2 newest ops in flight
wait_vmcnt<A_VMEM_OPS + 2>();
__syncthreads();
// Read A from LDS → quantize
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);
if constexpr (B_IN_LDS) {
// Read B from LDS
read_b_frag_lds(cur_b_slot, b_frag);
// Wait for B scale. Keep A_next + B_next = A_VMEM_OPS + 1
wait_vmcnt<A_VMEM_OPS + 1>();
} else {
// Wait for B data + B scale. Keep A_next = A_VMEM_OPS
wait_vmcnt<A_VMEM_OPS>();
}
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
}
// === Last iteration (peeled): no next A/B load ===
{
const int t = num_k_tiles - 1;
const int k = k_start + t * MFMA_K;
const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
[[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 3) * B_LDS_SLOT : 0;
v8i32 b_frag; int32_t sb;
if constexpr (B_IN_LDS) {
sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
} else {
load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k / 2, k / 32, g, l, sn_pad, b_frag, sb);
}
wait_vmcnt<0>();
__syncthreads();
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);
if constexpr (B_IN_LDS) {
read_b_frag_lds(cur_b_slot, b_frag);
}
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
}
} else {
// =====================================================================
// Ping-Pong K-loop: 8 waves, 2 per SIMD, staggered compute/memory
// =====================================================================
// wave_in_wg==0: compute-first (quantize+MFMA, then load A_next+B_next)
// wave_in_wg==1: memory-first (load B_scale+A_next+B_next, then quantize+MFMA)
// s_setprio(1) boosts compute wave, s_setprio(0) yields to partner
// sched_barrier prevents compiler from reordering across phase boundaries
for (int t = 0; t < num_k_tiles - 1; t++) {
const int k = k_start + t * MFMA_K;
const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
const uint32_t cur_b_slot = (t % 3) * B_LDS_SLOT;
const uint32_t next_a_slot = ((t + 1) % 3) * A_LDS_SLOT;
const uint32_t next_b_slot = ((t + 1) % 3) * B_LDS_SLOT;
const int k_next = k + MFMA_K;
// Wait for all outstanding VMEM from previous iteration
wait_vmcnt<0>();
__syncthreads();
if (wave_in_wg == 0) {
// --- Compute-first path ---
asm volatile("s_setprio 1" ::: "memory");
// Quantize A from LDS (no VMEM)
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);
// Read B from LDS (no VMEM)
v8i32 b_frag;
read_b_frag_lds(cur_b_slot, b_frag);
// Load B scale (1 VMEM, only outstanding op)
int32_t sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
wait_vmcnt<0>();
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
asm volatile("s_setprio 0" ::: "memory");
asm volatile("s_nop 0" ::: "memory");
// Memory phase: load next A and B into LDS
load_a_tile(next_a_slot, k_next * 2);
load_b_tile(next_b_slot, (k_next >> 6) << 9);
} else {
// --- Memory-first path ---
asm volatile("s_setprio 0" ::: "memory");
// Load B scale first (becomes oldest VMEM op)
int32_t sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
// Load next A and B into LDS
load_a_tile(next_a_slot, k_next * 2);
load_b_tile(next_b_slot, (k_next >> 6) << 9);
asm volatile("s_nop 0" ::: "memory");
asm volatile("s_setprio 1" ::: "memory");
// Quantize A from LDS (no VMEM)
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);
// Read B from LDS (no VMEM)
v8i32 b_frag;
read_b_frag_lds(cur_b_slot, b_frag);
// Wait for B scale (oldest). Keep A_next + B_next in flight.
// Outstanding: B_scale(1) + A_next(A_VMEM_OPS) + B_next(1) = A_VMEM_OPS + 2
// Retire 1 oldest (B_scale), keep A_VMEM_OPS + 1
wait_vmcnt<A_VMEM_OPS + 1>();
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
asm volatile("s_setprio 0" ::: "memory");
asm volatile("s_nop 0" ::: "memory");
}
}
// === Last iteration (peeled): both waves do compute, no next loads ===
{
const int t = num_k_tiles - 1;
const int k = k_start + t * MFMA_K;
const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
const uint32_t cur_b_slot = (t % 3) * B_LDS_SLOT;
wait_vmcnt<0>();
__syncthreads();
asm volatile("s_setprio 1" ::: "memory");
v8i32 a_frag; int32_t sa;
quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);
v8i32 b_frag;
read_b_frag_lds(cur_b_slot, b_frag);
int32_t sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
wait_vmcnt<0>();
do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
asm volatile("s_setprio 0" ::: "memory");
}
} // end PINGPONG
store_acc<MFMA_SIZE, OutType>(acc, C, block_m + wave_m * MFMA_SIZE, g, l, n_wave, N);
}
// ---- SplitK reduction kernel ----
__global__ void splitk_reduce(
const float* __restrict__ workspace, // [num_splits, M, N]
__hip_bfloat16* __restrict__ C, // [M, N]
int MN, int num_splits)
{
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= MN) return;
float sum = 0.0f;
for (int s = 0; s < num_splits; s++) {
sum += workspace[(int64_t)s * MN + idx];
}
C[idx] = __float2bfloat16(sum);
}
// ---- torch wrapper ----
#include <torch/extension.h>
template <int MFMA_SIZE, int TM, int TN, bool B_LDS = false, bool PP = false>
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);
dim3 block(NWAVES * WAVE_SIZE);
gemm_kernel<MFMA_SIZE, TM, TN, __hip_bfloat16, false, B_LDS, PP><<<grid, block>>>(A, B, Bs, C, m, n, k);
}
template <int MFMA_SIZE, int TM, int TN, bool B_LDS = false, bool PP = false>
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, int num_splits) {
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, num_splits);
dim3 block(NWAVES * WAVE_SIZE);
gemm_kernel<MFMA_SIZE, TM, TN, float, true, B_LDS, PP><<<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<<<reduce_blocks, REDUCE_THREADS>>>(workspace, C, mn, num_splits);
}
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(m % 16 == 0, "m must be divisible by 16");
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<16, 16, 16>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 16 && n == 2112 && k == 7168) {
int num_splits = 8;
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, 16, 64>(A_ptr, B_ptr, Bs_ptr, ws_ptr, C_ptr, m, n, k, num_splits);
}
else if (m == 32 && n == 4096 && k == 512) {
launch_gemm<16, 16, 32, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 32 && n == 2880 && k == 512) {
launch_gemm<16, 16, 32, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 64 && n == 7168 && k == 2048) {
launch_gemm<16, 16, 64>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else if (m == 256 && n == 3072 && k == 1536) {
// Best so far: <16, 16, 64>
launch_gemm<16, 16, 64>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
}
else {
launch_gemm<16, 16, 64>(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
TILE_M = 16
m_padded = ((m + TILE_M - 1) // TILE_M) * TILE_M
if m_padded != m:
A = torch.nn.functional.pad(A, (0, 0, 0, m_padded - m))
C = torch.empty((m_padded, n), dtype=torch.bfloat16, device="cuda")
out_gemm = module.mxfp4_gemm(
A,
B_shuffle.view(torch.uint8),
B_scale_sh.view(torch.uint8),
C,
m_padded,
n,
k,
)
return out_gemm[:m, :]
scrolls · 750 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