submission 647149
npip99 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3129 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-647149?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:2df11d0f8fc16cce4bdd05b3005dc42428e47c1a30ba5d65f98d787482fbecdc
license declaredunknown
license concludedunknown
authorsnpip99
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
constexpr int NUM_WARPS = 4;shared-memory
__shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];split-k
- SPLIT_K: Same as general, but with SPLIT_K implemented.tile-m = 1
constexpr uint32_t WARP_TILE_M = 1;tile-n = 1
constexpr uint32_t WARP_TILE_N = 1;Kernel source
submission.py3129 lines
"""
There are three main kernels implemented:
- General: Global->Register, 1 Warp : 1 Output tile
- SPLIT_K: Same as general, but with SPLIT_K implemented.
- Used for Shape2 (16, 2112, 7168). Hyperparameters are selected to minimize L2 traffic.
- Global->LDS, LDS->Register, Warp tiling + Threadblock tiling
- Used for the other 5 shapes
For the other 5 shapes, the C++ code is duplicated. The code is virtually identical, the only difference is the constexprs at the top (i.e. the tiling hyperparameters), and bounds checking.
TODO: Use -D and constexpr evaluations to prevent big copypaste. Constexpr could also evaluate the bounds checks at compile time.
"""
import os
from typing import Any
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_WRAPPER = """
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
);
"""
CUDA_SRC_SHAPE_GENERAL = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using f32x4_t = float __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using u8x32_t = uint8_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
__device__ inline uint16_t f32_to_bf16(float v) {
return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const uint16_t* row) {
uint16_t amax_u16 = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
uint32_t v = row32[i] & 0x7FFF7FFFu; // abs both bf16 in one AND
uint16_t lo = (uint16_t)(v);
uint16_t hi = (uint16_t)(v >> 16);
amax_u16 = max(amax_u16, max(lo, hi));
}
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
int tr = n_row / 32, r = n_row % 32;
int tc = k_group / 8, c = k_group % 8;
int base = tr * (32 * b_scale_stride) + tc * 256;
int inner = (r / 16) + ((c / 4) % 2) * 2 + (r % 16) * 4 + (c % 4) * 64;
return base + inner;
}
constexpr int MFMA_DIM_NM = 16;
constexpr int MFMA_DIM_K = 128;
constexpr int SCALE_GROUP_SIZE = 32;
constexpr int NUM_THREADS = 64;
template <int M, int N, int K>
__launch_bounds__(NUM_THREADS)
__global__ void kernel(
const uint16_t* a, const uint16_t* b,
const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
uint16_t* c
) {
constexpr bool BOUNDS_CHECK = (M % MFMA_DIM_NM != 0) || (N % MFMA_DIM_NM != 0);
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;
int tile_x = MFMA_DIM_NM * blockIdx.x; // N
int tile_y = MFMA_DIM_NM * blockIdx.y; // M
int lane_index = threadIdx.x;
int lane_row_nm = lane_index % MFMA_DIM_NM;
int lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.
// C[y][x] = Sum_{k=0}^{K-1} a[y][k] * b[x][k]
f32x4_t result = {};
#pragma unroll
for (int k_offset = 0; k_offset < K; k_offset += MFMA_DIM_K) {
int k_idx_start = k_offset + lane_col_k;
uint8_t a_scale = 0;
u32x8_t a_reg = {};
if (!BOUNDS_CHECK || tile_y + lane_row_nm < M) {
// Quantize A
const uint16_t* a_row = a + (tile_y + lane_row_nm) * K;
a_scale = bf16x32_to_scale_e8m0(a_row + k_idx_start);
float a_scale_f32 = e8m0_to_f32(a_scale);
#pragma unroll
for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
const bf16x2_t* a_row_v2 = (const bf16x2_t*)(a_row + k_idx_start + reg_idx * 8);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[0], a_scale_f32, 0);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
}
}
uint8_t b_scale = b_scale_sh[read_b_scale_base(tile_x + lane_row_nm, k_idx_start / SCALE_GROUP_SIZE, b_scale_stride)];
i32x8_t b_reg;
*(i32x4_t*)&b_reg = *(const i32x4_t*)&b_q[(tile_x + lane_row_nm) * (K / 2) + k_idx_start / 2];
// This is not faster unless we use LDS
// *(i32x4_t*)&b_reg = *(const i32x4_t*)&b_shuffle[
// tile_x * (K / 2)
// + (k_idx_start / 32) * 256
// + lane_row_nm * 16
// ];
result = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, b_reg, result, 4, 4, 0, a_scale, 0, b_scale);
}
int result_row_offset = 4 * (lane_index / 16);
int result_column_offset = lane_index % 16;
#pragma unroll
for (int i = 0; i < 4; i++) {
int row = tile_y + result_row_offset + i;
if (!BOUNDS_CHECK || row < M) {
c[row * N + (tile_x + result_column_offset)] = f32_to_bf16((float)result[i]);
}
}
}
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
) {
assert(K % MFMA_DIM_K == 0);
#define LAUNCH(m,n,k) kernel<m,n,k><<<dim3(CDIV(N, 16), CDIV(M, 16)), NUM_THREADS>>>( \
(const uint16_t*)a, (const uint16_t*)b, \
(const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh, \
(uint16_t*)c)
// Benchmark Shapes
if (M == M_DIM && N == N_DIM && K == K_DIM) {
LAUNCH(M_DIM, N_DIM, K_DIM);
} else {
assert(false && "Uncompiled (M, N, K) shape");
}
}
"""
CUDA_SRC_SHAPE1 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
__device__ inline uint16_t f32_to_bf16(float v) {
return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32_t* row) {
uint32_t amax_u16x2_packed = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row32[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
// CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
// FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
// denorm_exp = (127-1) + (23-1) + 1 = 149 => denorm_mask = 149 << 23 = 0x4A800000
// val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF (int32 wrapping)
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT = 0x4A800000u;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr uint32_t VAL_TO_ADD = 0xC11FFFFFu;
uint8_t s_bit = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
float abs_v = __builtin_bit_cast(float, abs_bits);
uint8_t result;
if (abs_v >= FP4_MAX_NORMAL) {
// saturate branch
result = 0x7u;
} else if (abs_v < FP4_MIN_NORMAL) {
// denormal branch: float trick to extract low bits
uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
result = (uint8_t)d;
} else {
// normal branch: bias-adjust exponent + round-to-nearest
uint32_t x = abs_bits;
uint32_t m_odd = (x >> 22) & 1u;
x = x + VAL_TO_ADD + m_odd;
result = (uint8_t)(x >> 22);
}
return s_bit | result;
}
// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
return f32_to_fp4_e2m1(v * f32_scale);
}
__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
uint8_t s = (nibble >> 3) & 1;
uint8_t e = (nibble >> 1) & 3;
uint8_t m = nibble & 1;
float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
return s ? -val : val;
}
// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
constexpr int POSITIONS = 256 / WIDTH; // slots per bank cycle
constexpr int MASK = POSITIONS - 1;
int c_group = c / WIDTH;
int c_offset = c % WIDTH; // both compile to shift/mask
int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
return swizzled * WIDTH + c_offset;
}
// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;
// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 1;
constexpr uint32_t NUM_CU_N = 9;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 5;
#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;
// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD_M * NUM_XCD_N * NUM_CUS_PER_XCD;
// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;
// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
int tr = n_row / 32, r = n_row % 32;
int tc = k_group / 8, c = k_group % 8;
int base = tr * (32 * b_scale_stride) + tc * 256;
int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
return base + inner;
}
// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
return x;
}
__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
const uint16_t* a,
const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
uint16_t* c
) {
// Timing
constexpr bool TIMEIT = false;
uint64_t _ts[16];
int _ti = 0;
auto TIME = [&]() {
if constexpr (TIMEIT) {
__builtin_amdgcn_sched_barrier(0);
_ts[_ti++] = __builtin_amdgcn_s_memrealtime();
__builtin_amdgcn_sched_barrier(0);
}
};
TIME();
constexpr uint32_t M = 16;
constexpr uint32_t M_ACTUAL = 4;
constexpr uint32_t N = 2880;
constexpr uint32_t K = 512;
constexpr uint32_t SPLIT_K = 1;
static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;
constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;
// Get XCD coordinate
uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
uint32_t xcd_n = xcd_id % NUM_XCD_N;
uint32_t xcd_m = xcd_id / NUM_XCD_N;
// Get CU coordinate
if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
return;
}
uint32_t cu_n = xcd_worker_id % NUM_CU_N;
uint32_t cu_m = xcd_worker_id / NUM_CU_N;
// Get warp id (Use SGPR for it)
uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
uint32_t warp_n = warp_id % NUM_WARPS_N;
uint32_t warp_m = warp_id / NUM_WARPS_N;
// Get lane assignment
uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.
// Get tile offset
uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;
// Shared Memory
__shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_a[BLOCK_M][K / 2];
__shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_b[BLOCK_N][K / 2];
auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };
// ======================
// A: Global -> Register
// ======================
TIME();
constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint32_t a_row_idx = tile_m_offset + a_transfer_m;
if (a_row_idx < M_ACTUAL) {
*(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
} else {
a_jobs[a_chunk_job_idx] = {};
}
}
auto a_process_job_step1 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
a_scale_jobs[a_chunk_job_idx] = a_scale;
__builtin_amdgcn_sched_barrier(0);
};
auto a_process_job_step2 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
u32x4_t a_reg;
#pragma unroll
for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
}
*(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
__builtin_amdgcn_sched_barrier(0);
};
uint32_t a_work_idx = 0;
// ======================
// B: Global -> Register
// ======================
TIME();
constexpr uint32_t B_BYTES_PER_JOB = 16;
constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);
u32x4_t b_jobs[B_NUM_JOBGROUPS];
uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
};
const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;
auto b_process_job_step0 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
uint32_t b_row_idx = tile_n_offset + b_transfer_n;
b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];
uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
void* src = (void*)(b_base + warp_base + lane_index * 16);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dwordx4 %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
#else
#pragma unroll
for (uint32_t sub = 0; sub < 4; sub++) {
uint32_t chunk_base = warp_base + sub * 256;
void* src = (void*)(b_base + chunk_base + lane_index * 4);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
}
#endif
__builtin_amdgcn_sched_barrier(0);
};
auto b_process_job_step1 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
asm volatile("" ::: "memory");
s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
asm volatile("" ::: "memory");
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step0(b_jobgroup_idx);
__builtin_amdgcn_sched_barrier(0);
if (b_jobgroup_idx == 1) {
a_process_job_step1(a_work_idx / 2);
} else if (b_jobgroup_idx == 2) {
a_process_job_step2(a_work_idx / 2);
} else if (b_jobgroup_idx == 3) {
b_process_job_step1(0);
}
__builtin_amdgcn_sched_barrier(0);
}
// ======================
// B: Register -> LDS
// ======================
TIME();
#pragma unroll
for (uint32_t b_jobgroup_idx = 1; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step1(b_jobgroup_idx);
}
// ======================
// A: Register -> Quantize -> LDS
// ======================
TIME();
// NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
// while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
// if (a_work_idx % 2 == 0) {
// a_process_job_step1(a_work_idx / 2);
// } else {
// a_process_job_step2(a_work_idx / 2);
// }
// a_work_idx += 1;
// }
// ======================
// LDS->MFMA
// ======================
__syncthreads(); // Ensure LDS is populated
TIME();
f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
constexpr uint32_t PREFETCH = 2;
uint8_t a_scale[PREFETCH][WARP_TILE_M];
u32x8_t a_reg[PREFETCH][WARP_TILE_M];
uint8_t b_scale[PREFETCH][WARP_TILE_N];
u32x8_t b_reg[PREFETCH][WARP_TILE_N];
auto load_lds = [&](uint32_t k_iter) {
__builtin_amdgcn_sched_barrier(0);
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
*(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
}
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_N; i++) {
uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
+ (k_idx_start / 32) * 256
+ lane_row_nm * 16;
*(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
}
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
load_lds(i);
}
#pragma unroll
for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
if (k_prefetch_iter < K_ITERS) {
load_lds(k_prefetch_iter);
}
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
f32x4_t mix;
for (uint32_t r = 0; r < 4; r++)
mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
result[i][j] += mix;
#endif
}
}
}
// ======================
// Reg->Write to Global
// ======================
TIME();
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for (uint32_t j = 0; j < WARP_TILE_N; j++) {
uint32_t result_row_offset = 4 * (lane_index / 16);
uint32_t result_column_offset = lane_index % 16;
#pragma unroll
for (uint32_t f = 0; f < 4; f++) {
uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
if (row < M_ACTUAL) {
c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
}
}
}
}
TIME();
if constexpr (TIMEIT) {
__shared__ uint32_t _run_idx;
if (threadIdx.x == 0) {
_run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
}
__syncthreads();
uint32_t h = hash(_run_idx);
if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
for (int i = 1; i < _ti; i++) {
printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
}
printf("\n");
}
}
}
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
) {
if (M == 4 && N == 2880 && K == 512) {
kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
(const uint16_t*)a,
(const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
(uint16_t*)c
);
} else {
// No impl
}
}
"""
CUDA_SRC_SHAPE2 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using f32x4_t = float __attribute__((ext_vector_type(4)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
__device__ inline uint16_t f32_to_bf16(float v) {
return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}
__device__ inline uint8_t bf16x32_to_scale_e8m0(uint32_t row[16]) {
uint32_t amax_u16x2_packed = 0;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
int tr = n_row / 32, r = n_row % 32;
int tc = k_group / 8, c = k_group % 8;
int base = tr * (32 * b_scale_stride) + tc * 256;
int inner = (r / 16) + ((c / 4) % 2) * 2 + (r % 16) * 4 + (c % 4) * 64;
return base + inner;
}
// Problem
constexpr int M = 16;
constexpr int N_ACTUAL = 2112;
constexpr int N = 2176;
constexpr int K = 7168;
constexpr int SPLIT_K = 7;
constexpr int BLOCK_K = K / SPLIT_K; // 1024
// Hardware
constexpr int MFMA_DIM_NM = 16;
constexpr int MFMA_DIM_K = 128;
constexpr int SCALE_GROUP_SIZE = 32;
constexpr int THREADS_PER_WARP = 64;
constexpr int NUM_WARPS = 4;
constexpr int NUM_THREADS = THREADS_PER_WARP * NUM_WARPS; // 256
constexpr int NUM_XCD_N = 8;
#ifdef __gfx950__
constexpr int NUM_CUS_PER_XCD = 32;
#else
constexpr int NUM_CUS_PER_XCD = 38;
#endif
// Tiling
constexpr int NUM_CU_N = 17; // N / (MFMA_DIM_NM * NUM_XCD_N) = 2176 / 128 = 17
constexpr int WORK_PER_XCD = NUM_CU_N * SPLIT_K; // 17 * 7 = 119
constexpr int MAX_WORK_PER_CU = CDIV(WORK_PER_XCD, NUM_CUS_PER_XCD); // 4
constexpr int K_ITERS = BLOCK_K / MFMA_DIM_K; // 8
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;
static_assert(N == MFMA_DIM_NM * NUM_CU_N * NUM_XCD_N);
static_assert(K % SPLIT_K == 0);
static_assert(BLOCK_K % MFMA_DIM_K == 0);
static_assert(MAX_WORK_PER_CU == NUM_WARPS);
constexpr int NUM_CUS = NUM_XCD_N * NUM_CUS_PER_XCD;
__launch_bounds__(NUM_THREADS, 1)
__global__ void kernel(
const uint16_t* a, const uint16_t* b,
const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
uint16_t* c, float* c_f32, int* c_counter
) {
// XCD coordinate
int xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
int xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
if (xcd_id >= NUM_XCD_N) return;
int xcd_n = xcd_id;
int warp_id = threadIdx.x / THREADS_PER_WARP;
int lane_index = threadIdx.x % THREADS_PER_WARP;
// Work assignment: k-major for L1 A reuse
// Consecutive work_ids share k_slice -> same A rows -> L1 hits
int work_id = xcd_worker_id * MAX_WORK_PER_CU + warp_id;
if (work_id >= WORK_PER_XCD) return;
int k_slice = work_id / NUM_CU_N;
int n_tile = work_id % NUM_CU_N;
int tile_x = xcd_n * (N / NUM_XCD_N) + n_tile * MFMA_DIM_NM;
int tile_y = 0; // M=16 = one tile
int k_base = k_slice * BLOCK_K;
// Branchless OOB: clamp to last valid tile for B loads
int tile_x_safe = min(tile_x, N_ACTUAL - MFMA_DIM_NM);
int lane_row_nm = lane_index % MFMA_DIM_NM;
int lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM);
// Output info
int result_row_offset = 4 * (lane_index / 16);
int result_column_offset = lane_index % 16;
int col = tile_x + result_column_offset;
if (col >= N_ACTUAL) {
return;
}
// Prefetch buffers: raw A bf16 (16 dwords = 32 bf16), B fp4 (4 dwords), B scale (1 byte)
constexpr uint32_t PREFETCH = 3;
uint32_t a_raw[PREFETCH][MFMA_DIM_NM];
i32x4_t b_raw[PREFETCH];
uint8_t b_sc[PREFETCH];
int b_row = tile_x_safe + lane_row_nm;
const uint16_t* a_base = a + (tile_y + lane_row_nm) * K;
auto loadGlobal = [&](int k_iter) {
__builtin_amdgcn_sched_barrier(0);
int buf = k_iter % PREFETCH;
int k_offset = k_base + k_iter * MFMA_DIM_K;
int k_idx_start = k_offset + lane_col_k;
const uint16_t* a_ptr = a_base + k_idx_start;
#pragma unroll
for (int r = 0; r < 4; r++) {
*(i32x4_t*)&a_raw[buf][r * 4] = *(const i32x4_t*)(a_ptr + r * 8);
}
b_raw[buf] = *(const i32x4_t*)&b_q[b_row * (K / 2) + k_idx_start / 2];
b_sc[buf] = b_scale_sh[read_b_scale_base(b_row, k_idx_start / SCALE_GROUP_SIZE, b_scale_stride)];
__builtin_amdgcn_sched_barrier(0);
};
// Fill pipeline
#pragma unroll
for (int i = 0; i < PREFETCH - 1; i++) {
if (i < K_ITERS) loadGlobal(i);
}
// Prefetch + Quant + MFMA
f32x4_t result = {};
#pragma unroll
for (int k_iter = 0; k_iter < K_ITERS; k_iter++) {
// Prefetch next iteration
int pf = k_iter + PREFETCH - 1;
if (pf < K_ITERS) {
asm volatile("" ::: "memory");
loadGlobal(pf);
}
// Consume from prefetch buffer
int buf = k_iter % PREFETCH;
// Quantize A from buffered bf16
uint8_t a_scale = bf16x32_to_scale_e8m0(a_raw[buf]);
float a_scale_f32 = e8m0_to_f32(a_scale);
u32x8_t a_reg = {};
#pragma unroll
for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
const bf16x2_t* v2 = (const bf16x2_t*)&a_raw[buf][reg_idx * 4];
#ifdef __gfx950__
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[0], a_scale_f32, 0);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[1], a_scale_f32, 1);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[2], a_scale_f32, 2);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[3], a_scale_f32, 3);
#else
a_reg[reg_idx] = __builtin_bit_cast(uint32_t, v2[0]) ^ __builtin_bit_cast(uint32_t, v2[1])
^ __builtin_bit_cast(uint32_t, v2[2]) ^ __builtin_bit_cast(uint32_t, v2[3])
^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
}
uint8_t b_scale = b_sc[buf];
i32x8_t b_reg;
*(i32x4_t*)&b_reg = b_raw[buf];
#ifdef __gfx950__
result = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, b_reg, result, 4, 4, 0, a_scale, 0, b_scale);
#else
(void)a_reg; (void)b_reg; (void)a_scale; (void)b_scale;
#endif
}
// Write output
int row_base = tile_y + result_row_offset;
#pragma unroll
for (int i = 0; i < 4; i++) {
int row = row_base + i;
atomicAdd(&c_f32[row * N_ACTUAL + col], (float)result[i]);
}
// See who's job it is to write.
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
int done = atomicAdd(&c_counter[row_base * N_ACTUAL + col], 1) + 1;
if (done == SPLIT_K) {
#pragma unroll
for (int i = 0; i < 4; i++) {
int row = row_base + i;
c[row * N_ACTUAL + col] = f32_to_bf16(c_f32[row * N_ACTUAL + col]);
c_f32[row * N_ACTUAL + col] = 0.0f;
}
c_counter[row_base * N_ACTUAL + col] = 0;
}
}
constexpr uint32_t c_max_elems = M * N;
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
) {
static float* c_f32 = nullptr;
if (!c_f32) {
assert(hipMalloc(&c_f32, c_max_elems * sizeof(float)) == hipSuccess);
assert(hipMemsetAsync(c_f32, 0, c_max_elems * sizeof(float)) == hipSuccess);
}
static int* c_counter = nullptr;
if (!c_counter) {
assert(hipMalloc(&c_counter, c_max_elems * sizeof(int)) == hipSuccess);
assert(hipMemsetAsync((void*)c_counter, 0, c_max_elems * sizeof(int)) == hipSuccess);
}
if (M == 16 && N == 2112 && K == 7168) {
kernel<<<dim3(NUM_CUS), dim3(NUM_THREADS)>>>(
(const uint16_t*)a, (const uint16_t*)b,
(const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
(uint16_t*)c, c_f32, c_counter
);
} else {
// Do nothing
}
}
"""
CUDA_SRC_SHAPE3 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
__device__ inline uint16_t f32_to_bf16(float v) {
return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32_t* row) {
uint32_t amax_u16x2_packed = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row32[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
// CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
// FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
// denorm_exp = (127-1) + (23-1) + 1 = 149 => denorm_mask = 149 << 23 = 0x4A800000
// val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF (int32 wrapping)
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT = 0x4A800000u;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr uint32_t VAL_TO_ADD = 0xC11FFFFFu;
uint8_t s_bit = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
float abs_v = __builtin_bit_cast(float, abs_bits);
uint8_t result;
if (abs_v >= FP4_MAX_NORMAL) {
// saturate branch
result = 0x7u;
} else if (abs_v < FP4_MIN_NORMAL) {
// denormal branch: float trick to extract low bits
uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
result = (uint8_t)d;
} else {
// normal branch: bias-adjust exponent + round-to-nearest
uint32_t x = abs_bits;
uint32_t m_odd = (x >> 22) & 1u;
x = x + VAL_TO_ADD + m_odd;
result = (uint8_t)(x >> 22);
}
return s_bit | result;
}
// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
return f32_to_fp4_e2m1(v * f32_scale);
}
__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
uint8_t s = (nibble >> 3) & 1;
uint8_t e = (nibble >> 1) & 3;
uint8_t m = nibble & 1;
float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
return s ? -val : val;
}
// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
constexpr int POSITIONS = 256 / WIDTH; // slots per bank cycle
constexpr int MASK = POSITIONS - 1;
int c_group = c / WIDTH;
int c_offset = c % WIDTH; // both compile to shift/mask
int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
return swizzled * WIDTH + c_offset;
}
// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;
// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 8;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;
// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD * NUM_CUS_PER_XCD;
// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;
// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
int tr = n_row / 32, r = n_row % 32;
int tc = k_group / 8, c = k_group % 8;
int base = tr * (32 * b_scale_stride) + tc * 256;
int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
return base + inner;
}
// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
return x;
}
__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
const uint16_t* a,
const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
uint16_t* c
) {
// Timing
constexpr bool TIMEIT = false;
uint64_t _ts[16];
int _ti = 0;
auto TIME = [&]() {
if constexpr (TIMEIT) {
__builtin_amdgcn_sched_barrier(0);
_ts[_ti++] = __builtin_amdgcn_s_memrealtime();
__builtin_amdgcn_sched_barrier(0);
}
};
TIME();
constexpr uint32_t M = 32;
constexpr uint32_t N = 4096;
constexpr uint32_t K = 512;
constexpr uint32_t SPLIT_K = 1;
static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;
constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;
// Get XCD coordinate
uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
uint32_t xcd_n = xcd_id % NUM_XCD_N;
uint32_t xcd_m = xcd_id / NUM_XCD_N;
// Get CU coordinate
if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
return;
}
uint32_t cu_n = xcd_worker_id % NUM_CU_N;
uint32_t cu_m = xcd_worker_id / NUM_CU_N;
// Get warp id (Use SGPR for it)
uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
uint32_t warp_n = warp_id % NUM_WARPS_N;
uint32_t warp_m = warp_id / NUM_WARPS_N;
// Get lane assignment
uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.
// Get tile offset
uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;
// Shared Memory
__shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_a[BLOCK_M][K / 2];
__shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_b[BLOCK_N][K / 2];
auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };
// ======================
// A: Global -> Register
// ======================
TIME();
constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint32_t a_row_idx = tile_m_offset + a_transfer_m;
*(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
}
auto a_process_job_step1 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
a_scale_jobs[a_chunk_job_idx] = a_scale;
__builtin_amdgcn_sched_barrier(0);
};
auto a_process_job_step2 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
u32x4_t a_reg;
#pragma unroll
for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
}
*(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
__builtin_amdgcn_sched_barrier(0);
};
uint32_t a_work_idx = 0;
// ======================
// B: Global -> Register
// ======================
TIME();
constexpr uint32_t B_BYTES_PER_JOB = 16;
constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);
u32x4_t b_jobs[B_NUM_JOBGROUPS];
uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
};
const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;
auto b_process_job_step0 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
uint32_t b_row_idx = tile_n_offset + b_transfer_n;
b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];
uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
void* src = (void*)(b_base + warp_base + lane_index * 16);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dwordx4 %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
#else
#pragma unroll
for (uint32_t sub = 0; sub < 4; sub++) {
uint32_t chunk_base = warp_base + sub * 256;
void* src = (void*)(b_base + chunk_base + lane_index * 4);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
}
#endif
__builtin_amdgcn_sched_barrier(0);
};
auto b_process_job_step1 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
asm volatile("" ::: "memory");
s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
asm volatile("" ::: "memory");
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step0(b_jobgroup_idx);
__builtin_amdgcn_sched_barrier(0);
if (b_jobgroup_idx == 1) {
a_process_job_step1(a_work_idx / 2);
} else if (b_jobgroup_idx == 2) {
a_process_job_step2(a_work_idx / 2);
} else if (b_jobgroup_idx == 3) {
b_process_job_step1(0);
}
__builtin_amdgcn_sched_barrier(0);
}
// ======================
// B: Register -> LDS
// ======================
TIME();
#pragma unroll
for (uint32_t b_jobgroup_idx = 1; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step1(b_jobgroup_idx);
}
// ======================
// A: Register -> Quantize -> LDS
// ======================
TIME();
// NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
// while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
// if (a_work_idx % 2 == 0) {
// a_process_job_step1(a_work_idx / 2);
// } else {
// a_process_job_step2(a_work_idx / 2);
// }
// a_work_idx += 1;
// }
// ======================
// LDS->MFMA
// ======================
__syncthreads(); // Ensure LDS is populated
TIME();
f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
constexpr uint32_t PREFETCH = 2;
uint8_t a_scale[PREFETCH][WARP_TILE_M];
u32x8_t a_reg[PREFETCH][WARP_TILE_M];
uint8_t b_scale[PREFETCH][WARP_TILE_N];
u32x8_t b_reg[PREFETCH][WARP_TILE_N];
auto load_lds = [&](uint32_t k_iter) {
__builtin_amdgcn_sched_barrier(0);
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
*(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
}
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_N; i++) {
uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
+ (k_idx_start / 32) * 256
+ lane_row_nm * 16;
*(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
}
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
load_lds(i);
}
#pragma unroll
for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
if (k_prefetch_iter < K_ITERS) {
load_lds(k_prefetch_iter);
}
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
f32x4_t mix;
for (uint32_t r = 0; r < 4; r++)
mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
result[i][j] += mix;
#endif
}
}
}
// ======================
// Reg->Write to Global
// ======================
TIME();
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for (uint32_t j = 0; j < WARP_TILE_N; j++) {
uint32_t result_row_offset = 4 * (lane_index / 16);
uint32_t result_column_offset = lane_index % 16;
#pragma unroll
for (uint32_t f = 0; f < 4; f++) {
uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
}
}
}
TIME();
if constexpr (TIMEIT) {
__shared__ uint32_t _run_idx;
if (threadIdx.x == 0) {
_run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
}
__syncthreads();
uint32_t h = hash(_run_idx);
if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
for (int i = 1; i < _ti; i++) {
printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
}
printf("\n");
}
}
}
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
) {
if (M == 32 && N == 4096 && K == 512) {
kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
(const uint16_t*)a,
(const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
(uint16_t*)c
);
} else {
// No impl
}
}
"""
CUDA_SRC_SHAPE4 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
__device__ inline uint16_t f32_to_bf16(float v) {
return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32_t* row) {
uint32_t amax_u16x2_packed = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row32[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
// CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
// FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
// denorm_exp = (127-1) + (23-1) + 1 = 149 => denorm_mask = 149 << 23 = 0x4A800000
// val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF (int32 wrapping)
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT = 0x4A800000u;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr uint32_t VAL_TO_ADD = 0xC11FFFFFu;
uint8_t s_bit = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
float abs_v = __builtin_bit_cast(float, abs_bits);
uint8_t result;
if (abs_v >= FP4_MAX_NORMAL) {
// saturate branch
result = 0x7u;
} else if (abs_v < FP4_MIN_NORMAL) {
// denormal branch: float trick to extract low bits
uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
result = (uint8_t)d;
} else {
// normal branch: bias-adjust exponent + round-to-nearest
uint32_t x = abs_bits;
uint32_t m_odd = (x >> 22) & 1u;
x = x + VAL_TO_ADD + m_odd;
result = (uint8_t)(x >> 22);
}
return s_bit | result;
}
// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
return f32_to_fp4_e2m1(v * f32_scale);
}
__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
uint8_t s = (nibble >> 3) & 1;
uint8_t e = (nibble >> 1) & 3;
uint8_t m = nibble & 1;
float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
return s ? -val : val;
}
// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
constexpr int POSITIONS = 256 / WIDTH; // slots per bank cycle
constexpr int MASK = POSITIONS - 1;
int c_group = c / WIDTH;
int c_offset = c % WIDTH; // both compile to shift/mask
int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
return swizzled * WIDTH + c_offset;
}
// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;
// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 9;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 5;
#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;
// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD_M * NUM_XCD_N * NUM_CUS_PER_XCD;
// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;
// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
int tr = n_row / 32, r = n_row % 32;
int tc = k_group / 8, c = k_group % 8;
int base = tr * (32 * b_scale_stride) + tc * 256;
int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
return base + inner;
}
// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
return x;
}
__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
const uint16_t* a,
const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
uint16_t* c
) {
// Timing
constexpr bool TIMEIT = false;
uint64_t _ts[16];
int _ti = 0;
auto TIME = [&]() {
if constexpr (TIMEIT) {
__builtin_amdgcn_sched_barrier(0);
_ts[_ti++] = __builtin_amdgcn_s_memrealtime();
__builtin_amdgcn_sched_barrier(0);
}
};
TIME();
constexpr uint32_t M = 32;
constexpr uint32_t N = 2880;
constexpr uint32_t K = 512;
constexpr uint32_t SPLIT_K = 1;
static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;
constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;
// Get XCD coordinate
uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
uint32_t xcd_n = xcd_id % NUM_XCD_N;
uint32_t xcd_m = xcd_id / NUM_XCD_N;
// Get CU coordinate
if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
return;
}
uint32_t cu_n = xcd_worker_id % NUM_CU_N;
uint32_t cu_m = xcd_worker_id / NUM_CU_N;
// Get warp id (Use SGPR for it)
uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
uint32_t warp_n = warp_id % NUM_WARPS_N;
uint32_t warp_m = warp_id / NUM_WARPS_N;
// Get lane assignment
uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.
// Get tile offset
uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;
// Shared Memory
__shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_a[BLOCK_M][K / 2];
__shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_b[BLOCK_N][K / 2];
auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };
// ======================
// A: Global -> Register
// ======================
TIME();
constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint32_t a_row_idx = tile_m_offset + a_transfer_m;
*(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
}
auto a_process_job_step1 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
a_scale_jobs[a_chunk_job_idx] = a_scale;
__builtin_amdgcn_sched_barrier(0);
};
auto a_process_job_step2 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
u32x4_t a_reg;
#pragma unroll
for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
}
*(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
__builtin_amdgcn_sched_barrier(0);
};
uint32_t a_work_idx = 0;
// ======================
// B: Global -> Register
// ======================
TIME();
constexpr uint32_t B_BYTES_PER_JOB = 16;
constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);
u32x4_t b_jobs[B_NUM_JOBGROUPS];
uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
};
const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;
auto b_process_job_step0 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
uint32_t b_row_idx = tile_n_offset + b_transfer_n;
b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];
uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
void* src = (void*)(b_base + warp_base + lane_index * 16);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dwordx4 %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
#else
#pragma unroll
for (uint32_t sub = 0; sub < 4; sub++) {
uint32_t chunk_base = warp_base + sub * 256;
void* src = (void*)(b_base + chunk_base + lane_index * 4);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
}
#endif
__builtin_amdgcn_sched_barrier(0);
};
auto b_process_job_step1 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
asm volatile("" ::: "memory");
s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
asm volatile("" ::: "memory");
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step0(b_jobgroup_idx);
__builtin_amdgcn_sched_barrier(0);
if (b_jobgroup_idx == 1) {
a_process_job_step1(a_work_idx / 2);
} else if (b_jobgroup_idx == 2) {
a_process_job_step2(a_work_idx / 2);
} else if (b_jobgroup_idx == 3) {
b_process_job_step1(0);
}
__builtin_amdgcn_sched_barrier(0);
}
// ======================
// B: Register -> LDS
// ======================
TIME();
#pragma unroll
for (uint32_t b_jobgroup_idx = 1; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step1(b_jobgroup_idx);
}
// ======================
// A: Register -> Quantize -> LDS
// ======================
TIME();
// NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
// while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
// if (a_work_idx % 2 == 0) {
// a_process_job_step1(a_work_idx / 2);
// } else {
// a_process_job_step2(a_work_idx / 2);
// }
// a_work_idx += 1;
// }
// ======================
// LDS->MFMA
// ======================
__syncthreads(); // Ensure LDS is populated
TIME();
f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
constexpr uint32_t PREFETCH = 2;
uint8_t a_scale[PREFETCH][WARP_TILE_M];
u32x8_t a_reg[PREFETCH][WARP_TILE_M];
uint8_t b_scale[PREFETCH][WARP_TILE_N];
u32x8_t b_reg[PREFETCH][WARP_TILE_N];
auto load_lds = [&](uint32_t k_iter) {
__builtin_amdgcn_sched_barrier(0);
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
*(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
}
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_N; i++) {
uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
+ (k_idx_start / 32) * 256
+ lane_row_nm * 16;
*(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
}
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
load_lds(i);
}
#pragma unroll
for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
if (k_prefetch_iter < K_ITERS) {
load_lds(k_prefetch_iter);
}
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
f32x4_t mix;
for (uint32_t r = 0; r < 4; r++)
mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
result[i][j] += mix;
#endif
}
}
}
// ======================
// Reg->Write to Global
// ======================
TIME();
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for (uint32_t j = 0; j < WARP_TILE_N; j++) {
uint32_t result_row_offset = 4 * (lane_index / 16);
uint32_t result_column_offset = lane_index % 16;
#pragma unroll
for (uint32_t f = 0; f < 4; f++) {
uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
}
}
}
TIME();
if constexpr (TIMEIT) {
__shared__ uint32_t _run_idx;
if (threadIdx.x == 0) {
_run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
}
__syncthreads();
uint32_t h = hash(_run_idx);
if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
for (int i = 1; i < _ti; i++) {
printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
}
printf("\n");
}
}
}
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
) {
if (M == 32 && N == 2880 && K == 512) {
kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
(const uint16_t*)a,
(const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
(uint16_t*)c
);
} else {
// No impl
}
}
"""
CUDA_SRC_SHAPE5 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
__device__ inline uint16_t f32_to_bf16(float v) {
return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32_t* row) {
uint32_t amax_u16x2_packed = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row32[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
// CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
// FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
// denorm_exp = (127-1) + (23-1) + 1 = 149 => denorm_mask = 149 << 23 = 0x4A800000
// val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF (int32 wrapping)
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT = 0x4A800000u;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr uint32_t VAL_TO_ADD = 0xC11FFFFFu;
uint8_t s_bit = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
float abs_v = __builtin_bit_cast(float, abs_bits);
uint8_t result;
if (abs_v >= FP4_MAX_NORMAL) {
// saturate branch
result = 0x7u;
} else if (abs_v < FP4_MIN_NORMAL) {
// denormal branch: float trick to extract low bits
uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
result = (uint8_t)d;
} else {
// normal branch: bias-adjust exponent + round-to-nearest
uint32_t x = abs_bits;
uint32_t m_odd = (x >> 22) & 1u;
x = x + VAL_TO_ADD + m_odd;
result = (uint8_t)(x >> 22);
}
return s_bit | result;
}
// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
return f32_to_fp4_e2m1(v * f32_scale);
}
__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
uint8_t s = (nibble >> 3) & 1;
uint8_t e = (nibble >> 1) & 3;
uint8_t m = nibble & 1;
float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
return s ? -val : val;
}
// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
constexpr int POSITIONS = 256 / WIDTH; // slots per bank cycle
constexpr int MASK = POSITIONS - 1;
int c_group = c / WIDTH;
int c_offset = c % WIDTH; // both compile to shift/mask
int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
return swizzled * WIDTH + c_offset;
}
// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;
// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 2;
constexpr uint32_t NUM_WARPS_M = 2;
constexpr uint32_t NUM_WARPS_N = 2;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;
// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD * NUM_CUS_PER_XCD;
// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;
// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
int tr = n_row / 32, r = n_row % 32;
int tc = k_group / 8, c = k_group % 8;
int base = tr * (32 * b_scale_stride) + tc * 256;
int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
return base + inner;
}
// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
return x;
}
__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
const uint16_t* a,
const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
uint16_t* c
) {
// Timing
constexpr bool TIMEIT = false;
uint64_t _ts[16];
int _ti = 0;
auto TIME = [&]() {
if constexpr (TIMEIT) {
__builtin_amdgcn_sched_barrier(0);
_ts[_ti++] = __builtin_amdgcn_s_memrealtime();
__builtin_amdgcn_sched_barrier(0);
}
};
TIME();
constexpr uint32_t M = 64;
constexpr uint32_t N = 7168;
constexpr uint32_t K = 2048;
constexpr uint32_t SPLIT_K = 1;
static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;
constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;
// Get XCD coordinate
uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
uint32_t xcd_n = xcd_id % NUM_XCD_N;
uint32_t xcd_m = xcd_id / NUM_XCD_N;
// Get CU coordinate
if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
return;
}
uint32_t cu_n = xcd_worker_id % NUM_CU_N;
uint32_t cu_m = xcd_worker_id / NUM_CU_N;
// Get warp id (Use SGPR for it)
uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
uint32_t warp_n = warp_id % NUM_WARPS_N;
uint32_t warp_m = warp_id / NUM_WARPS_N;
// Get lane assignment
uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.
// Get tile offset
uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;
// Shared Memory
__shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_a[BLOCK_M][K / 2];
__shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_b[BLOCK_N][K / 2];
auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };
// ======================
// A: Global -> Register
// ======================
TIME();
constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint32_t a_row_idx = tile_m_offset + a_transfer_m;
*(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
}
auto a_process_job_step1 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
a_scale_jobs[a_chunk_job_idx] = a_scale;
__builtin_amdgcn_sched_barrier(0);
};
auto a_process_job_step2 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
u32x4_t a_reg;
#pragma unroll
for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
}
*(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
__builtin_amdgcn_sched_barrier(0);
};
uint32_t a_work_idx = 0;
// ======================
// B: Global -> Register
// ======================
TIME();
constexpr uint32_t B_BYTES_PER_JOB = 16;
constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);
u32x4_t b_jobs[B_NUM_JOBGROUPS];
uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
};
const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;
auto b_process_job_step0 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
uint32_t b_row_idx = tile_n_offset + b_transfer_n;
b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];
uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
void* src = (void*)(b_base + warp_base + lane_index * 16);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dwordx4 %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
#else
#pragma unroll
for (uint32_t sub = 0; sub < 4; sub++) {
uint32_t chunk_base = warp_base + sub * 256;
void* src = (void*)(b_base + chunk_base + lane_index * 4);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
}
#endif
__builtin_amdgcn_sched_barrier(0);
};
auto b_process_job_step1 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
asm volatile("" ::: "memory");
s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
asm volatile("" ::: "memory");
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step0(b_jobgroup_idx);
__builtin_amdgcn_sched_barrier(0);
if (a_work_idx % 2 == 0) {
a_process_job_step1(a_work_idx / 2);
} else {
a_process_job_step2(a_work_idx / 2);
}
a_work_idx += 1;
__builtin_amdgcn_sched_barrier(0);
}
// ======================
// B: Register -> LDS
// ======================
TIME();
#pragma unroll
for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step1(b_jobgroup_idx);
}
// ======================
// A: Register -> Quantize -> LDS
// ======================
TIME();
// NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
// while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
// if (a_work_idx % 2 == 0) {
// a_process_job_step1(a_work_idx / 2);
// } else {
// a_process_job_step2(a_work_idx / 2);
// }
// a_work_idx += 1;
// }
// ======================
// LDS->MFMA
// ======================
__syncthreads(); // Ensure LDS is populated
TIME();
f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
constexpr uint32_t PREFETCH = 2;
uint8_t a_scale[PREFETCH][WARP_TILE_M];
u32x8_t a_reg[PREFETCH][WARP_TILE_M];
uint8_t b_scale[PREFETCH][WARP_TILE_N];
u32x8_t b_reg[PREFETCH][WARP_TILE_N];
auto load_lds = [&](uint32_t k_iter) {
__builtin_amdgcn_sched_barrier(0);
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
*(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
}
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_N; i++) {
uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
+ (k_idx_start / 32) * 256
+ lane_row_nm * 16;
*(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
}
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
load_lds(i);
}
#pragma unroll
for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
if (k_prefetch_iter < K_ITERS) {
load_lds(k_prefetch_iter);
}
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
f32x4_t mix;
for (uint32_t r = 0; r < 4; r++)
mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
result[i][j] += mix;
#endif
}
}
}
// ======================
// Reg->Write to Global
// ======================
TIME();
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for (uint32_t j = 0; j < WARP_TILE_N; j++) {
uint32_t result_row_offset = 4 * (lane_index / 16);
uint32_t result_column_offset = lane_index % 16;
#pragma unroll
for (uint32_t f = 0; f < 4; f++) {
uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
}
}
}
TIME();
if constexpr (TIMEIT) {
__shared__ uint32_t _run_idx;
if (threadIdx.x == 0) {
_run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
}
__syncthreads();
uint32_t h = hash(_run_idx);
if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
for (int i = 1; i < _ti; i++) {
printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
}
printf("\n");
}
}
}
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
) {
if (M == 64 && N == 7168 && K == 2048) {
kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
(const uint16_t*)a,
(const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
(uint16_t*)c
);
} else {
// No impl
}
}
"""
CUDA_SRC_SHAPE6 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
__device__ inline uint16_t f32_to_bf16(float v) {
return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32_t* row) {
uint32_t amax_u16x2_packed = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row32[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
// CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
// FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
// denorm_exp = (127-1) + (23-1) + 1 = 149 => denorm_mask = 149 << 23 = 0x4A800000
// val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF (int32 wrapping)
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT = 0x4A800000u;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr uint32_t VAL_TO_ADD = 0xC11FFFFFu;
uint8_t s_bit = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
float abs_v = __builtin_bit_cast(float, abs_bits);
uint8_t result;
if (abs_v >= FP4_MAX_NORMAL) {
// saturate branch
result = 0x7u;
} else if (abs_v < FP4_MIN_NORMAL) {
// denormal branch: float trick to extract low bits
uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
result = (uint8_t)d;
} else {
// normal branch: bias-adjust exponent + round-to-nearest
uint32_t x = abs_bits;
uint32_t m_odd = (x >> 22) & 1u;
x = x + VAL_TO_ADD + m_odd;
result = (uint8_t)(x >> 22);
}
return s_bit | result;
}
// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
return f32_to_fp4_e2m1(v * f32_scale);
}
__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
uint8_t s = (nibble >> 3) & 1;
uint8_t e = (nibble >> 1) & 3;
uint8_t m = nibble & 1;
float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
return s ? -val : val;
}
// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
constexpr int POSITIONS = 256 / WIDTH; // slots per bank cycle
constexpr int MASK = POSITIONS - 1;
int c_group = c / WIDTH;
int c_offset = c % WIDTH; // both compile to shift/mask
int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
return swizzled * WIDTH + c_offset;
}
// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;
// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 3;
constexpr uint32_t NUM_WARPS_M = 2;
constexpr uint32_t NUM_WARPS_N = 2;
constexpr uint32_t NUM_CU_M = 4;
constexpr uint32_t NUM_CU_N = 8;
constexpr uint32_t NUM_XCD_M = 2;
constexpr uint32_t NUM_XCD_N = 4;
#else
constexpr uint32_t NUM_CUS_PER_XCD = 38 * 2;
// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 3;
constexpr uint32_t NUM_WARPS_M = 2;
constexpr uint32_t NUM_WARPS_N = 1;
constexpr uint32_t NUM_CU_M = 4;
constexpr uint32_t NUM_CU_N = 16;
constexpr uint32_t NUM_XCD_M = 2;
constexpr uint32_t NUM_XCD_N = 4;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD * NUM_CUS_PER_XCD;
// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;
// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
int tr = n_row / 32, r = n_row % 32;
int tc = k_group / 8, c = k_group % 8;
int base = tr * (32 * b_scale_stride) + tc * 256;
int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
return base + inner;
}
// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
x *= 0x45d9f3b;
x ^= x >> 16;
return x;
}
__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
const uint16_t* a,
const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
uint16_t* c
) {
// Timing
constexpr bool TIMEIT = false;
uint64_t _ts[16];
int _ti = 0;
auto TIME = [&]() {
if constexpr (TIMEIT) {
__builtin_amdgcn_sched_barrier(0);
_ts[_ti++] = __builtin_amdgcn_s_memrealtime();
__builtin_amdgcn_sched_barrier(0);
}
};
TIME();
constexpr uint32_t M = 256;
constexpr uint32_t N = 3072;
constexpr uint32_t K = 1536;
constexpr uint32_t SPLIT_K = 1;
static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;
constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;
// Get XCD coordinate
uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
uint32_t xcd_n = xcd_id % NUM_XCD_N;
uint32_t xcd_m = xcd_id / NUM_XCD_N;
// Get CU coordinate
if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
return;
}
uint32_t cu_n = xcd_worker_id % NUM_CU_N;
uint32_t cu_m = xcd_worker_id / NUM_CU_N;
// Get warp id (Use SGPR for it)
uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
uint32_t warp_n = warp_id % NUM_WARPS_N;
uint32_t warp_m = warp_id / NUM_WARPS_N;
// Get lane assignment
uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.
// Get tile offset
uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;
// Shared Memory
__shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_a[BLOCK_M][K / 2];
__shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
__shared__ uint8_t s_b[BLOCK_N][K / 2];
auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };
// ======================
// A: Global -> Register
// ======================
TIME();
constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint32_t a_row_idx = tile_m_offset + a_transfer_m;
*(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
}
auto a_process_job_step1 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
a_scale_jobs[a_chunk_job_idx] = a_scale;
__builtin_amdgcn_sched_barrier(0);
};
auto a_process_job_step2 = [&](int a_chunk_job_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
u32x4_t a_reg;
#pragma unroll
for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
}
*(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
__builtin_amdgcn_sched_barrier(0);
};
uint32_t a_work_idx = 0;
// ======================
// B: Global -> Register
// ======================
TIME();
constexpr uint32_t B_BYTES_PER_JOB = 16;
constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);
u32x4_t b_jobs[B_NUM_JOBGROUPS];
uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
};
const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;
auto b_process_job_step0 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
uint32_t b_row_idx = tile_n_offset + b_transfer_n;
b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];
uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
void* src = (void*)(b_base + warp_base + lane_index * 16);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dwordx4 %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
#else
#pragma unroll
for (uint32_t sub = 0; sub < 4; sub++) {
uint32_t chunk_base = warp_base + sub * 256;
void* src = (void*)(b_base + chunk_base + lane_index * 4);
uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_off), "v"(src)
: "memory", "m0"
);
}
#endif
__builtin_amdgcn_sched_barrier(0);
};
auto b_process_job_step1 = [&](int b_jobgroup_idx) {
__builtin_amdgcn_sched_barrier(0);
uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
uint32_t b_transfer_n = b_job_offset / (K / 2);
uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
b_process_job_step0(b_jobgroup_idx);
if (b_jobgroup_idx % 3 == 2) {
if (a_work_idx % 2 == 0) {
a_process_job_step1(a_work_idx / 2);
} else {
a_process_job_step2(a_work_idx / 2);
}
a_work_idx += 1;
}
}
// ======================
// B: Register -> LDS
// ======================
TIME();
#pragma unroll
for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
if (b_jobgroup_idx % 3 == 2) {
if (a_work_idx % 2 == 0) {
a_process_job_step1(a_work_idx / 2);
} else {
a_process_job_step2(a_work_idx / 2);
}
a_work_idx += 1;
}
b_process_job_step1(b_jobgroup_idx);
}
// ======================
// A: Register -> Quantize -> LDS
// ======================
TIME();
// NOTE: This task was interwoven above. This while loop should not be hit.
#pragma unroll
while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
if (a_work_idx % 2 == 0) {
a_process_job_step1(a_work_idx / 2);
} else {
a_process_job_step2(a_work_idx / 2);
}
a_work_idx += 1;
}
// ======================
// LDS->MFMA
// ======================
__syncthreads(); // Ensure LDS is populated
TIME();
f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
constexpr uint32_t PREFETCH = 2;
uint8_t a_scale[PREFETCH][WARP_TILE_M];
u32x8_t a_reg[PREFETCH][WARP_TILE_M];
uint8_t b_scale[PREFETCH][WARP_TILE_N];
u32x8_t b_reg[PREFETCH][WARP_TILE_N];
auto load_lds = [&](uint32_t k_iter) {
__builtin_amdgcn_sched_barrier(0);
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
*(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
}
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_N; i++) {
uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
+ (k_idx_start / 32) * 256
+ lane_row_nm * 16;
*(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
}
__builtin_amdgcn_sched_barrier(0);
};
#pragma unroll
for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
load_lds(i);
}
#pragma unroll
for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
if (k_prefetch_iter < K_ITERS) {
load_lds(k_prefetch_iter);
}
uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
uint32_t k_idx_start = k_offset + lane_col_k;
uint32_t buf_idx = k_iter % PREFETCH;
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
f32x4_t mix;
for (uint32_t r = 0; r < 4; r++)
mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
result[i][j] += mix;
#endif
}
}
}
// ======================
// Reg->Write to Global
// ======================
TIME();
#pragma unroll
for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
for (uint32_t j = 0; j < WARP_TILE_N; j++) {
uint32_t result_row_offset = 4 * (lane_index / 16);
uint32_t result_column_offset = lane_index % 16;
#pragma unroll
for (uint32_t f = 0; f < 4; f++) {
uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
}
}
}
TIME();
if constexpr (TIMEIT) {
__shared__ uint32_t _run_idx;
if (threadIdx.x == 0) {
_run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
}
__syncthreads();
uint32_t h = hash(_run_idx);
if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
for (int i = 1; i < _ti; i++) {
printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
}
printf("\n");
}
}
}
void entry(
const uintptr_t a, const uintptr_t b,
const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
uintptr_t c,
int M, int N, int K
) {
if (M == 256 && N == 3072 && K == 1536) {
kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
(const uint16_t*)a,
(const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
(uint16_t*)c
);
} else {
// No impl
}
}
"""
class CompiledModule:
M: int
N: int
K: int
module: Any
out: torch.Tensor
def __init__(self, M, N, K):
CUDA_PRELUDE = """
#ifndef __gfx950__
#define __gfx950__
#endif
"""
cflags = ["--offload-arch=gfx950", "-std=c++20", "-O3", "-ffast-math", "-march=native", "-funroll-loops", "-fomit-frame-pointer"]
cflags.extend([f"-DM_DIM={M}", f"-DN_DIM={N}", f"-DK_DIM={K}"])
self.M = M
self.N = N
self.K = K
self.out = torch.empty((M, N), dtype=torch.bfloat16, device='cuda')
cuda_src = None
if M == 4 and N == 2880 and K == 512:
cuda_src = CUDA_SRC_SHAPE1
elif M == 16 and N == 2112 and K == 7168:
cuda_src = CUDA_SRC_SHAPE2
elif M == 32 and N == 4096 and K == 512:
cuda_src = CUDA_SRC_SHAPE3
elif M == 32 and N == 2880 and K == 512:
cuda_src = CUDA_SRC_SHAPE4
elif M == 64 and N == 7168 and K == 2048:
cuda_src = CUDA_SRC_SHAPE5
elif M == 256 and N == 3072 and K == 1536:
cuda_src = CUDA_SRC_SHAPE6
else:
cuda_src = CUDA_SRC_SHAPE_GENERAL
self.module = load_inline(
name=f"solution_{M}_{N}_{K}",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[CUDA_PRELUDE + cuda_src],
functions=['entry'],
with_cuda=True,
verbose=False,
extra_cuda_cflags=cflags,
extra_cflags=cflags,
)
def inference(self, A, B, B_q, B_shuffle, B_scale_sh):
self.module.entry(
A.data_ptr(), B.data_ptr(),
B_q.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
self.out.data_ptr(),
self.M, self.N, self.K,
)
return self.out
_compiled_modules: dict[tuple[int, int, int], CompiledModule] = {}
def custom_kernel(data: input_t) -> output_t:
global _compiled_modules
A, B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
N, _ = B.shape
key = (M, N, K)
if key not in _compiled_modules:
_compiled_modules[key] = CompiledModule(M, N, K)
return _compiled_modules[key].inference(A, B, B_q, B_shuffle, B_scale_sh)
scrolls · 3129 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