submission 754770
Brian Sun · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1679 lines, June 9 Researcher Reciprocity License v1.0.
homestretch_splitk.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754770?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:df79d7d495574ad6dac76da0ee9bf8dbdeaad3f2f338454a6f1dd52469468fff
license declaredunknown
license concludedunknown
authorsBrian Sun
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm -> bf16 C.shared-memory
__shared__ float lds_reduce[M_MAX][BLOCK_N][N_WF]; // [16][64][8] = 32KBsplit-k
template <int K_TOTAL, int M_MAX, int N_TOTAL, int SPLIT_K>tile-k = 128
constexpr int TILE_K = 128;tile-m = 16
constexpr int TILE_M = 16;tile-n = 16
constexpr int TILE_N = 16;Kernel source
homestretch_splitk.py1679 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm -> bf16 C.
"""
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
from aiter import QuantType,dtypes
import aiter
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
HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <pybind11/pybind11.h>
// ============================================================
// Types
// ============================================================
typedef int v4i32 __attribute__((ext_vector_type(4)));
typedef int v8i32 __attribute__((ext_vector_type(8)));
typedef float v4f32 __attribute__((ext_vector_type(4)));
typedef float v16f32 __attribute__((ext_vector_type(16)));
typedef unsigned int ext_u32x2 __attribute__((ext_vector_type(2)));
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
v4i32 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__ __forceinline__ v4i32 make_srsrc(const void* ptr, uint32_t range_bytes) {
buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
return *reinterpret_cast<const v4i32*>(&rsrc);
}
__device__ __forceinline__
v8i32 to_v8(v4i32 v) {
// fp4 MFMA only reads lower 4 elements; upper 4 are ignored.
// Union avoids 4 wasted v_mov_b32 zero-inits per call.
union { v4i32 lo; v8i32 full; } u;
u.lo = v;
return u.full;
}
// ============================================================
// Constants
// ============================================================
constexpr int TILE_M = 16;
constexpr int TILE_N = 16;
constexpr int TILE_K = 128;
constexpr int GROUP_SZ = 32;
constexpr int WF_SIZE = 64;
// ============================================================
// Device helpers
// ============================================================
__device__ __forceinline__
uint8_t compute_e8m0_scale(const __hip_bfloat16* vals) {
// Find max |val| using bf16 bit representation (no float conversion)
uint16_t mx_bits = 0;
#pragma unroll
for (int i = 0; i < GROUP_SZ; i++) {
uint16_t b = *reinterpret_cast<const uint16_t*>(&vals[i]) & 0x7FFF;
mx_bits = max(mx_bits, b);
}
if (mx_bits == 0) return 0;
// Promote to f32 bit space for rounding: bf16 is top 16 bits of f32
uint32_t bits = (uint32_t)mx_bits << 16;
bits = (bits + 0x200000u) & 0xFF800000u;
int e8m0 = (int)(bits >> 23) - 2;
return (uint8_t)max(0, min(254, e8m0));
}
__device__ __forceinline__
void quantize_group(const __hip_bfloat16* vals, v4i32& out, uint32_t& scale_out) {
uint8_t scale = compute_e8m0_scale(vals);
scale_out = scale;
float scale_f = __uint_as_float((uint32_t)scale << 23);
#define PACK_WORD(w) do { \
unsigned int packed = 0; \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+0], vals[(w)*8+1]}, scale_f, 0); \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+2], vals[(w)*8+3]}, scale_f, 1); \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+4], vals[(w)*8+5]}, scale_f, 2); \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+6], vals[(w)*8+7]}, scale_f, 3); \
out[w] = (int)packed; \
} while(0)
PACK_WORD(0);
PACK_WORD(1);
PACK_WORD(2);
PACK_WORD(3);
#undef PACK_WORD
}
// ============================================================
// All quant-done flags in a single cache line (128 bytes)
// ============================================================
struct alignas(128) FlagLine { unsigned char flags[128]; }; // 7 used, rest padding
// ============================================================
// Shape 2 split-K GEMM: 8 wf split K within CTA, each does N_SUBS=4 MFMAs.
// 16×64 output tile. A quantized once per wf, reused across 4 B tiles.
// LDS reduction across 8 K-split wf. Then workspace for cross-CTA split-K.
// Grid: (N_BLOCKS * SPLIT_K, 1, 1), Block: 8 * 64 = 512
// ============================================================
template <int K_TOTAL, int M_MAX, int N_TOTAL, int SPLIT_K>
__global__ __attribute__((amdgpu_flat_work_group_size(8 * WF_SIZE, 8 * WF_SIZE)))
void mxfp4_gemm_splitk_gemm(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_scale,
int B_scale_stride,
float* __restrict__ workspace // [N_BLOCKS * SPLIT_K * M_MAX * BLOCK_N]
) {
constexpr int K = K_TOTAL;
constexpr int N = N_TOTAL;
constexpr int N_WF = 8;
constexpr int N_SUBS = 4; // B tiles per wf
constexpr int BLOCK_N = TILE_N * N_SUBS; // 64
constexpr int N_BLOCKS = (N + BLOCK_N - 1) / BLOCK_N; // 33
constexpr int K_PER_SPLIT = K / SPLIT_K; // 1024
// 8 wf split K_PER_SPLIT: each wf handles 128 K = 1 MFMA per B tile
const int tid = threadIdx.x;
const int wf_id = tid / WF_SIZE; // 0..7 — which K slice
const int tid_in_wf = tid % WF_SIZE;
const int lane = tid_in_wf % 16;
const int k_group = tid_in_wf / 16;
const int n_block = blockIdx.x / SPLIT_K;
const int split_idx = blockIdx.x % SPLIT_K;
const int a_row = lane;
// A: each wf handles a different 128-element K slice
const int k_start = split_idx * K_PER_SPLIT + wf_id * TILE_K;
const __hip_bfloat16* a_ptr = A + a_row * K + k_start + k_group * GROUP_SZ;
// B addressing base for this wf's K position
const int c_cur = (k_start / 32) + k_group;
const int bs_k_part = (c_cur >> 3) * 256 + (c_cur & 3) * 64 + ((c_cur >> 2) & 1) * 2;
v4i32 b_buf[2]; uint32_t sb_buf[2];
auto load_b_ns = [&](int buf, int ns) __attribute__((always_inline)) {
int n_tile = n_block * N_SUBS + ns;
if (n_tile >= N / TILE_N) return;
int b_tile_n_base = n_tile * (K / 32);
int b_n_sub = n_tile & 1;
int bs_n_base = (n_tile >> 1) * (B_scale_stride * 32);
int bs_lane_off = lane * 4 + b_n_sub;
int b_off = (b_tile_n_base + c_cur) * 256 + lane * 16;
b_buf[buf] = *reinterpret_cast<const v4i32*>(B_shuf + b_off);
sb_buf[buf] = B_scale[bs_n_base + bs_k_part + bs_lane_off];
};
// Issue A load + first 2 B loads BEFORE quantize — B arrives during quantize
v4i32 a_raw[4];
const v4i32* sv = reinterpret_cast<const v4i32*>(a_ptr);
#pragma unroll
for (int v = 0; v < 4; v++) a_raw[v] = sv[v];
load_b_ns(0, 0);
if (N_SUBS > 1) load_b_ns(1, 1);
// Quantize A (~103 cycles — B loads arrive during this window)
v4i32 a_buf; uint32_t sa_buf;
__hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw);
quantize_group(vals, a_buf, sa_buf);
v4f32 acc[N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = {0.0f, 0.0f, 0.0f, 0.0f};
int cur = 0;
// Double-buffered N_SUBS loop: MFMA → prefetch next B → repeat
#pragma unroll
for (int ns = 0; ns < N_SUBS - 1; ns++) {
int nxt = cur ^ 1;
int n_tile = n_block * N_SUBS + ns;
if (n_tile < N / TILE_N) {
acc[ns] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf), to_v8(b_buf[cur]), acc[ns],
4, 4, 0, sa_buf, 0, sb_buf[cur]);
}
if (ns + 2 < N_SUBS) load_b_ns(cur, ns + 2);
cur = nxt;
}
// Last N_SUB
{
int n_tile = n_block * N_SUBS + (N_SUBS - 1);
if (n_tile < N / TILE_N) {
acc[N_SUBS - 1] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf), to_v8(b_buf[cur]), acc[N_SUBS - 1],
4, 4, 0, sa_buf, 0, sb_buf[cur]);
}
}
// ---- LDS reduction across 8 K-split wavefronts ----
// Layout [row][col][wf]: 8 wf values contiguous for fast reduction reads
__shared__ float lds_reduce[M_MAX][BLOCK_N][N_WF]; // [16][64][8] = 32KB
{
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
int col = ns * TILE_N + lane;
#pragma unroll
for (int i = 0; i < 4; i++) {
int row = k_group * 4 + i;
lds_reduce[row][col][wf_id] = acc[ns][i];
}
}
}
__syncthreads();
// All 512 threads reduce + store: 1024 elements / 512 threads = 2 per thread
// Reduction reads 8 contiguous floats per element (2 × v4f32)
{
int ws_off = (n_block * SPLIT_K + split_idx) * M_MAX * BLOCK_N;
#pragma unroll
for (int e = tid; e < M_MAX * BLOCK_N; e += N_WF * WF_SIZE) {
int row = e / BLOCK_N;
int col = e % BLOCK_N;
int abs_col = n_block * BLOCK_N + col;
if (abs_col < N) {
const float* src = &lds_reduce[row][col][0];
v4f32 v0 = *reinterpret_cast<const v4f32*>(src);
v4f32 v1 = *reinterpret_cast<const v4f32*>(src + 4);
float sum = v0[0] + v0[1] + v0[2] + v0[3]
+ v1[0] + v1[1] + v1[2] + v1[3];
workspace[ws_off + e] = sum;
}
}
}
}
// ============================================================
// Shape 2 split-K reduce: sum SPLIT_K partials for 16×64 tiles → bf16.
// Grid: (N_BLOCKS, 1, 1), Block: 256 (= 16 * 64 / 4 elements per thread)
// ============================================================
template <int M_MAX, int N_TOTAL, int SPLIT_K>
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void mxfp4_splitk_reduce(
const float* __restrict__ workspace,
__hip_bfloat16* __restrict__ C_bf16
) {
constexpr int N = N_TOTAL;
constexpr int BLOCK_N = 64;
constexpr int N_BLOCKS = (N + BLOCK_N - 1) / BLOCK_N;
constexpr int TILE_FLOATS = M_MAX * BLOCK_N; // 1024
const int n_block = blockIdx.x;
const int tid = threadIdx.x; // 0..511
const int base = tid * 4; // each thread handles 4 consecutive floats
if (base >= TILE_FLOATS) return;
int row = base / BLOCK_N;
int col = base % BLOCK_N;
int abs_col = n_block * BLOCK_N + col;
float s0 = 0, s1 = 0, s2 = 0, s3 = 0;
#pragma unroll
for (int s = 0; s < SPLIT_K; s++) {
const float* src = &workspace[(n_block * SPLIT_K + s) * TILE_FLOATS + base];
v4f32 v = *reinterpret_cast<const v4f32*>(src);
s0 += v[0]; s1 += v[1]; s2 += v[2]; s3 += v[3];
}
// Bounds check for N not divisible by 128 (2112 % 128 = 64)
if (abs_col + 3 < N) {
int pk0, pk1;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk0) : "v"(s0), "v"(s1));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk1) : "v"(s2), "v"(s3));
int gbl = row * N + abs_col;
*reinterpret_cast<int*>(&C_bf16[gbl]) = pk0;
*reinterpret_cast<int*>(&C_bf16[gbl + 2]) = pk1;
} else {
// Scalar fallback for edge
if (abs_col < N) C_bf16[row * N + abs_col] = __float2bfloat16(s0);
if (abs_col + 1 < N) C_bf16[row * N + abs_col + 1] = __float2bfloat16(s1);
if (abs_col + 2 < N) C_bf16[row * N + abs_col + 2] = __float2bfloat16(s2);
if (abs_col + 3 < N) C_bf16[row * N + abs_col + 3] = __float2bfloat16(s3);
}
}
// ============================================================
// Shape 2 phased kernel: Phase 1 = all blocks quantize A,
// Phase 2 = blocks 0..131 run GEMM with prequantized A.
// Waterfall sync: block 0→1→2→...→6→master, all blocks wait on master.
// Grid: (256, 1, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_MAX, int N_WF, int N_GEMM_BLOCKS, int N_GEMM_BLOCKS_EXTRA = 24, int FLAG_VAL = 1, uint64_t FLAG_PAT8 = 0x0101010101010101ULL>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_shape2_phased(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_scale,
__hip_bfloat16* __restrict__ C,
int N, int B_scale_stride,
v4i32* __restrict__ A_q_out,
uint32_t* __restrict__ A_scale_out,
FlagLine* __restrict__ flag_line
) {
constexpr int K = K_TOTAL;
constexpr int GROUPS_PER_ROW = K / GROUP_SZ;
constexpr int TOTAL_GROUPS = M_MAX * GROUPS_PER_ROW;
constexpr int N_QUANT_BLOCKS = N_GEMM_BLOCKS_EXTRA; // must be multiple of 8
const int tid = threadIdx.x;
const int threads_per_block = N_WF * WF_SIZE;
// ======== Blocks 0..6: Quantize + signal + return ========
if (blockIdx.x < N_QUANT_BLOCKS) {
int idx = blockIdx.x * threads_per_block + tid;
if (idx < TOTAL_GROUPS) {
int row = idx / GROUPS_PER_ROW;
int k_grp = idx % GROUPS_PER_ROW;
const __hip_bfloat16* src = A + row * K + k_grp * GROUP_SZ;
__hip_bfloat16 vals[GROUP_SZ];
const v4i32* src_vec = reinterpret_cast<const v4i32*>(src);
v4i32* dst_vec = reinterpret_cast<v4i32*>(vals);
#pragma unroll
for (int v = 0; v < 4; v++) dst_vec[v] = src_vec[v];
v4i32 out; uint32_t scale;
quantize_group(vals, out, scale);
int out_idx = k_grp * M_MAX + row;
*reinterpret_cast<volatile v4i32*>(&A_q_out[out_idx]) = out;
*reinterpret_cast<volatile uint32_t*>(&A_scale_out[out_idx]) = scale;
}
__builtin_amdgcn_s_waitcnt(0);
__syncthreads();
if (tid == 0) {
*reinterpret_cast<volatile unsigned char*>(&flag_line->flags[blockIdx.x]) = FLAG_VAL;
}
return;
}
// ======== Blocks 7..138: 4 bf16 iters, prefetch flag, check, branch ========
const int n_tile = blockIdx.x - N_QUANT_BLOCKS;
constexpr int K_ITERS_PER_WF = K / N_WF / TILE_K; // 7
constexpr int K_WF_TILES = N_WF * (TILE_K / 32);
constexpr int WF_K_OFF = TILE_K / 32;
constexpr int A_ADVANCE = N_WF * TILE_K;
constexpr int B_BYTE_INC = K_WF_TILES * 256;
const int wf_id = tid / WF_SIZE;
const int tid_in_wf = tid % WF_SIZE;
const int lane = tid_in_wf % 16;
const int k_group = tid_in_wf / 16;
const int m_tile = 0;
const int m_base = 0;
const int a_row = lane % M_MAX;
const int k_rotation = n_tile % K_ITERS_PER_WF;
const int k_start = k_rotation * A_ADVANCE;
const __hip_bfloat16* a_ptr = A + a_row * K + k_start + wf_id * TILE_K + k_group * GROUP_SZ;
const int b_tile_n_base = n_tile * (K / 32);
const int b_n_sub = n_tile & 1;
const int bs_n_base = (n_tile >> 1) * (B_scale_stride * 32);
const int bs_lane_off = lane * 4 + b_n_sub;
int c_cur = (k_start / 32) + wf_id * WF_K_OFF + k_group;
int b_byte_off = (b_tile_n_base + c_cur) * 256 + lane * 16;
v4f32 acc = {0.0f, 0.0f, 0.0f, 0.0f};
v4i32 a_buf[2];
uint32_t sa_buf[2];
v4i32 b_buf[2];
uint32_t sb_buf[2];
v4i32 a_raw[2][4];
auto issue_a_loads = [&](int buf) __attribute__((always_inline)) {
const v4i32* src_vec = reinterpret_cast<const v4i32*>(a_ptr);
#pragma unroll
for (int v = 0; v < 4; v++) a_raw[buf][v] = src_vec[v];
};
auto quantize_a = [&](int buf) __attribute__((always_inline)) {
__hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw[buf]);
quantize_group(vals, a_buf[buf], sa_buf[buf]);
if constexpr (M_MAX < TILE_M) { if (lane >= M_MAX) sa_buf[buf] = 0; }
};
auto load_a_pq = [&](int buf) __attribute__((always_inline)) {
int pq_idx = c_cur * M_MAX + a_row;
a_buf[buf] = A_q_out[pq_idx];
sa_buf[buf] = A_scale_out[pq_idx];
if constexpr (M_MAX < TILE_M) { if (lane >= M_MAX) sa_buf[buf] = 0; }
};
auto load_b = [&](int buf) __attribute__((always_inline)) {
b_buf[buf] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off);
int k_inn = c_cur & 3;
int k_sub_v = (c_cur >> 2) & 1;
int k_blk = c_cur >> 3;
sb_buf[buf] = B_scale[bs_n_base + k_blk * 256 + k_inn * 64 + bs_lane_off + k_sub_v * 2];
};
auto advance_k_bf16 = [&]() __attribute__((always_inline)) {
a_ptr += A_ADVANCE;
b_byte_off += B_BYTE_INC;
c_cur += K_WF_TILES;
if (a_ptr >= A + a_row * K + K) {
a_ptr -= K;
b_byte_off -= (K / 32) * 256;
c_cur -= K / 32;
}
};
auto advance_k_pq = [&]() __attribute__((always_inline)) {
b_byte_off += B_BYTE_INC;
c_cur += K_WF_TILES;
if (c_cur >= K / 32) {
b_byte_off -= (K / 32) * 256;
c_cur -= K / 32;
}
};
constexpr int K_THRESHOLD = 3;
// ---- Iterations 0 to K_THRESHOLD-1: bf16 ----
issue_a_loads(0); load_b(0);
int cur = 0;
#pragma unroll
for (int ki = 0; ki < K_THRESHOLD; ki++) {
int nxt = cur ^ 1;
quantize_a(cur);
advance_k_bf16();
issue_a_loads(nxt); load_b(nxt);
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
cur = nxt;
}
// ---- Load flags as 3 × int64 (24 bytes), quantize (masks latency), compare ----
constexpr int N_FLAG_WORDS = N_QUANT_BLOCKS / 8; // must divide evenly
const uint64_t* f64 = reinterpret_cast<const uint64_t*>(flag_line->flags);
uint64_t fv[N_FLAG_WORDS];
#pragma unroll
for (int i = 0; i < N_FLAG_WORDS; i++) fv[i] = f64[i];
quantize_a(cur); // ~130 cycles, hides flag load latency
bool use_pq = true;
#pragma unroll
for (int i = 0; i < N_FLAG_WORDS; i++) {
if (fv[i] != FLAG_PAT8) { use_pq = false; break; }
}
if (use_pq) {
// ---- Iteration K_THRESHOLD+1: cur already quantized, prefetch prequant ----
{
int nxt = cur ^ 1;
advance_k_pq();
load_a_pq(nxt); load_b(nxt);
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
cur = nxt;
}
// ---- Remaining prequant iterations ----
#pragma unroll
for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
int nxt = cur ^ 1;
advance_k_pq();
load_a_pq(nxt); load_b(nxt);
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
cur = nxt;
}
// Epilogue
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
} else {
// ---- Remaining bf16 iterations (cur already quantized) ----
// First: advance + load + MFMA (quantize already done above)
{
int nxt = cur ^ 1;
advance_k_bf16();
issue_a_loads(nxt); load_b(nxt);
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
cur = nxt;
}
// Remaining loop
#pragma unroll
for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
int nxt = cur ^ 1;
quantize_a(cur);
advance_k_bf16();
issue_a_loads(nxt); load_b(nxt);
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
cur = nxt;
}
// Epilogue
quantize_a(cur);
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
}
// LDS reduction across N_WF wavefronts
__shared__ float lds_reduce[N_WF][4][WF_SIZE];
#pragma unroll
for (int i = 0; i < 4; i++)
lds_reduce[wf_id][i][tid_in_wf] = acc[i];
__syncthreads();
if (wf_id == 0) {
float sum[4];
#pragma unroll
for (int i = 0; i < 4; i++) {
sum[i] = 0.0f;
#pragma unroll
for (int w = 0; w < N_WF; w++)
sum[i] += lds_reduce[w][i][tid_in_wf];
}
int out_col = n_tile * TILE_N + lane;
#pragma unroll
for (int i = 0; i < 4; i++) {
int out_row = m_base + k_group * 4 + i;
if ((M_MAX >= TILE_M) || (out_row < m_base + M_MAX))
C[out_row * N + out_col] = __float2bfloat16(sum[i]);
}
}
}
// ============================================================
// Small-M kernel: N_WF wavefronts split K, LDS reduction.
// Each WG handles a single 16x16 output tile.
// Grid: (N/16, num_m_tiles, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_MAX, int N_WF>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_smallm(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_scale,
__hip_bfloat16* __restrict__ C,
int N, int B_scale_stride
) {
constexpr int K = K_TOTAL;
constexpr int K_ITERS_PER_WF = K / N_WF / TILE_K;
const int tid = threadIdx.x;
const int wf_id = tid / WF_SIZE;
const int tid_in_wf = tid % WF_SIZE;
const int lane = tid_in_wf % 16;
const int k_group = tid_in_wf / 16;
const int n_tile = blockIdx.x;
const int m_tile = blockIdx.y;
const int m_base = m_tile * TILE_M;
constexpr int WF_K_OFF = TILE_K / 32;
const int b_tile_n_base = n_tile * (K / 32);
const int b_n_sub = n_tile & 1;
const int bs_n_base = (n_tile >> 1) * (B_scale_stride * 32);
const int bs_lane_off = lane * 4 + b_n_sub;
int c_cur = wf_id * WF_K_OFF + k_group;
int b_byte_off = (b_tile_n_base + c_cur) * 256 + lane * 16;
v4f32 acc = {0.0f, 0.0f, 0.0f, 0.0f};
v4i32 a_buf;
uint32_t sa_buf;
v4i32 b_buf;
uint32_t sb_buf;
// ---- Load ALL of A into LDS via buffer_load_lds ----
// A = 4 rows × 512 bf16 = 4KB. 256 threads × 16 bytes = 4KB. One load per thread.
// Each wf loads one row (wf_id → row), perfectly coalesced.
__shared__ __hip_bfloat16 lds_a[M_MAX][K]; // [4][512] = 4KB
{
v4i32 a_srsrc = make_srsrc(A, N * K * 2);
int a_soff = (m_base + wf_id) * K * 2;
int a_voff = tid_in_wf * 16;
llvm_amdgcn_raw_buffer_load_lds(a_srsrc,
reinterpret_cast<as3_uint32_ptr>(
reinterpret_cast<uintptr_t>(&lds_a[wf_id][0])),
16, a_voff, a_soff, 0, 0);
// A load: 1 vmcnt in flight
}
// Issue B loads while A is in flight (2 more vmcnt)
b_buf = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off);
int k_inn = c_cur & 3;
int k_sub_v = (c_cur >> 2) & 1;
int k_blk = c_cur >> 3;
sb_buf = B_scale[bs_n_base + k_blk * 256 + k_inn * 64 + bs_lane_off + k_sub_v * 2];
// Wait for A buffer_load_lds (vmcnt(2) = let 2 B loads still be outstanding)
asm volatile("s_waitcnt vmcnt(2)");
__syncthreads();
// Read A from LDS → registers, then quantize
// Each thread reads its 32 bf16 from (row, K slice) in the full A tile
{
int my_row = lane % M_MAX;
int my_k = wf_id * TILE_K + k_group * GROUP_SZ;
const __hip_bfloat16* lds_src = &lds_a[my_row][my_k];
v4i32 a_raw[4];
const v4i32* src_vec = reinterpret_cast<const v4i32*>(lds_src);
#pragma unroll
for (int v = 0; v < 4; v++) a_raw[v] = src_vec[v];
__hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw);
quantize_group(vals, a_buf, sa_buf);
sa_buf = (lane < M_MAX) ? sa_buf : 0u;
}
// Single MFMA
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_buf), to_v8(b_buf), acc,
4, 4, 0, sa_buf, 0, sb_buf);
// ---- LDS reduction across N_WF wavefronts ----
if constexpr (M_MAX < TILE_M) {
// Only k_group=0 holds valid rows. Compact LDS: [N_WF-1][M_MAX][16].
// wf0 keeps its acc in registers — no LDS write or read for itself.
__shared__ float lds_reduce[N_WF - 1][M_MAX][TILE_N];
if (k_group == 0 && wf_id > 0) {
#pragma unroll
for (int i = 0; i < M_MAX; i++)
lds_reduce[wf_id - 1][i][lane] = acc[i];
}
__syncthreads();
if (wf_id == 0 && k_group == 0) {
int out_col = n_tile * TILE_N + lane;
#pragma unroll
for (int i = 0; i < M_MAX; i++) {
float sum = acc[i]; // start with own accumulator
#pragma unroll
for (int w = 0; w < N_WF - 1; w++)
sum += lds_reduce[w][i][lane];
C[(m_base + i) * N + out_col] = __float2bfloat16(sum);
}
}
} else {
// wf0 skips its own LDS write/read — uses acc directly.
__shared__ float lds_reduce[N_WF - 1][4][WF_SIZE];
if (wf_id > 0) {
#pragma unroll
for (int i = 0; i < 4; i++)
lds_reduce[wf_id - 1][i][tid_in_wf] = acc[i];
}
__syncthreads();
if (wf_id == 0) {
float sum[4];
#pragma unroll
for (int i = 0; i < 4; i++) {
sum[i] = acc[i]; // start with own accumulator
#pragma unroll
for (int w = 0; w < N_WF - 1; w++)
sum[i] += lds_reduce[w][i][tid_in_wf];
}
int out_col = n_tile * TILE_N + lane;
#pragma unroll
for (int i = 0; i < 4; i++) {
int out_row = m_base + k_group * 4 + i;
C[out_row * N + out_col] = __float2bfloat16(sum[i]);
}
}
}
}
// 32×32×64 MFMA kernel: 32×(32*N_SUBS) output tile, N_WF wf split K.
// 2 K-chunks per step → 4 MFMAs per step (2 kc × N_SUBS).
// A quantized once per kc, reused across N_SUBS MFMAs.
// Pipelined: issue loads → MFMAs → quantize arrived A.
// LDS reduction across N_WF wf at end. Direct bf16 store.
// Grid: (N/(32*N_SUBS), M/32, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_TOTAL, int N_TOTAL, int N_WF, int N_SUBS_T = 2>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_32x32(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_scale,
__hip_bfloat16* __restrict__ C,
int B_scale_stride
) {
constexpr int K = K_TOTAL;
constexpr int M = M_TOTAL;
constexpr int N = N_TOTAL;
constexpr int TILE_32 = 32;
constexpr int TILE_K_32 = 64;
constexpr int GROUP_SZ_32 = 32;
constexpr int N_SUBS = N_SUBS_T;
constexpr int K_CHUNKS = 2;
constexpr int K_STEP = K_CHUNKS * TILE_K_32; // 128 fp4 per step
constexpr int K_ITERS_PER_WF = K / N_WF / K_STEP;
const int tid = threadIdx.x;
const int wf_id = tid / WF_SIZE;
const int tid_in_wf = tid % WF_SIZE;
const int lane = tid_in_wf % 32;
const int k_chunk = tid_in_wf / 32;
const int n_block = blockIdx.x;
const int m_tile = blockIdx.y;
const int n_base = n_block * (TILE_32 * N_SUBS);
const int m_base = m_tile * TILE_32;
const int my_a_row = m_base + lane;
const __hip_bfloat16* A_row_ptr = A + my_a_row * K;
const int b_lane_in_tile = lane % 16;
const int b_n_sub_tile_off = lane / 16;
int b_tile_n_idx[N_SUBS];
int bs_n_blk_base[N_SUBS];
int bs_n_sub_val[N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
int n_off = n_base + ns * TILE_32 + lane;
b_tile_n_idx[ns] = (n_base + ns * TILE_32) / 16 + b_n_sub_tile_off;
bs_n_blk_base[ns] = (n_off / 32) * (B_scale_stride * 32);
bs_n_sub_val[ns] = (n_off / 16) % 2;
}
const int bs_n_inn = (lane % 16) * 4;
v16f32 acc[N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = {0,0,0,0, 0,0,0,0, 0,0,0,0, 0,0,0,0};
// A: base pointer for this thread's K range (wf spaced by K_STEP, not TILE_K_32)
const __hip_bfloat16* a_load_ptr = A_row_ptr + wf_id * K_STEP + k_chunk * GROUP_SZ_32;
// B: c_base uses K_STEP spacing
const int c_base = (wf_id * K_STEP) / 32 + k_chunk;
int b_byte_off[N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
int tile_idx_0 = b_tile_n_idx[ns] * (K / 32) + c_base;
b_byte_off[ns] = tile_idx_0 * 256 + b_lane_in_tile * 16;
}
// Double-buffered: quantized A + B data for 2 k_chunks
v4i32 a_q[2][K_CHUNKS]; // [buf][kc] quantized fp4
uint32_t a_sc_buf[2][K_CHUNKS]; // [buf][kc] scales
v4i32 b_reg[2][K_CHUNKS][N_SUBS]; // [buf][kc][ns]
uint32_t b_sc_buf[2][K_CHUNKS][N_SUBS];
v4i32 a_raw[K_CHUNKS][4]; // raw bf16 before quantize
// B scale: c%4, (c/4)%2 constant across steps. k_blk tracked incrementally.
const int k_inn_kc0 = c_base % 4;
const int k_inn_kc1 = (c_base + 2) % 4;
const int k_sub_kc0 = (c_base / 4) % 2;
const int k_sub_kc1 = ((c_base + 2) / 4) % 2;
int bs_const[K_CHUNKS][N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
bs_const[0][ns] = bs_n_blk_base[ns] + k_inn_kc0 * 64 + bs_n_inn
+ k_sub_kc0 * 2 + bs_n_sub_val[ns];
bs_const[1][ns] = bs_n_blk_base[ns] + k_inn_kc1 * 64 + bs_n_inn
+ k_sub_kc1 * 2 + bs_n_sub_val[ns];
}
constexpr int K_ADVANCE = N_WF * K_STEP;
constexpr int C_INC = K_ADVANCE / 32;
constexpr int B_BYTE_INC = C_INC * 256;
constexpr int K_BLK_INC = C_INC / 8;
int k_blk_kc0 = c_base / 8;
int k_blk_kc1 = (c_base + 2) / 8;
auto load_a_raw = [&]() __attribute__((always_inline)) {
const v4i32* src0 = reinterpret_cast<const v4i32*>(a_load_ptr);
const v4i32* src1 = reinterpret_cast<const v4i32*>(a_load_ptr + TILE_K_32);
#pragma unroll
for (int v = 0; v < 4; v++) a_raw[0][v] = src0[v];
#pragma unroll
for (int v = 0; v < 4; v++) a_raw[1][v] = src1[v];
};
auto load_b = [&](int buf) __attribute__((always_inline)) {
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
b_reg[buf][0][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns]);
b_sc_buf[buf][0][ns] = B_scale[bs_const[0][ns] + k_blk_kc0 * 256];
b_reg[buf][1][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns] + 512);
b_sc_buf[buf][1][ns] = B_scale[bs_const[1][ns] + k_blk_kc1 * 256];
}
};
auto quantize_to = [&](int kc, int buf) __attribute__((always_inline)) {
__hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw[kc]);
quantize_group(vals, a_q[buf][kc], a_sc_buf[buf][kc]);
};
auto advance_ptrs = [&]() __attribute__((always_inline)) {
a_load_ptr += K_ADVANCE;
k_blk_kc0 += K_BLK_INC;
k_blk_kc1 += K_BLK_INC;
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
b_byte_off[ns] += B_BYTE_INC;
};
// Prologue: load + quantize step 0 into buf 0
load_a_raw();
load_b(0);
quantize_to(0, 0);
quantize_to(1, 0);
int cur = 0;
// Main K loop: MFMA first, then issue loads (overlap), then quantize
#pragma unroll
for (int ki = 0; ki < K_ITERS_PER_WF - 1; ki++) {
int nxt = cur ^ 1;
// MFMA[kc=0, ns=0] — start matrix core
acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]
);
// Issue loads for next step while MFMA[0] in flight
advance_ptrs();
load_a_raw();
load_b(nxt);
// Remaining kc=0 MFMAs (A kc=0 reused across n_subs)
#pragma unroll
for (int ns = 1; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]
);
// All kc=1 MFMAs (A kc=1 reused across n_subs)
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]
);
// Quantize arrived A data
quantize_to(0, nxt);
quantize_to(1, nxt);
cur = nxt;
}
// Final step: just MFMAs
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]
);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]
);
// Write acc directly to LDS in row-major layout [wf_id][row][col] as f32
__shared__ float lds_f32[N_WF][TILE_32][TILE_32 * N_SUBS];
{
int k_group = tid_in_wf / 32;
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
int col_local = ns * TILE_32 + lane;
#pragma unroll
for (int i = 0; i < 16; i++) {
int row_local = k_group * 4 + (i / 4) * 8 + (i % 4);
lds_f32[wf_id][row_local][col_local] = acc[ns][i];
}
}
}
__syncthreads();
// All threads: 128-bit LDS reads, reduce, convert, 128-bit global write
{
int local_idx = tid * 8;
int row_local = local_idx / (TILE_32 * N_SUBS);
int col_local = local_idx % (TILE_32 * N_SUBS);
// Read 2 groups of 4 f32 from each wf slot, sum, convert to bf16
__hip_bfloat16 out[8];
#pragma unroll
for (int g = 0; g < 2; g++) {
float sum[4] = {0, 0, 0, 0};
#pragma unroll
for (int w = 0; w < N_WF; w++) {
const float* src = &lds_f32[w][row_local][col_local + g * 4];
#pragma unroll
for (int j = 0; j < 4; j++)
sum[j] += src[j];
}
#pragma unroll
for (int j = 0; j < 4; j++)
out[g * 4 + j] = __float2bfloat16(sum[j]);
}
// 128-bit global write (8 bf16 = 16 bytes)
*reinterpret_cast<v4i32*>(&C[(m_base + row_local) * N + n_base + col_local]) =
*reinterpret_cast<v4i32*>(out);
}
}
// ============================================================
// Shape 5 phased kernel: 32×32 GEMM + prequant on idle CUs.
// Blocks 0..N_QUANT_BLOCKS-1: quantize A → global + volatile store + flag.
// Blocks N_QUANT_BLOCKS..N_QUANT_BLOCKS+N_GEMM_BLOCKS-1: GEMM with toggle flags.
// Grid: (N_QUANT_BLOCKS + N_GEMM_BLOCKS, 1, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_TOTAL, int N_TOTAL, int N_WF, int N_SUBS_T,
int N_GEMM_BLOCKS, int N_QUANT_BLOCKS = 16,
int FLAG_VAL = 1, uint64_t FLAG_PAT8 = 0x0101010101010101ULL>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_shape5_phased(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_scale,
__hip_bfloat16* __restrict__ C,
int B_scale_stride,
v4i32* __restrict__ A_q_out,
uint32_t* __restrict__ A_scale_out,
FlagLine* __restrict__ flag_line
) {
constexpr int K = K_TOTAL;
constexpr int M = M_TOTAL;
constexpr int N = N_TOTAL;
constexpr int GROUPS_PER_ROW = K / GROUP_SZ;
constexpr int TOTAL_GROUPS = M * GROUPS_PER_ROW;
const int tid = threadIdx.x;
const int threads_per_block = N_WF * WF_SIZE;
// ======== Quant blocks: quantize A, volatile store, set flag, return ========
if (blockIdx.x < N_QUANT_BLOCKS) {
int idx = blockIdx.x * threads_per_block + tid;
if (idx < TOTAL_GROUPS) {
int row = idx / GROUPS_PER_ROW;
int k_grp = idx % GROUPS_PER_ROW;
const __hip_bfloat16* src = A + row * K + k_grp * GROUP_SZ;
__hip_bfloat16 vals[GROUP_SZ];
const v4i32* src_vec = reinterpret_cast<const v4i32*>(src);
v4i32* dst_vec = reinterpret_cast<v4i32*>(vals);
#pragma unroll
for (int v = 0; v < 4; v++) dst_vec[v] = src_vec[v];
v4i32 out; uint32_t scale;
quantize_group(vals, out, scale);
int out_idx = k_grp * M + row;
*reinterpret_cast<volatile v4i32*>(&A_q_out[out_idx]) = out;
*reinterpret_cast<volatile uint32_t*>(&A_scale_out[out_idx]) = scale;
}
__builtin_amdgcn_s_waitcnt(0);
__syncthreads();
if (tid == 0) {
*reinterpret_cast<volatile unsigned char*>(&flag_line->flags[blockIdx.x]) = FLAG_VAL;
}
return;
}
// ======== GEMM blocks with phased toggle ========
constexpr int TILE_32 = 32;
constexpr int TILE_K_32 = 64;
constexpr int GROUP_SZ_32 = 32;
constexpr int N_SUBS = N_SUBS_T;
constexpr int K_CHUNKS = 2;
constexpr int K_STEP = K_CHUNKS * TILE_K_32; // 128 fp4 per step
constexpr int K_ITERS_PER_WF = K / N_WF / K_STEP;
const int wf_id = tid / WF_SIZE;
const int tid_in_wf = tid % WF_SIZE;
const int lane = tid_in_wf % 32;
const int k_chunk = tid_in_wf / 32;
const int gemm_idx = blockIdx.x - N_QUANT_BLOCKS;
const int n_block = gemm_idx % (N / (TILE_32 * N_SUBS));
const int m_tile = gemm_idx / (N / (TILE_32 * N_SUBS));
const int n_base = n_block * (TILE_32 * N_SUBS);
const int m_base = m_tile * TILE_32;
const int my_a_row = m_base + lane;
const __hip_bfloat16* A_row_ptr = A + my_a_row * K;
const int b_lane_in_tile = lane % 16;
const int b_n_sub_tile_off = lane / 16;
int b_tile_n_idx[N_SUBS];
int bs_n_blk_base[N_SUBS];
int bs_n_sub_val[N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
int n_off = n_base + ns * TILE_32 + lane;
b_tile_n_idx[ns] = (n_base + ns * TILE_32) / 16 + b_n_sub_tile_off;
bs_n_blk_base[ns] = (n_off / 32) * (B_scale_stride * 32);
bs_n_sub_val[ns] = (n_off / 16) % 2;
}
const int bs_n_inn = (lane % 16) * 4;
v16f32 acc[N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = {0,0,0,0, 0,0,0,0, 0,0,0,0, 0,0,0,0};
const __hip_bfloat16* a_load_ptr = A_row_ptr + wf_id * K_STEP + k_chunk * GROUP_SZ_32;
const int c_base = (wf_id * K_STEP) / 32 + k_chunk;
int b_byte_off[N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
int tile_idx_0 = b_tile_n_idx[ns] * (K / 32) + c_base;
b_byte_off[ns] = tile_idx_0 * 256 + b_lane_in_tile * 16;
}
v4i32 a_q[2][K_CHUNKS];
uint32_t a_sc_buf[2][K_CHUNKS];
v4i32 b_reg[2][K_CHUNKS][N_SUBS];
uint32_t b_sc_buf[2][K_CHUNKS][N_SUBS];
v4i32 a_raw[K_CHUNKS][4];
const int k_inn_kc0 = c_base % 4;
const int k_inn_kc1 = (c_base + 2) % 4;
const int k_sub_kc0 = (c_base / 4) % 2;
const int k_sub_kc1 = ((c_base + 2) / 4) % 2;
int bs_const[K_CHUNKS][N_SUBS];
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
bs_const[0][ns] = bs_n_blk_base[ns] + k_inn_kc0 * 64 + bs_n_inn
+ k_sub_kc0 * 2 + bs_n_sub_val[ns];
bs_const[1][ns] = bs_n_blk_base[ns] + k_inn_kc1 * 64 + bs_n_inn
+ k_sub_kc1 * 2 + bs_n_sub_val[ns];
}
constexpr int K_ADVANCE = N_WF * K_STEP;
constexpr int C_INC = K_ADVANCE / 32;
constexpr int B_BYTE_INC = C_INC * 256;
constexpr int K_BLK_INC = C_INC / 8;
int k_blk_kc0 = c_base / 8;
int k_blk_kc1 = (c_base + 2) / 8;
// Track absolute K-group for prequant indexing
int c_kc0 = c_base;
int c_kc1 = c_base + 2;
auto load_a_raw = [&]() __attribute__((always_inline)) {
const v4i32* src0 = reinterpret_cast<const v4i32*>(a_load_ptr);
const v4i32* src1 = reinterpret_cast<const v4i32*>(a_load_ptr + TILE_K_32);
#pragma unroll
for (int v = 0; v < 4; v++) a_raw[0][v] = src0[v];
#pragma unroll
for (int v = 0; v < 4; v++) a_raw[1][v] = src1[v];
};
auto load_b = [&](int buf) __attribute__((always_inline)) {
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
b_reg[buf][0][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns]);
b_sc_buf[buf][0][ns] = B_scale[bs_const[0][ns] + k_blk_kc0 * 256];
b_reg[buf][1][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns] + 512);
b_sc_buf[buf][1][ns] = B_scale[bs_const[1][ns] + k_blk_kc1 * 256];
}
};
auto quantize_to = [&](int kc, int buf) __attribute__((always_inline)) {
__hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw[kc]);
quantize_group(vals, a_q[buf][kc], a_sc_buf[buf][kc]);
};
auto read_prequant_a = [&](int buf) __attribute__((always_inline)) {
int pq0 = c_kc0 * M + my_a_row;
int pq1 = c_kc1 * M + my_a_row;
a_q[buf][0] = A_q_out[pq0];
a_sc_buf[buf][0] = A_scale_out[pq0];
a_q[buf][1] = A_q_out[pq1];
a_sc_buf[buf][1] = A_scale_out[pq1];
};
auto advance_ptrs_bf16 = [&]() __attribute__((always_inline)) {
a_load_ptr += K_ADVANCE;
k_blk_kc0 += K_BLK_INC;
k_blk_kc1 += K_BLK_INC;
c_kc0 += C_INC;
c_kc1 += C_INC;
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
b_byte_off[ns] += B_BYTE_INC;
};
auto advance_ptrs_pq = [&]() __attribute__((always_inline)) {
k_blk_kc0 += K_BLK_INC;
k_blk_kc1 += K_BLK_INC;
c_kc0 += C_INC;
c_kc1 += C_INC;
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
b_byte_off[ns] += B_BYTE_INC;
};
constexpr int K_THRESHOLD = 1;
// ---- Iterations 0 to K_THRESHOLD-1: bf16 ----
load_a_raw();
load_b(0);
quantize_to(0, 0);
quantize_to(1, 0);
int cur = 0;
#pragma unroll
for (int ki = 0; ki < K_THRESHOLD; ki++) {
int nxt = cur ^ 1;
acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);
advance_ptrs_bf16();
load_a_raw();
load_b(nxt);
#pragma unroll
for (int ns = 1; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
quantize_to(0, nxt);
quantize_to(1, nxt);
cur = nxt;
}
// ---- Load flags, quantize current (masks flag latency), check ----
constexpr int N_FLAG_WORDS = N_QUANT_BLOCKS / 8;
const uint64_t* f64 = reinterpret_cast<const uint64_t*>(flag_line->flags);
uint64_t fv[N_FLAG_WORDS];
#pragma unroll
for (int i = 0; i < N_FLAG_WORDS; i++) fv[i] = f64[i];
// cur is already quantized from the loop above
bool use_pq = true;
#pragma unroll
for (int i = 0; i < N_FLAG_WORDS; i++) {
if (fv[i] != FLAG_PAT8) { use_pq = false; break; }
}
if constexpr (K_ITERS_PER_WF - K_THRESHOLD == 1) {
// Only 1 iteration left after threshold — just epilogue, no transition needed
// cur already has quantized data from the K_THRESHOLD loop
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
} else if (use_pq) {
// ---- Transition iter: cur already quantized, prefetch prequant for next ----
{
int nxt = cur ^ 1;
acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);
advance_ptrs_pq();
read_prequant_a(nxt);
load_b(nxt);
#pragma unroll
for (int ns = 1; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
cur = nxt;
}
// ---- Remaining prequant iterations ----
#pragma unroll
for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
int nxt = cur ^ 1;
acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);
advance_ptrs_pq();
read_prequant_a(nxt);
load_b(nxt);
#pragma unroll
for (int ns = 1; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
cur = nxt;
}
// Epilogue
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
} else {
// ---- Transition iter: cur already quantized, bf16 load for next ----
{
int nxt = cur ^ 1;
acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);
advance_ptrs_bf16();
load_a_raw();
load_b(nxt);
#pragma unroll
for (int ns = 1; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
quantize_to(0, nxt);
quantize_to(1, nxt);
cur = nxt;
}
// ---- Remaining bf16 iterations ----
#pragma unroll
for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
int nxt = cur ^ 1;
acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);
advance_ptrs_bf16();
load_a_raw();
load_b(nxt);
#pragma unroll
for (int ns = 1; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
quantize_to(0, nxt);
quantize_to(1, nxt);
cur = nxt;
}
// Epilogue
quantize_to(0, cur);
quantize_to(1, cur);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++)
acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
}
// LDS reduction (same as mxfp4_gemm_32x32)
__shared__ float lds_f32[N_WF][TILE_32][TILE_32 * N_SUBS];
{
int k_group = tid_in_wf / 32;
#pragma unroll
for (int ns = 0; ns < N_SUBS; ns++) {
int col_local = ns * TILE_32 + lane;
#pragma unroll
for (int i = 0; i < 16; i++) {
int row_local = k_group * 4 + (i / 4) * 8 + (i % 4);
lds_f32[wf_id][row_local][col_local] = acc[ns][i];
}
}
}
__syncthreads();
{
constexpr int TOTAL_OUT = TILE_32 * TILE_32 * N_SUBS;
constexpr int THREADS = N_WF * WF_SIZE;
constexpr int ELEMS = TOTAL_OUT / THREADS;
if constexpr (ELEMS >= 8) {
// 8 elements/thread, 128-bit write
int local_idx = tid * 8;
int row_local = local_idx / (TILE_32 * N_SUBS);
int col_local = local_idx % (TILE_32 * N_SUBS);
__hip_bfloat16 out[8];
#pragma unroll
for (int g = 0; g < 2; g++) {
float sum[4] = {0, 0, 0, 0};
#pragma unroll
for (int w = 0; w < N_WF; w++) {
const float* src = &lds_f32[w][row_local][col_local + g * 4];
#pragma unroll
for (int j = 0; j < 4; j++)
sum[j] += src[j];
}
#pragma unroll
for (int j = 0; j < 4; j++)
out[g * 4 + j] = __float2bfloat16(sum[j]);
}
*reinterpret_cast<v4i32*>(&C[(m_base + row_local) * N + n_base + col_local]) =
*reinterpret_cast<v4i32*>(out);
} else {
// 4 elements/thread, 64-bit write
int local_idx = tid * 4;
if (local_idx < TOTAL_OUT) {
int row_local = local_idx / (TILE_32 * N_SUBS);
int col_local = local_idx % (TILE_32 * N_SUBS);
float s0 = 0, s1 = 0, s2 = 0, s3 = 0;
#pragma unroll
for (int w = 0; w < N_WF; w++) {
v4f32 v = *reinterpret_cast<const v4f32*>(&lds_f32[w][row_local][col_local]);
s0 += v[0]; s1 += v[1]; s2 += v[2]; s3 += v[3];
}
int pk0, pk1;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk0) : "v"(s0), "v"(s1));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk1) : "v"(s2), "v"(s3));
ext_u32x2 out_vec = {(uint32_t)pk0, (uint32_t)pk1};
int gbl_off = (m_base + row_local) * N + n_base + col_local;
*reinterpret_cast<ext_u32x2*>(&C[gbl_off]) = out_vec;
}
}
}
}
// ============================================================
// M=32 kernel: 8 wf = 4 K-splits × 2 N-subs. 16×32 output tile.
// A reused across n_sub via L1 cache. Each wf loads its own B.
// Grid: (N/32, 2, 1), Block: 512
// ============================================================
template <int K_TOTAL, int N_TOTAL>
__global__ __attribute__((amdgpu_flat_work_group_size(512, 512)))
void mxfp4_gemm_m32(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_scale,
__hip_bfloat16* __restrict__ C,
int B_scale_stride
) {
constexpr int K = K_TOTAL;
constexpr int N = N_TOTAL;
constexpr int K_SPLITS = K / TILE_K; // 4
constexpr int N_SUBS = 2;
const int tid = threadIdx.x;
const int wf_id = tid / WF_SIZE; // 0..7
const int tid_in_wf = tid % WF_SIZE;
const int lane = tid_in_wf % TILE_N; // 0..15
const int k_group = tid_in_wf / TILE_N; // 0..3
const int k_split = wf_id & 3; // 0..3
const int n_sub = wf_id >> 2; // 0 or 1
const int m_base = blockIdx.y * TILE_M; // 0 or 16
const int n_tile = blockIdx.x * N_SUBS + n_sub;
const int a_row = m_base + lane;
const int k_off = k_split * TILE_K + k_group * GROUP_SZ;
const int scale_thread_off = k_group * 64 + lane * 4;
// Load A (both n_sub wf load same A — n_sub=1 hits L1 cache)
__hip_bfloat16 vals[GROUP_SZ];
const v4i32* src_vec = reinterpret_cast<const v4i32*>(A + a_row * K + k_off);
v4i32* dst_vec = reinterpret_cast<v4i32*>(vals);
#pragma unroll
for (int v = 0; v < 4; v++) dst_vec[v] = src_vec[v];
// Load B (each wf loads its own n_tile's B)
int tile_idx = n_tile * (K / 32) + k_split * 4 + k_group;
int byte_offset = tile_idx * 256 + lane * 16;
v4i32 b_reg = *reinterpret_cast<const v4i32*>(B_shuf + byte_offset);
int scale_base = (n_tile >> 1) * (B_scale_stride * 32) + ((k_split >> 1) * 256);
uint32_t scale_b = B_scale[scale_base + scale_thread_off + ((k_split & 1) << 1) + (n_tile & 1)];
// Quantize A
v4i32 a_reg;
uint32_t scale_a;
quantize_group(vals, a_reg, scale_a);
// Single MFMA
v4f32 acc = {0.0f, 0.0f, 0.0f, 0.0f};
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(a_reg), to_v8(b_reg), acc,
4, 4, 0, scale_a, 0, scale_b
);
// K-split reduction via LDS
__shared__ float reduce_lds[N_SUBS][K_SPLITS][4][WF_SIZE];
#pragma unroll
for (int i = 0; i < 4; i++)
reduce_lds[n_sub][k_split][i][tid_in_wf] = acc[i];
__syncthreads();
if (k_split == 0) {
float sum[4];
#pragma unroll
for (int i = 0; i < 4; i++) {
sum[i] = reduce_lds[n_sub][0][i][tid_in_wf]
+ reduce_lds[n_sub][1][i][tid_in_wf]
+ reduce_lds[n_sub][2][i][tid_in_wf]
+ reduce_lds[n_sub][3][i][tid_in_wf];
}
int out_col = n_tile * TILE_N + lane;
#pragma unroll
for (int i = 0; i < 4; i++) {
int out_row = m_base + k_group * 4 + i;
C[out_row * N + out_col] = __float2bfloat16(sum[i]);
}
}
}
// Host entry point — raw pointers, C++ dispatch
// ============================================================
static int d_call_count = 0;
void run(
uintptr_t a_ptr,
uintptr_t b_shuf_ptr,
uintptr_t b_scale_ptr,
uintptr_t c_ptr,
int M, int N, int K,
int B_scale_stride,
uintptr_t aq_ptr,
uintptr_t as_ptr,
uintptr_t flag_ptr,
uintptr_t aq5_ptr,
uintptr_t as5_ptr,
uintptr_t flag5_ptr,
uintptr_t ws_ptr
) {
d_call_count++;
const auto* A = reinterpret_cast<const __hip_bfloat16*>(a_ptr);
const auto* B_shuf = reinterpret_cast<const uint8_t*>(b_shuf_ptr);
const auto* B_scale = reinterpret_cast<const uint8_t*>(b_scale_ptr);
auto* C = reinterpret_cast<__hip_bfloat16*>(c_ptr);
// Shape 1: M<=4, K=512 — 4 wavefronts split K
if (M <= 4 && K == 512) {
int num_n_tiles = (N + TILE_N - 1) / TILE_N;
mxfp4_gemm_smallm<512, 4, 4><<<dim3(num_n_tiles, 1, 1), 4 * WF_SIZE>>>(
A, B_shuf, B_scale, C,
N, B_scale_stride
);
return;
}
// Shape 2: M<=16, K=7168 — split-K=7, 8wf K-split, N_SUBS=4 (16×64 tile)
if (M <= 16 && K == 7168) {
auto* ws = reinterpret_cast<float*>(ws_ptr);
constexpr int SPLIT_K = 7;
constexpr int BLOCK_N = 64;
constexpr int N_BLOCKS = (2112 + BLOCK_N - 1) / BLOCK_N; // 33
// Kernel 1: 8 wf split K, each does 4 MFMAs (N_SUBS=4), LDS reduce
mxfp4_gemm_splitk_gemm<7168, 16, 2112, SPLIT_K>
<<<dim3(N_BLOCKS * SPLIT_K, 1, 1), 8 * WF_SIZE>>>(
A, B_shuf, B_scale, B_scale_stride, ws);
// Kernel 2: reduce 7 partials per 16×64 block
mxfp4_splitk_reduce<16, 2112, SPLIT_K>
<<<dim3(N_BLOCKS, 1, 1), 256>>>(ws, C);
return;
}
// Shapes 3 & 4: M=32, K=512 — 8 wf (2 m_tiles × 4 K-splits), B in LDS
if (M == 32 && K == 512 && N == 4096) {
mxfp4_gemm_m32<512, 4096><<<dim3(4096 / (TILE_N * 2), 2, 1), 512>>>(
A, B_shuf, B_scale, C, B_scale_stride
);
return;
}
if (M == 32 && K == 512) {
mxfp4_gemm_m32<512, 2880><<<dim3(2880 / (TILE_N * 2), 2, 1), 512>>>(
A, B_shuf, B_scale, C, B_scale_stride
);
return;
}
// Shape 5: M=64, N=7168, K=2048 — phased: 16 quant blocks + 224 GEMM blocks
if (M == 64 && K == 2048) {
auto* A_q_out5 = reinterpret_cast<v4i32*>(aq5_ptr);
auto* A_scale_out5 = reinterpret_cast<uint32_t*>(as5_ptr);
auto* flag_line5 = reinterpret_cast<FlagLine*>(flag5_ptr);
constexpr int N_GEMM_BLOCKS5 = 224; // (7168/64) * (64/32) = 112 * 2
constexpr int N_QUANT5 = 16; // multiple of 8
constexpr int TOTAL5 = N_QUANT5 + N_GEMM_BLOCKS5;
if (d_call_count & 1)
mxfp4_gemm_shape5_phased<2048, 64, 7168, 8, 2, N_GEMM_BLOCKS5, N_QUANT5, 1, 0x0101010101010101ULL>
<<<dim3(TOTAL5, 1, 1), 8 * WF_SIZE>>>(
A, B_shuf, B_scale, C, B_scale_stride, A_q_out5, A_scale_out5, flag_line5);
else
mxfp4_gemm_shape5_phased<2048, 64, 7168, 8, 2, N_GEMM_BLOCKS5, N_QUANT5, 2, 0x0202020202020202ULL>
<<<dim3(TOTAL5, 1, 1), 8 * WF_SIZE>>>(
A, B_shuf, B_scale, C, B_scale_stride, A_q_out5, A_scale_out5, flag_line5);
return;
}
// Shape 6: M=256, N=3072, K=1536 — 32×96 tile, 6 wf K-split, 3 n_subs
if (M == 256 && K == 1536) {
mxfp4_gemm_32x32<1536, 256, 3072, 6, 3><<<dim3(3072 / 96, 256 / 32, 1), 6 * WF_SIZE>>>(
A, B_shuf, B_scale, C, B_scale_stride
);
return;
}
}
PYBIND11_MODULE(mxfp4_gemm_pt, m) {
m.def("run", &run, "MXFP4 GEMM kernel with phased quantization");
}
"""
module = load_inline(
name='mxfp4_gemm_pt',
cpp_sources='',
cuda_sources=HIP_SRC,
with_cuda=True,
verbose=True,
extra_cuda_cflags=["-std=c++20", "-O3", "--offload-arch=gfx950"],
no_implicit_headers=True,
)
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)
# Phased buffers (allocated once)
_aq_buf = None # Shape 2: v4i32 per group (3584 × 16 bytes)
_as_buf = None # Shape 2: uint32 per group (3584 × 4 bytes)
_flag_buf = None # Shape 2: FlagLine: 128 bytes
_aq5_buf = None # Shape 5: v4i32 per group (4096 × 16 bytes)
_as5_buf = None # Shape 5: uint32 per group (4096 × 4 bytes)
_flag5_buf = None # Shape 5: FlagLine: 128 bytes
_ws_buf = None # Shape 2 split-K: workspace [14 * 132 * 16 * 16] floats
def _ensure_phased_bufs(device):
global _aq_buf, _as_buf, _flag_buf, _aq5_buf, _as5_buf, _flag5_buf, _ws_buf
if _aq_buf is None:
n = 16 * (7168 // 32) # 3584
_aq_buf = torch.empty(n * 4, dtype=torch.int32, device=device)
_as_buf = torch.empty(n, dtype=torch.int32, device=device)
_flag_buf = torch.empty(128, dtype=torch.uint8, device=device)
if _aq5_buf is None:
n5 = 64 * (2048 // 32) # 4096
_aq5_buf = torch.empty(n5 * 4, dtype=torch.int32, device=device)
_as5_buf = torch.empty(n5, dtype=torch.int32, device=device)
_flag5_buf = torch.empty(128, dtype=torch.uint8, device=device)
if _ws_buf is None:
_ws_buf = torch.empty(33 * 7 * 16 * 64, dtype=torch.float32, device=device) # N_BLOCKS * SPLIT_K * M * BLOCK_N
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.size(0)
C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
if (M == 8) or (M == 16 and N == 3072) or \
(M == 64 and N == 3072) or (M == 256 and N == 2880):
A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
out_gemm = aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True)
return out_gemm
_ensure_phased_bufs(A.device)
module.run(
A.data_ptr(),
B_shuffle.data_ptr(),
B_scale_sh.data_ptr(),
C.data_ptr(),
M, N, K,
B_scale_sh.size(1),
_aq_buf.data_ptr(),
_as_buf.data_ptr(),
_flag_buf.data_ptr(),
_aq5_buf.data_ptr(),
_as5_buf.data_ptr(),
_flag5_buf.data_ptr(),
_ws_buf.data_ptr(),
)
return C
scrolls · 1679 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