submission 596238
RyanWillie · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2034 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-596238?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:60e41574fc9a0817ff8500e26e5fb9190a7ebfb3dc460ae9f38cabf6de3b025d
license declaredunknown
license concludedunknown
authorsRyanWillie
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"Fused MXFP4 quant+GEMM");shared-memory
__shared__ uint8_t lds_a[8192]; // 2 x (16 rows * 128 BF16 * 2 bytes)split-k
void mxfp4_fused_gemm_m16_splitk(Kernel source
submission.py2034 lines
# Zeus Variant: v177_k512_2wave_only
# Experiment: #177
# Technique: Cherry-pick K=512 2-wave unrolled kernel for S3/S4 (M=32) from v174
# Hypothesis: v176 showed K=512 2-wave unrolled kernel improved S3 by -3.6% and S4 by -3.9%.
# S1 (M=4) 1-wave regressed, S2 (M=16) 4-wave multiwave regressed badly.
# This cherry-picks ONLY the winning change: K=512 2-wave for M>16 shapes.
# Base: submissions/v125_ds_bpermute_combined.py
"""
Hybrid HIP Quant + AITER ASM GEMM + Multi-Wave Fused M=16 Kernel + K=512 2-Wave:
Architecture:
Shape 1 (M=4, K=512): Original crossbuf kernel (unchanged from v125)
Shape 2 (M=16, N=2112, K=7168): Multi-wave fused 8-wave 16x16 MFMA kernel
Uses ds_bpermute_b32 for A-data lane remapping instead of LDS.
Shape 3/4 (M=32, K=512): NEW K=512 2-wave unrolled kernel (from v174)
- Fully unrolled 4 K_STEP=128 iterations for K=512
- Cross-buffer vmcnt(6) pipeline overlapping loads with compute
- 2 waves (128 threads), TILE_M_GRID=32
Shape 5/6 (K>=1024): Custom HIP quant kernel -> AITER ASM GEMM
"""
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from task import input_t, output_t
import torch
# ═══════════════════════════════════════════════════════════════
# HIP C++ Source
# ═══════════════════════════════════════════════════════════════
_HIP_SOURCE = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <torch/extension.h>
// ═══════════════════════════════════════════════════════════════
// Type definitions for MFMA and CVT
// ═══════════════════════════════════════════════════════════════
typedef float __attribute__((ext_vector_type(4))) floatx4;
typedef int __attribute__((ext_vector_type(4))) intx4;
typedef __bf16 __attribute__((ext_vector_type(2))) bf16x2_t;
// ═══════════════════════════════════════════════════════════════
// Constants
// ═══════════════════════════════════════════════════════════════
#define TILE_M_PER_WAVE 16
#define TILE_M_GRID 32 // 2 waves × 16 rows each
#define TILE_N 16
#define K_STEP 128
#define SCALE_GROUP 32
#define NUM_K_GROUPS 4 // K_STEP / SCALE_GROUP = 128 / 32
#define WAVESIZE 64
#define BLOCK_SIZE 128 // 2 waves
// ═══════════════════════════════════════════════════════════════
// Device helpers
// ═══════════════════════════════════════════════════════════════
__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
uint32_t bits = ((uint32_t)x) << 16;
float f;
__builtin_memcpy(&f, &bits, 4);
return f;
}
__device__ __forceinline__ uint16_t f32_to_bf16(float f) {
uint32_t bits;
__builtin_memcpy(&bits, &f, 4);
uint32_t lsb = (bits >> 16) & 1;
uint32_t rounding_bias = 0x7FFF + lsb;
bits += rounding_bias;
return (uint16_t)(bits >> 16);
}
__device__ __forceinline__ uint32_t float_as_uint(float f) {
uint32_t u;
__builtin_memcpy(&u, &f, 4);
return u;
}
__device__ __forceinline__ float uint_as_float(uint32_t u) {
float f;
__builtin_memcpy(&f, &u, 4);
return f;
}
// ═══════════════════════════════════════════════════════════════
// Standalone MXFP4 Quantization Kernel (from v91)
// ═══════════════════════════════════════════════════════════════
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void mxfp4_quant_kernel(
const uint16_t* __restrict__ A,
uint8_t* __restrict__ A_q,
uint8_t* __restrict__ A_scale_sh,
int M, int K,
int scaleN,
int scaleN_valid
) {
int tid = blockIdx.x * 256 + threadIdx.x;
int m = tid / scaleN_valid;
int s = tid % scaleN_valid;
if (m >= M) return;
int k_start = s * SCALE_GROUP;
int a_words[16];
if (k_start + SCALE_GROUP <= K) {
const intx4* a_vec = reinterpret_cast<const intx4*>(A + m * K + k_start);
intx4 v0 = a_vec[0], v1 = a_vec[1], v2 = a_vec[2], v3 = a_vec[3];
a_words[0] = v0[0]; a_words[1] = v0[1]; a_words[2] = v0[2]; a_words[3] = v0[3];
a_words[4] = v1[0]; a_words[5] = v1[1]; a_words[6] = v1[2]; a_words[7] = v1[3];
a_words[8] = v2[0]; a_words[9] = v2[1]; a_words[10] = v2[2]; a_words[11] = v2[3];
a_words[12] = v3[0]; a_words[13] = v3[1]; a_words[14] = v3[2]; a_words[15] = v3[3];
} else {
for (int i = 0; i < 16; i++) a_words[i] = 0;
for (int i = 0; i < SCALE_GROUP && k_start + i < K; i++) {
uint16_t val = A[m * K + k_start + i];
int w = i / 2;
if (i % 2 == 0)
a_words[w] = (int)((uint32_t)val);
else
a_words[w] |= (int)(((uint32_t)val) << 16);
}
}
float av[32];
#pragma unroll
for (int w = 0; w < 16; w++) {
uint32_t abs_word = ((uint32_t)a_words[w]) & 0x7FFF7FFFu;
av[2*w] = bf16_to_f32((uint16_t)(abs_word & 0xFFFF));
av[2*w+1] = bf16_to_f32((uint16_t)(abs_word >> 16));
}
float L1_0, L1_1, L1_2, L1_3, L1_4, L1_5, L1_6, L1_7, L1_8, L1_9;
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_0) : "v"(av[0]), "v"(av[1]), "v"(av[2]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_1) : "v"(av[3]), "v"(av[4]), "v"(av[5]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_2) : "v"(av[6]), "v"(av[7]), "v"(av[8]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_3) : "v"(av[9]), "v"(av[10]), "v"(av[11]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_4) : "v"(av[12]), "v"(av[13]), "v"(av[14]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_5) : "v"(av[15]), "v"(av[16]), "v"(av[17]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_6) : "v"(av[18]), "v"(av[19]), "v"(av[20]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_7) : "v"(av[21]), "v"(av[22]), "v"(av[23]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_8) : "v"(av[24]), "v"(av[25]), "v"(av[26]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_9) : "v"(av[27]), "v"(av[28]), "v"(av[29]));
float L2_0, L2_1, L2_2, L2_3;
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_0) : "v"(L1_0), "v"(L1_1), "v"(L1_2));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_1) : "v"(L1_3), "v"(L1_4), "v"(L1_5));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_2) : "v"(L1_6), "v"(L1_7), "v"(L1_8));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_3) : "v"(L1_9), "v"(av[30]), "v"(av[31]));
float L3_0;
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L3_0) : "v"(L2_0), "v"(L2_1), "v"(L2_2));
float amax_val;
asm volatile("v_max_f32 %0, %1, %2" : "=v"(amax_val) : "v"(L3_0), "v"(L2_3));
uint32_t amax_bits = float_as_uint(amax_val);
amax_bits = ((amax_bits + 0x200000u) & 0xFF800000u);
int scale_unbiased;
if (amax_bits == 0) {
scale_unbiased = -127;
} else {
int exp_biased = (int)((amax_bits >> 23) & 0xFF);
scale_unbiased = exp_biased - 127 - 2;
if (scale_unbiased < -127) scale_unbiased = -127;
if (scale_unbiased > 127) scale_unbiased = 127;
}
uint8_t a_e8m0 = (uint8_t)(scale_unbiased + 127);
float cvt_scale = uint_as_float(((uint32_t)a_e8m0) << 23);
unsigned int dst0 = 0, dst1 = 0, dst2 = 0, dst3 = 0;
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[0]), cvt_scale, 0);
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[1]), cvt_scale, 1);
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[2]), cvt_scale, 2);
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[3]), cvt_scale, 3);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[4]), cvt_scale, 0);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[5]), cvt_scale, 1);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[6]), cvt_scale, 2);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[7]), cvt_scale, 3);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[8]), cvt_scale, 0);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[9]), cvt_scale, 1);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[10]), cvt_scale, 2);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[11]), cvt_scale, 3);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[12]), cvt_scale, 0);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[13]), cvt_scale, 1);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[14]), cvt_scale, 2);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[15]), cvt_scale, 3);
uint8_t* out_ptr = A_q + m * (K / 2) + k_start / 2;
reinterpret_cast<uint32_t*>(out_ptr)[0] = dst0;
reinterpret_cast<uint32_t*>(out_ptr)[1] = dst1;
reinterpret_cast<uint32_t*>(out_ptr)[2] = dst2;
reinterpret_cast<uint32_t*>(out_ptr)[3] = dst3;
int m_mod32 = m & 31;
int m_div32 = m >> 5;
int s_mod8 = s & 7;
int s_div8 = s >> 3;
int sh_idx = (m_mod32 >> 4)
+ ((s_mod8 >> 2) << 1)
+ ((m_mod32 & 15) << 2)
+ ((s_mod8 & 3) << 6)
+ (s_div8 << 8)
+ m_div32 * (32 * scaleN);
A_scale_sh[sh_idx] = a_e8m0;
}
void launch_mxfp4_quant(
torch::Tensor A,
torch::Tensor A_q,
torch::Tensor A_scale_sh,
int M, int K,
int scaleN,
int scaleN_valid
) {
TORCH_CHECK(A.is_cuda() && A_q.is_cuda() && A_scale_sh.is_cuda());
int total_groups = M * scaleN_valid;
int threads = 256;
int blocks = (total_groups + threads - 1) / threads;
mxfp4_quant_kernel<<<blocks, threads>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
A_q.data_ptr<uint8_t>(),
A_scale_sh.data_ptr<uint8_t>(),
M, K, scaleN, scaleN_valid
);
}
// ═══════════════════════════════════════════════════════════════
// LOAD macros and helpers for fused kernel (same as v91)
// ═══════════════════════════════════════════════════════════════
#define LOAD_A_DATA(a_tmp0, a_tmp1, a_tmp2, a_tmp3, a_path_var, \
A_ptr, a_row_global, a_k_start, stride_a, M, K) \
do { \
if ((a_row_global) < M && ((a_k_start) + SCALE_GROUP) <= K) { \
const intx4* a_vec_ptr = reinterpret_cast<const intx4*>( \
(A_ptr) + (a_row_global) * (stride_a) + (a_k_start)); \
(a_tmp0) = a_vec_ptr[0]; \
(a_tmp1) = a_vec_ptr[1]; \
(a_tmp2) = a_vec_ptr[2]; \
(a_tmp3) = a_vec_ptr[3]; \
(a_path_var) = 0; \
} else if ((a_row_global) < M) { \
(a_tmp0) = intx4{0,0,0,0}; \
(a_tmp1) = intx4{0,0,0,0}; \
(a_tmp2) = intx4{0,0,0,0}; \
(a_tmp3) = intx4{0,0,0,0}; \
for (int i = 0; i < SCALE_GROUP; i++) { \
int k_idx = (a_k_start) + i; \
uint16_t bf16_val = 0; \
if (k_idx < K) { \
bf16_val = (A_ptr)[(a_row_global) * (stride_a) + k_idx]; \
} \
int w = i / 2; \
int slot = i % 2; \
int* words; \
if (w < 4) words = reinterpret_cast<int*>(&(a_tmp0)); \
else if (w < 8) { words = reinterpret_cast<int*>(&(a_tmp1)); w -= 4; } \
else if (w < 12) { words = reinterpret_cast<int*>(&(a_tmp2)); w -= 8; } \
else { words = reinterpret_cast<int*>(&(a_tmp3)); w -= 12; } \
if (slot == 0) { \
words[w] = (int)((uint32_t)bf16_val); \
} else { \
words[w] |= (int)(((uint32_t)bf16_val) << 16); \
} \
} \
(a_path_var) = 1; \
} else { \
(a_tmp0) = intx4{0,0,0,0}; \
(a_tmp1) = intx4{0,0,0,0}; \
(a_tmp2) = intx4{0,0,0,0}; \
(a_tmp3) = intx4{0,0,0,0}; \
(a_path_var) = 2; \
} \
} while(0)
#define LOAD_B_DATA(b_mfma_var, b_tile_row, k_base, lane, n_base, N, K, row_in_tile, k_group) \
do { \
const uint8_t* tile_base = (b_tile_row) + ((k_base) / 2) * 16; \
if ((n_base) + TILE_N <= N && ((k_base) + K_STEP) <= K) { \
const intx4* b_vec_ptr = reinterpret_cast<const intx4*>( \
tile_base + (lane) * 16); \
(b_mfma_var) = *b_vec_ptr; \
} else { \
int b_n = (n_base) + (row_in_tile); \
if (b_n < N && ((k_base) + (k_group) * SCALE_GROUP + SCALE_GROUP) <= K) { \
const intx4* b_vec_ptr = reinterpret_cast<const intx4*>( \
tile_base + (lane) * 16); \
(b_mfma_var) = *b_vec_ptr; \
} else { \
(b_mfma_var)[0] = 0; (b_mfma_var)[1] = 0; \
(b_mfma_var)[2] = 0; (b_mfma_var)[3] = 0; \
} \
} \
} while(0)
#define LOAD_B_SCALE_SHUFFLED(b_e8m0_var, B_scale_sh_ptr, b_row_global, b_scale_k_idx, \
N, num_scale_k) \
do { \
if ((b_row_global) < N && (b_scale_k_idx) < (num_scale_k)) { \
int _n = (b_row_global); \
int _k = (b_scale_k_idx); \
int _n_group = _n >> 5; \
int _n_half = (_n >> 4) & 1; \
int _n_row16 = _n & 15; \
int _k_group8 = _k >> 3; \
int _k_half = (_k >> 2) & 1; \
int _k_mod4 = _k & 3; \
int _sh_idx = _n_group * ((num_scale_k) * 32) \
+ _k_group8 * 256 \
+ _k_mod4 * 64 \
+ _n_row16 * 4 \
+ _k_half * 2 \
+ _n_half; \
(b_e8m0_var) = (B_scale_sh_ptr)[_sh_idx]; \
} else { \
(b_e8m0_var) = 0; \
} \
} while(0)
#define PROCESS_TILE(cur_a0, cur_a1, cur_a2, cur_a3, cur_b, cur_bs, \
cur_path, acc_var) \
do { \
int a_words[16]; \
a_words[0] = (cur_a0)[0]; a_words[1] = (cur_a0)[1]; \
a_words[2] = (cur_a0)[2]; a_words[3] = (cur_a0)[3]; \
a_words[4] = (cur_a1)[0]; a_words[5] = (cur_a1)[1]; \
a_words[6] = (cur_a1)[2]; a_words[7] = (cur_a1)[3]; \
a_words[8] = (cur_a2)[0]; a_words[9] = (cur_a2)[1]; \
a_words[10] = (cur_a2)[2]; a_words[11] = (cur_a2)[3]; \
a_words[12] = (cur_a3)[0]; a_words[13] = (cur_a3)[1]; \
a_words[14] = (cur_a3)[2]; a_words[15] = (cur_a3)[3]; \
\
float amax_val = 0.0f; \
if ((cur_path) == 0 || (cur_path) == 1) { \
float av[32]; \
_Pragma("unroll") \
for (int w = 0; w < 16; w++) { \
uint32_t abs_word = ((uint32_t)a_words[w]) & 0x7FFF7FFFu; \
av[2*w] = bf16_to_f32((uint16_t)(abs_word & 0xFFFF)); \
av[2*w+1] = bf16_to_f32((uint16_t)(abs_word >> 16)); \
} \
\
float L1_0, L1_1, L1_2, L1_3, L1_4, L1_5, L1_6, L1_7, L1_8, L1_9; \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_0) : "v"(av[0]), "v"(av[1]), "v"(av[2])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_1) : "v"(av[3]), "v"(av[4]), "v"(av[5])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_2) : "v"(av[6]), "v"(av[7]), "v"(av[8])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_3) : "v"(av[9]), "v"(av[10]), "v"(av[11])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_4) : "v"(av[12]), "v"(av[13]), "v"(av[14])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_5) : "v"(av[15]), "v"(av[16]), "v"(av[17])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_6) : "v"(av[18]), "v"(av[19]), "v"(av[20])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_7) : "v"(av[21]), "v"(av[22]), "v"(av[23])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_8) : "v"(av[24]), "v"(av[25]), "v"(av[26])); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_9) : "v"(av[27]), "v"(av[28]), "v"(av[29])); \
\
float L2_0, L2_1, L2_2, L2_3; \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_0) : "v"(L1_0), "v"(L1_1), "v"(L1_2)); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_1) : "v"(L1_3), "v"(L1_4), "v"(L1_5)); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_2) : "v"(L1_6), "v"(L1_7), "v"(L1_8)); \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_3) : "v"(L1_9), "v"(av[30]), "v"(av[31])); \
\
float L3_0; \
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L3_0) : "v"(L2_0), "v"(L2_1), "v"(L2_2)); \
\
asm volatile("v_max_f32 %0, %1, %2" : "=v"(amax_val) : "v"(L3_0), "v"(L2_3)); \
} \
\
uint32_t amax_bits = float_as_uint(amax_val); \
amax_bits = ((amax_bits + 0x200000u) & 0xFF800000u); \
int scale_unbiased; \
if (amax_bits == 0) { \
scale_unbiased = -127; \
} else { \
int exp_biased = (int)((amax_bits >> 23) & 0xFF); \
scale_unbiased = exp_biased - 127 - 2; \
if (scale_unbiased < -127) scale_unbiased = -127; \
if (scale_unbiased > 127) scale_unbiased = 127; \
} \
uint8_t a_e8m0 = (uint8_t)(scale_unbiased + 127); \
float cvt_scale = uint_as_float(((uint32_t)a_e8m0) << 23); \
\
unsigned int dst0 = 0, dst1 = 0, dst2 = 0, dst3 = 0; \
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[0]), cvt_scale, 0); \
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[1]), cvt_scale, 1); \
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[2]), cvt_scale, 2); \
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[3]), cvt_scale, 3); \
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[4]), cvt_scale, 0); \
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[5]), cvt_scale, 1); \
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[6]), cvt_scale, 2); \
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[7]), cvt_scale, 3); \
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[8]), cvt_scale, 0); \
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[9]), cvt_scale, 1); \
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[10]), cvt_scale, 2); \
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[11]), cvt_scale, 3); \
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[12]), cvt_scale, 0); \
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[13]), cvt_scale, 1); \
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[14]), cvt_scale, 2); \
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[15]), cvt_scale, 3); \
\
intx4 a_mfma; \
a_mfma[0] = (int)dst0; \
a_mfma[1] = (int)dst1; \
a_mfma[2] = (int)dst2; \
a_mfma[3] = (int)dst3; \
\
int scale_a_packed = (int)((uint32_t)a_e8m0); \
int scale_b_packed = (int)((uint32_t)(cur_bs)); \
\
asm volatile("s_setprio 1" ::: "memory"); \
asm volatile( \
"v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4" \
: "+v"(acc_var) \
: "v"(a_mfma), "v"(cur_b), "v"(scale_a_packed), "v"(scale_b_packed) \
); \
asm volatile("s_setprio 0" ::: "memory"); \
} while(0)
// ═══════════════════════════════════════════════════════════════
// NEW: Fused M=16 kernel with DOUBLE-BUFFERED K-loop
// For Shape 2 (M=16, N=2112, K=7168) only
//
// Architecture:
// Grid: ceil(N/16) = 132 workgroups
// Block: 64 threads (1 wavefront)
// Tile: 16x16 output (one MFMA instruction per K=128 step)
// K-loop: 56 iterations of K=128, DOUBLE-BUFFERED
//
// KEY CHANGE from v103 (serial):
// v103: load_A → wait → LDS → read_LDS → quant → load_B → load_scale → MFMA → repeat
// v108: Prefetch A[0] & B[0], then loop: {prefetch[n+1] → process[n] → wait → repeat}
// This overlaps VMEM loads with MFMA compute.
//
// Pipeline per K-step:
// 1. Issue global loads for A[n+1] (non-blocking)
// 2. Issue global loads for B[n+1] and B_scale[n+1] (non-blocking)
// 3. Process current data: LDS write→wait→LDS read→quant→MFMA
// 4. s_waitcnt for next iteration's loads
//
// LDS: 2 x 4096 = 8192 bytes (double-buffered A data), trivially fits
// ═══════════════════════════════════════════════════════════════
// Helper: quantize 32 BF16 values (in 16 int words = 4 intx4) to FP4 (4 VGPRs)
// Returns a_mfma and a_e8m0 via out params
__device__ __forceinline__ void quantize_a_tile(
const intx4& a_data0, const intx4& a_data1,
const intx4& a_data2, const intx4& a_data3,
intx4& a_mfma, uint8_t& a_e8m0
) {
int a_words[16];
a_words[0] = a_data0[0]; a_words[1] = a_data0[1];
a_words[2] = a_data0[2]; a_words[3] = a_data0[3];
a_words[4] = a_data1[0]; a_words[5] = a_data1[1];
a_words[6] = a_data1[2]; a_words[7] = a_data1[3];
a_words[8] = a_data2[0]; a_words[9] = a_data2[1];
a_words[10] = a_data2[2]; a_words[11] = a_data2[3];
a_words[12] = a_data3[0]; a_words[13] = a_data3[1];
a_words[14] = a_data3[2]; a_words[15] = a_data3[3];
// Compute amax via V_MAX3_F32 tree reduction
float av[32];
#pragma unroll
for (int w = 0; w < 16; w++) {
uint32_t abs_word = ((uint32_t)a_words[w]) & 0x7FFF7FFFu;
av[2*w] = bf16_to_f32((uint16_t)(abs_word & 0xFFFF));
av[2*w+1] = bf16_to_f32((uint16_t)(abs_word >> 16));
}
float L1_0, L1_1, L1_2, L1_3, L1_4, L1_5, L1_6, L1_7, L1_8, L1_9;
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_0) : "v"(av[0]), "v"(av[1]), "v"(av[2]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_1) : "v"(av[3]), "v"(av[4]), "v"(av[5]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_2) : "v"(av[6]), "v"(av[7]), "v"(av[8]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_3) : "v"(av[9]), "v"(av[10]), "v"(av[11]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_4) : "v"(av[12]), "v"(av[13]), "v"(av[14]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_5) : "v"(av[15]), "v"(av[16]), "v"(av[17]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_6) : "v"(av[18]), "v"(av[19]), "v"(av[20]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_7) : "v"(av[21]), "v"(av[22]), "v"(av[23]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_8) : "v"(av[24]), "v"(av[25]), "v"(av[26]));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_9) : "v"(av[27]), "v"(av[28]), "v"(av[29]));
float L2_0, L2_1, L2_2, L2_3;
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_0) : "v"(L1_0), "v"(L1_1), "v"(L1_2));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_1) : "v"(L1_3), "v"(L1_4), "v"(L1_5));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_2) : "v"(L1_6), "v"(L1_7), "v"(L1_8));
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_3) : "v"(L1_9), "v"(av[30]), "v"(av[31]));
float L3_0;
asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L3_0) : "v"(L2_0), "v"(L2_1), "v"(L2_2));
float amax_val;
asm volatile("v_max_f32 %0, %1, %2" : "=v"(amax_val) : "v"(L3_0), "v"(L2_3));
uint32_t amax_bits = float_as_uint(amax_val);
amax_bits = ((amax_bits + 0x200000u) & 0xFF800000u);
int scale_unbiased;
if (amax_bits == 0) {
scale_unbiased = -127;
} else {
int exp_biased = (int)((amax_bits >> 23) & 0xFF);
scale_unbiased = exp_biased - 127 - 2;
if (scale_unbiased < -127) scale_unbiased = -127;
if (scale_unbiased > 127) scale_unbiased = 127;
}
a_e8m0 = (uint8_t)(scale_unbiased + 127);
float cvt_scale = uint_as_float(((uint32_t)a_e8m0) << 23);
unsigned int dst0 = 0, dst1 = 0, dst2 = 0, dst3 = 0;
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[0]), cvt_scale, 0);
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[1]), cvt_scale, 1);
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[2]), cvt_scale, 2);
dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[3]), cvt_scale, 3);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[4]), cvt_scale, 0);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[5]), cvt_scale, 1);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[6]), cvt_scale, 2);
dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[7]), cvt_scale, 3);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[8]), cvt_scale, 0);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[9]), cvt_scale, 1);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[10]), cvt_scale, 2);
dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[11]), cvt_scale, 3);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[12]), cvt_scale, 0);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[13]), cvt_scale, 1);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[14]), cvt_scale, 2);
dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[15]), cvt_scale, 3);
a_mfma[0] = (int)dst0;
a_mfma[1] = (int)dst1;
a_mfma[2] = (int)dst2;
a_mfma[3] = (int)dst3;
}
// ═══════════════════════════════════════════════════════════════
// Split-K Fused M=16 kernel: increases CU occupancy from 0.52 to 4.1 WGs/CU
//
// Grid: dim3(num_n_tiles, split_k) = (132, 8) = 1056 WGs for Shape 2
// Block: 64 threads (1 wavefront)
// Each WG processes K/split_k K-elements (e.g., 896 for K=7168, split_k=8)
// Output: F32 partial sums to workspace[split_idx * M * N + m * N + n]
// ═══════════════════════════════════════════════════════════════
__global__ __attribute__((amdgpu_flat_work_group_size(64, 64)))
void mxfp4_fused_gemm_m16_splitk(
const uint16_t* __restrict__ A, // [M, K] BF16 (M<=16)
const uint8_t* __restrict__ B_shuffle, // [N//16, K/2*16] preshuffled
const uint8_t* __restrict__ B_scale_sh, // shuffled e8m0 scales
float* __restrict__ workspace, // [split_k, M, N] F32 partial sums
int M, int N, int K,
int k_per_split // K-elements per split (aligned to K_STEP)
) {
// Double-buffered LDS for A data redistribution
__shared__ uint8_t lds_a[8192]; // 2 x (16 rows * 128 BF16 * 2 bytes)
const int n_tile = blockIdx.x;
const int split_idx = blockIdx.y;
const int n_base = n_tile * 16;
if (n_base >= N) return;
const int lane = threadIdx.x; // 0..63
const int row_in_tile = lane & 15; // N-index within tile (0..15)
const int k_group = lane >> 4; // K-group index (0..3)
const int K_half = K >> 1;
const int num_scale_k = (K + 31) / 32;
// B_shuffle tile base for this N-tile
const uint8_t* b_tile_base = B_shuffle + (long long)n_tile * K_half * 16;
// Accumulator
floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};
const int stride_a = K;
// ═══════════════════════════════════════════════════
// Compute K-range for this split
// ═══════════════════════════════════════════════════
int k_start = split_idx * k_per_split;
int k_end = k_start + k_per_split;
if (k_end > K) k_end = K;
if (k_start >= K) return;
// Number of K_STEP=128 iterations for this split
int first_k_step = k_start / K_STEP;
int last_k_step = (k_end + K_STEP - 1) / K_STEP;
int num_k_steps = last_k_step - first_k_step;
if (num_k_steps <= 0) return;
// ═══════════════════════════════════════════════════
// Precompute lane's global A address components
// Lane L loads A at: row = lane/4, col_base = (lane%4)*32
// ═══════════════════════════════════════════════════
const int a_row = lane >> 2; // 0..15
const int a_col_base = (lane & 3) * 32; // 0, 32, 64, 96
// ═══════════════════════════════════════════════════
// Precompute lane's LDS read offset
// ═══════════════════════════════════════════════════
const int lds_read_base = row_in_tile * 256 + k_group * 64;
// ═══════════════════════════════════════════════════
// Precompute B_scale shuffle index components
// ═══════════════════════════════════════════════════
const int b_row_global = n_base + row_in_tile;
const int _n_group = b_row_global >> 5;
const int _n_half = (b_row_global >> 4) & 1;
const int _n_row16 = b_row_global & 15;
const int _scale_n_base = _n_group * (num_scale_k * 32) + _n_row16 * 4 + _n_half;
// ═══════════════════════════════════════════════════
// PROLOGUE: Issue loads for first K-step (non-blocking)
// ═══════════════════════════════════════════════════
int lds_buf = 0;
int first_k_base = first_k_step * K_STEP;
// Load A[first_k] into VGPRs
intx4 a_load0, a_load1, a_load2, a_load3;
{
int a_k_global = first_k_base + a_col_base;
if (a_row < M && a_k_global + 31 < K) {
const intx4* a_vec = reinterpret_cast<const intx4*>(
A + a_row * stride_a + a_k_global);
a_load0 = a_vec[0];
a_load1 = a_vec[1];
a_load2 = a_vec[2];
a_load3 = a_vec[3];
} else {
a_load0 = intx4{0,0,0,0};
a_load1 = intx4{0,0,0,0};
a_load2 = intx4{0,0,0,0};
a_load3 = intx4{0,0,0,0};
}
}
// Load B[first_k] directly to register
intx4 b_mfma;
{
const uint8_t* b_k_ptr = b_tile_base + (long long)(first_k_base / 2) * 16;
const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_ptr + lane * 16);
b_mfma = *b_vec;
}
// Load B_scale[first_k]
uint8_t b_e8m0;
{
int b_scale_k_idx = first_k_base / SCALE_GROUP + k_group;
int _k = b_scale_k_idx;
int _k_group8 = _k >> 3;
int _k_half = (_k >> 2) & 1;
int _k_mod4 = _k & 3;
int _sh_idx = _scale_n_base
+ _k_group8 * 256
+ _k_mod4 * 64
+ _k_half * 2;
b_e8m0 = B_scale_sh[_sh_idx];
}
// ═══════════════════════════════════════════════════
// MAIN K-LOOP: Double-buffered over this split's K-range
// ═══════════════════════════════════════════════════
for (int step = 0; step < num_k_steps; step++) {
int abs_k_step = first_k_step + step;
int k_base = abs_k_step * K_STEP;
int lds_offset = lds_buf * 4096;
// ── Step 1: Write A data to LDS ──
{
intx4* lds_dst = reinterpret_cast<intx4*>(lds_a + lds_offset + lane * 64);
lds_dst[0] = a_load0;
lds_dst[1] = a_load1;
lds_dst[2] = a_load2;
lds_dst[3] = a_load3;
}
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
// ── Step 2: Read A from LDS with MFMA lane pattern ──
intx4 a_data0, a_data1, a_data2, a_data3;
{
const intx4* lds_src = reinterpret_cast<const intx4*>(
lds_a + lds_offset + lds_read_base);
a_data0 = lds_src[0];
a_data1 = lds_src[1];
a_data2 = lds_src[2];
a_data3 = lds_src[3];
}
// ── Step 3: Issue prefetch for NEXT K-step ──
intx4 next_a_load0, next_a_load1, next_a_load2, next_a_load3;
intx4 next_b_mfma;
uint8_t next_b_e8m0;
if (step + 1 < num_k_steps) {
int next_abs_k_step = first_k_step + step + 1;
int next_k_base = next_abs_k_step * K_STEP;
// Prefetch A[k+1]
{
int a_k_global = next_k_base + a_col_base;
if (a_row < M && a_k_global + 31 < K) {
const intx4* a_vec = reinterpret_cast<const intx4*>(
A + a_row * stride_a + a_k_global);
next_a_load0 = a_vec[0];
next_a_load1 = a_vec[1];
next_a_load2 = a_vec[2];
next_a_load3 = a_vec[3];
} else {
next_a_load0 = intx4{0,0,0,0};
next_a_load1 = intx4{0,0,0,0};
next_a_load2 = intx4{0,0,0,0};
next_a_load3 = intx4{0,0,0,0};
}
}
// Prefetch B[k+1]
{
const uint8_t* b_k_base_ptr = b_tile_base + (long long)(next_k_base / 2) * 16;
const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_base_ptr + lane * 16);
next_b_mfma = *b_vec;
}
// Prefetch B_scale[k+1]
{
int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
int _k = b_scale_k_idx;
int _k_group8 = _k >> 3;
int _k_half = (_k >> 2) & 1;
int _k_mod4 = _k & 3;
int _sh_idx = _scale_n_base
+ _k_group8 * 256
+ _k_mod4 * 64
+ _k_half * 2;
next_b_e8m0 = B_scale_sh[_sh_idx];
}
}
// ── Step 4: Quantize current A data (BF16 → FP4) ──
intx4 a_mfma;
uint8_t a_e8m0;
quantize_a_tile(a_data0, a_data1, a_data2, a_data3, a_mfma, a_e8m0);
// ── Step 5: MFMA execution ──
int scale_a_packed = (int)((uint32_t)a_e8m0);
int scale_b_packed = (int)((uint32_t)b_e8m0);
asm volatile("s_setprio 1" ::: "memory");
asm volatile(
"v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
: "+v"(acc)
: "v"(a_mfma), "v"(b_mfma), "v"(scale_a_packed), "v"(scale_b_packed)
);
asm volatile("s_setprio 0" ::: "memory");
// ── Step 6: Wait for prefetched data and swap buffers ──
if (step + 1 < num_k_steps) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
a_load0 = next_a_load0;
a_load1 = next_a_load1;
a_load2 = next_a_load2;
a_load3 = next_a_load3;
b_mfma = next_b_mfma;
b_e8m0 = next_b_e8m0;
}
lds_buf ^= 1;
}
// ════════════════════════════════════════════════
// Epilogue: Write F32 partial sums to workspace
// workspace layout: [split_k, M, N] in row-major
// Lane L has acc[0..3] = C[row_base+0..3, col]
// where col = L%16, row_base = (L/16)*4
// ════════════════════════════════════════════════
{
int col = lane & 15;
int row_base = (lane >> 4) * 4;
int n_global = n_base + col;
if (n_global < N) {
float* ws_slice = workspace + (long long)split_idx * M * N;
for (int v = 0; v < 4; v++) {
int m_global = row_base + v;
if (m_global < M) {
ws_slice[m_global * N + n_global] = acc[v];
}
}
}
}
}
// ═══════════════════════════════════════════════════════════════
// Reduction kernel: sum F32 partial sums across split_k, convert to BF16
// Grid: ceil(M * N / 256), Block: 256
// ═══════════════════════════════════════════════════════════════
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void reduce_splitk_to_bf16(
const float* __restrict__ workspace, // [split_k, M, N]
uint16_t* __restrict__ C, // [M, N] BF16
int M, int N, int split_k
) {
int idx = blockIdx.x * 256 + threadIdx.x;
int total = M * N;
if (idx >= total) return;
float sum = 0.0f;
for (int s = 0; s < split_k; s++) {
sum += workspace[(long long)s * M * N + idx];
}
C[idx] = f32_to_bf16(sum);
}
// ═══════════════════════════════════════════════════════════════
// NEW: Multi-wave fused M=16 kernel with ds_bpermute A-redistribution
//
// Grid: dim3(num_n_tiles) = 132 WGs for Shape 2
// Block: 512 threads (8 wavefronts per WG)
//
// Each wave independently processes K/8 K-range.
// KEY CHANGE (v124): Uses ds_bpermute_b32 for A-data lane remapping
// instead of LDS write + barrier + LDS read.
// - Eliminates per-wave LDS A double-buffers (saves 64KB LDS)
// - No lgkmcnt barrier needed for A redistribution
// - 16 ds_bpermute_b32 per K-step (one per int32 of 4 intx4)
//
// LDS layout (v124):
// Reduction only: 8 × 4 × 64 × 4 = 8192 bytes (8KB)
// Total: 8192 bytes — down from 72KB!
// ═══════════════════════════════════════════════════════════════
#define NUM_WAVES 8
#define MULTIWAVE_BLOCK_SIZE 512
__global__ __attribute__((amdgpu_flat_work_group_size(512, 512)))
void mxfp4_fused_gemm_m16_multiwave(
const uint16_t* __restrict__ A, // [M, K] BF16 (M<=16)
const uint8_t* __restrict__ B_shuffle, // [N//16, K/2*16] preshuffled
const uint8_t* __restrict__ B_scale_sh, // shuffled e8m0 scales
uint16_t* __restrict__ C, // [M, N] BF16 output (direct!)
int M, int N, int K,
int k_per_wave // K-elements per wave (aligned to K_STEP)
) {
// LDS: reduction buffer only (no A double-buffers needed with ds_bpermute!)
__shared__ uint8_t lds_raw[8192];
// Reduction buffer at offset 0
float* lds_reduce = reinterpret_cast<float*>(lds_raw);
// Layout: lds_reduce[wave_id * 256 + v * 64 + lane]
const int n_tile = blockIdx.x;
const int n_base = n_tile * 16;
if (n_base >= N) return;
const int wave_id = threadIdx.x / 64; // 0..7
const int lane = threadIdx.x % 64; // 0..63
const int row_in_tile = lane & 15; // N-index within tile (0..15)
const int k_group = lane >> 4; // K-group index (0..3)
const int K_half = K >> 1;
const int num_scale_k = (K + 31) / 32;
const int stride_a = K;
// B_shuffle tile base for this N-tile
const uint8_t* b_tile_base = B_shuffle + (long long)n_tile * K_half * 16;
// Accumulator (per-wave, independent)
floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};
// ═══════════════════════════════════════════════════
// Precompute ds_bpermute source lane offset
//
// A load pattern: lane L has row=L/4, kgrp=L%4
// MFMA needs: lane L wants row=L%16, kgrp=L/16
// Source lane: src = (L%16)*4 + (L/16)
// ds_bpermute takes byte offset = src_lane * 4
// ═══════════════════════════════════════════════════
const int bpermute_src_lane = (row_in_tile << 2) | k_group; // (L%16)*4 + (L/16)
const int bpermute_byte_offset = bpermute_src_lane << 2; // * 4 for byte offset
// ═══════════════════════════════════════════════════
// Compute K-range for this wave
// ═══════════════════════════════════════════════════
int k_start = wave_id * k_per_wave;
int k_end = k_start + k_per_wave;
if (k_end > K) k_end = K;
if (k_start >= K) goto do_reduction; // This wave has no work
{
int first_k_step = k_start / K_STEP;
int last_k_step = (k_end + K_STEP - 1) / K_STEP;
int num_k_steps = last_k_step - first_k_step;
if (num_k_steps <= 0) goto do_reduction;
// ═══════════════════════════════════════════════════
// Precompute lane's global A address components
// Lane L loads A at: row = lane/4, col_base = (lane%4)*32
// ═══════════════════════════════════════════════════
const int a_row = lane >> 2; // 0..15
const int a_col_base = (lane & 3) * 32; // 0, 32, 64, 96
// Precompute B_scale shuffle index components
const int b_row_global = n_base + row_in_tile;
const int _n_group = b_row_global >> 5;
const int _n_half = (b_row_global >> 4) & 1;
const int _n_row16 = b_row_global & 15;
const int _scale_n_base = _n_group * (num_scale_k * 32) + _n_row16 * 4 + _n_half;
// ═══════════════════════════════════════════════════
// PROLOGUE: Load first K-step data
// ═══════════════════════════════════════════════════
int first_k_base = first_k_step * K_STEP;
// Load A[first_k]
intx4 a_load0, a_load1, a_load2, a_load3;
{
int a_k_global = first_k_base + a_col_base;
if (a_row < M && a_k_global + 31 < K) {
const intx4* a_vec = reinterpret_cast<const intx4*>(
A + a_row * stride_a + a_k_global);
a_load0 = a_vec[0];
a_load1 = a_vec[1];
a_load2 = a_vec[2];
a_load3 = a_vec[3];
} else {
a_load0 = intx4{0,0,0,0};
a_load1 = intx4{0,0,0,0};
a_load2 = intx4{0,0,0,0};
a_load3 = intx4{0,0,0,0};
}
}
// Load B[first_k]
intx4 b_mfma;
{
const uint8_t* b_k_ptr = b_tile_base + (long long)(first_k_base / 2) * 16;
const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_ptr + lane * 16);
b_mfma = *b_vec;
}
// Load B_scale[first_k]
uint8_t b_e8m0;
{
int b_scale_k_idx = first_k_base / SCALE_GROUP + k_group;
int _k = b_scale_k_idx;
int _k_group8 = _k >> 3;
int _k_half = (_k >> 2) & 1;
int _k_mod4 = _k & 3;
int _sh_idx = _scale_n_base
+ _k_group8 * 256
+ _k_mod4 * 64
+ _k_half * 2;
b_e8m0 = B_scale_sh[_sh_idx];
}
// ═══════════════════════════════════════════════════
// MAIN K-LOOP: Each wave processes its own K-range
// Uses ds_bpermute for A-data redistribution (no LDS needed!)
// ═══════════════════════════════════════════════════
for (int step = 0; step < num_k_steps; step++) {
// ── Redistribute A data via ds_bpermute ──
// Each lane has a_load[0..3] in coalesced pattern (row=L/4, kgrp=L%4)
// MFMA needs (row=L%16, kgrp=L/16)
// ds_bpermute_b32: lane L reads from source lane = (L%16)*4 + (L/16)
intx4 a_data0, a_data1, a_data2, a_data3;
{
int d0, d1, d2, d3;
// Permute a_load0 (4 int32s)
d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[0]);
d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[1]);
d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[2]);
d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[3]);
a_data0 = intx4{d0, d1, d2, d3};
// Permute a_load1
d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[0]);
d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[1]);
d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[2]);
d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[3]);
a_data1 = intx4{d0, d1, d2, d3};
// Permute a_load2
d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[0]);
d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[1]);
d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[2]);
d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[3]);
a_data2 = intx4{d0, d1, d2, d3};
// Permute a_load3
d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[0]);
d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[1]);
d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[2]);
d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[3]);
a_data3 = intx4{d0, d1, d2, d3};
}
// ── Issue prefetch for NEXT K-step ──
intx4 next_a_load0, next_a_load1, next_a_load2, next_a_load3;
intx4 next_b_mfma;
uint8_t next_b_e8m0;
if (step + 1 < num_k_steps) {
int next_abs_k_step = first_k_step + step + 1;
int next_k_base = next_abs_k_step * K_STEP;
// Prefetch A[k+1]
{
int a_k_global = next_k_base + a_col_base;
if (a_row < M && a_k_global + 31 < K) {
const intx4* a_vec = reinterpret_cast<const intx4*>(
A + a_row * stride_a + a_k_global);
next_a_load0 = a_vec[0];
next_a_load1 = a_vec[1];
next_a_load2 = a_vec[2];
next_a_load3 = a_vec[3];
} else {
next_a_load0 = intx4{0,0,0,0};
next_a_load1 = intx4{0,0,0,0};
next_a_load2 = intx4{0,0,0,0};
next_a_load3 = intx4{0,0,0,0};
}
}
// Prefetch B[k+1]
{
const uint8_t* b_k_base_ptr = b_tile_base + (long long)(next_k_base / 2) * 16;
const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_base_ptr + lane * 16);
next_b_mfma = *b_vec;
}
// Prefetch B_scale[k+1]
{
int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
int _k = b_scale_k_idx;
int _k_group8 = _k >> 3;
int _k_half = (_k >> 2) & 1;
int _k_mod4 = _k & 3;
int _sh_idx = _scale_n_base
+ _k_group8 * 256
+ _k_mod4 * 64
+ _k_half * 2;
next_b_e8m0 = B_scale_sh[_sh_idx];
}
}
// ── Quantize current A data (BF16 → FP4) ──
intx4 a_mfma;
uint8_t a_e8m0;
quantize_a_tile(a_data0, a_data1, a_data2, a_data3, a_mfma, a_e8m0);
// ── MFMA execution ──
int scale_a_packed = (int)((uint32_t)a_e8m0);
int scale_b_packed = (int)((uint32_t)b_e8m0);
asm volatile("s_setprio 1" ::: "memory");
asm volatile(
"v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
: "+v"(acc)
: "v"(a_mfma), "v"(b_mfma), "v"(scale_a_packed), "v"(scale_b_packed)
);
asm volatile("s_setprio 0" ::: "memory");
// ── Swap to next prefetched data ──
if (step + 1 < num_k_steps) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
a_load0 = next_a_load0;
a_load1 = next_a_load1;
a_load2 = next_a_load2;
a_load3 = next_a_load3;
b_mfma = next_b_mfma;
b_e8m0 = next_b_e8m0;
}
}
} // end K-loop scope
do_reduction:
// ════════════════════════════════════════════════
// LDS Reduction: Sum accumulators across 8 waves
//
// Each wave's acc[0..3] holds partial sums for its K-range.
// We need to sum across all 8 waves for the same output position.
//
// Step 1: Each wave writes its acc[4] to LDS reduction buffer
// Step 2: __syncthreads() to ensure all waves have written
// Step 3: Wave 0 reads and sums all 8 waves' values
// Step 4: Wave 0 writes BF16 to output
// ════════════════════════════════════════════════
{
// Write acc to LDS reduction buffer
// Layout: lds_reduce[wave_id * 256 + v * 64 + lane]
int reduce_base = wave_id * 256 + lane;
lds_reduce[reduce_base + 0 * 64] = acc[0];
lds_reduce[reduce_base + 1 * 64] = acc[1];
lds_reduce[reduce_base + 2 * 64] = acc[2];
lds_reduce[reduce_base + 3 * 64] = acc[3];
__syncthreads();
// Only wave 0 does the final reduction and writes output
if (wave_id == 0) {
float sum0 = 0.0f, sum1 = 0.0f, sum2 = 0.0f, sum3 = 0.0f;
#pragma unroll
for (int w = 0; w < NUM_WAVES; w++) {
int w_base = w * 256 + lane;
sum0 += lds_reduce[w_base + 0 * 64];
sum1 += lds_reduce[w_base + 1 * 64];
sum2 += lds_reduce[w_base + 2 * 64];
sum3 += lds_reduce[w_base + 3 * 64];
}
// Write BF16 directly to output
int col = lane & 15;
int row_base = (lane >> 4) * 4;
int n_global = n_base + col;
if (n_global < N) {
for (int v = 0; v < 4; v++) {
int m_global = row_base + v;
if (m_global < M) {
float val = (v == 0) ? sum0 : (v == 1) ? sum1 : (v == 2) ? sum2 : sum3;
C[m_global * N + n_global] = f32_to_bf16(val);
}
}
}
}
}
}
// ═══════════════════════════════════════════════════════════════
// PyTorch wrapper for multi-wave fused M=16 kernel (single kernel!)
// ═══════════════════════════════════════════════════════════════
void launch_mxfp4_fused_gemm_m16_multiwave(
torch::Tensor A,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh,
torch::Tensor C, // [M, N] bfloat16 (direct output!)
int M, int N, int K
) {
TORCH_CHECK(A.is_cuda() && B_shuffle.is_cuda() && B_scale_sh.is_cuda() && C.is_cuda());
int num_n_tiles = (N + 15) / 16;
// Compute k_per_wave aligned to K_STEP=128
int k_steps_total = (K + K_STEP - 1) / K_STEP;
int k_steps_per_wave = (k_steps_total + NUM_WAVES - 1) / NUM_WAVES;
int k_per_wave = k_steps_per_wave * K_STEP;
dim3 grid(num_n_tiles);
dim3 block(MULTIWAVE_BLOCK_SIZE); // 512 threads = 8 wavefronts
mxfp4_fused_gemm_m16_multiwave<<<grid, block>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
B_shuffle.data_ptr<uint8_t>(),
B_scale_sh.data_ptr<uint8_t>(),
reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
M, N, K,
k_per_wave
);
}
// ═══════════════════════════════════════════════════════════════
// PyTorch wrapper for split-K fused M=16 kernel (legacy, kept for fallback)
// Launches main kernel then reduction kernel
// ═══════════════════════════════════════════════════════════════
void launch_mxfp4_fused_gemm_m16_splitk(
torch::Tensor A,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh,
torch::Tensor workspace, // [split_k, M, N] float32
torch::Tensor C, // [M, N] bfloat16
int M, int N, int K,
int split_k
) {
TORCH_CHECK(A.is_cuda() && B_shuffle.is_cuda() && B_scale_sh.is_cuda());
TORCH_CHECK(workspace.is_cuda() && C.is_cuda());
TORCH_CHECK(workspace.dtype() == torch::kFloat32, "workspace must be float32");
int num_n_tiles = (N + 15) / 16;
// Compute k_per_split aligned to K_STEP=128
int k_steps_total = (K + K_STEP - 1) / K_STEP;
int k_steps_per_split = (k_steps_total + split_k - 1) / split_k;
int k_per_split = k_steps_per_split * K_STEP;
// Main kernel: grid(num_n_tiles, split_k)
dim3 grid_main(num_n_tiles, split_k);
dim3 block_main(64); // 1 wavefront
mxfp4_fused_gemm_m16_splitk<<<grid_main, block_main>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
B_shuffle.data_ptr<uint8_t>(),
B_scale_sh.data_ptr<uint8_t>(),
workspace.data_ptr<float>(),
M, N, K,
k_per_split
);
// Reduction kernel: sum across split_k, convert to BF16
int total_elements = M * N;
int red_blocks = (total_elements + 255) / 256;
reduce_splitk_to_bf16<<<red_blocks, 256>>>(
workspace.data_ptr<float>(),
reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
M, N, split_k
);
}
// ═══════════════════════════════════════════════════════════════
// K=512 specialized 2-wave kernel with fully unrolled K-loop
// For S3/S4 (M=32, K=512): 4 K_STEP=128 iterations, cross-buffer pipeline
// ═══════════════════════════════════════════════════════════════
__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void mxfp4_fused_gemm_k512_2wave(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_preshuffle,
uint16_t* __restrict__ C,
const uint8_t* __restrict__ B_scale_sh,
int M, int N, int K,
int stride_bps
) {
const int wave_id = threadIdx.x / 64;
const int lane = threadIdx.x % 64;
int num_tile_n = (N + TILE_N - 1) / TILE_N;
int tile_m = blockIdx.x / num_tile_n;
int tile_n = blockIdx.x % num_tile_n;
int m_base = tile_m * TILE_M_GRID;
int m_wave = m_base + wave_id * TILE_M_PER_WAVE;
int n_base = tile_n * TILE_N;
int stride_a = K;
int num_scale_k = (K + 31) / 32;
floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};
int row_in_tile = lane & 15;
int k_group = lane >> 4;
const uint8_t* b_tile_row = B_preshuffle + (n_base / 16) * stride_bps;
// Precompute invariant A row
int a_row_global = m_wave + row_in_tile;
// ═══════════════════════════════════════════════════
// K=512: exactly 4 K_STEP=128 iterations, fully unrolled
// ═══════════════════════════════════════════════════
// ── Load step 0 ──
intx4 a_tmp0_0, a_tmp1_0, a_tmp2_0, a_tmp3_0;
intx4 b_mfma_0;
uint8_t b_e8m0_0;
int a_path_0;
{
int a_k_start = k_group * SCALE_GROUP;
LOAD_A_DATA(a_tmp0_0, a_tmp1_0, a_tmp2_0, a_tmp3_0,
a_path_0, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(b_mfma_0, b_tile_row, 0, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = k_group;
LOAD_B_SCALE_SHUFFLED(b_e8m0_0, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
}
// ── Load step 1, process step 0 ──
intx4 a_tmp0_1, a_tmp1_1, a_tmp2_1, a_tmp3_1;
intx4 b_mfma_1;
uint8_t b_e8m0_1;
int a_path_1;
{
int k_base_1 = K_STEP;
int a_k_start = k_base_1 + k_group * SCALE_GROUP;
LOAD_A_DATA(a_tmp0_1, a_tmp1_1, a_tmp2_1, a_tmp3_1,
a_path_1, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(b_mfma_1, b_tile_row, k_base_1, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = k_base_1 / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(b_e8m0_1, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
PROCESS_TILE(a_tmp0_0, a_tmp1_0, a_tmp2_0, a_tmp3_0,
b_mfma_0, b_e8m0_0, a_path_0, acc);
}
// ── Load step 2, process step 1 ──
intx4 a_tmp0_2, a_tmp1_2, a_tmp2_2, a_tmp3_2;
intx4 b_mfma_2;
uint8_t b_e8m0_2;
int a_path_2;
{
int k_base_2 = 2 * K_STEP;
int a_k_start = k_base_2 + k_group * SCALE_GROUP;
LOAD_A_DATA(a_tmp0_2, a_tmp1_2, a_tmp2_2, a_tmp3_2,
a_path_2, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(b_mfma_2, b_tile_row, k_base_2, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = k_base_2 / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(b_e8m0_2, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
PROCESS_TILE(a_tmp0_1, a_tmp1_1, a_tmp2_1, a_tmp3_1,
b_mfma_1, b_e8m0_1, a_path_1, acc);
}
// ── Load step 3, process step 2 ──
intx4 a_tmp0_3, a_tmp1_3, a_tmp2_3, a_tmp3_3;
intx4 b_mfma_3;
uint8_t b_e8m0_3;
int a_path_3;
{
int k_base_3 = 3 * K_STEP;
int a_k_start = k_base_3 + k_group * SCALE_GROUP;
LOAD_A_DATA(a_tmp0_3, a_tmp1_3, a_tmp2_3, a_tmp3_3,
a_path_3, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(b_mfma_3, b_tile_row, k_base_3, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = k_base_3 / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(b_e8m0_3, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
PROCESS_TILE(a_tmp0_2, a_tmp1_2, a_tmp2_2, a_tmp3_2,
b_mfma_2, b_e8m0_2, a_path_2, acc);
}
// ── Process step 3 (final) ──
{
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
PROCESS_TILE(a_tmp0_3, a_tmp1_3, a_tmp2_3, a_tmp3_3,
b_mfma_3, b_e8m0_3, a_path_3, acc);
}
// ── Store output ──
{
int col = lane & 15;
int row_base = (lane >> 4) * 4;
int n_global = n_base + col;
if (n_global < N) {
for (int v = 0; v < 4; v++) {
int m_global = m_wave + row_base + v;
if (m_global < M) {
C[m_global * N + n_global] = f32_to_bf16(acc[v]);
}
}
}
}
}
// ═══════════════════════════════════════════════════════════════
// FUSED kernel: bf16 A quantize + FP4×FP4 MFMA GEMM (from v47/v91)
// Cross-buffer vmcnt(6) pipeline + 2-wave + s_setprio
// ═══════════════════════════════════════════════════════════════
__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void mxfp4_fused_gemm_crossbuf(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_preshuffle,
uint16_t* __restrict__ C,
const uint8_t* __restrict__ B_scale_sh,
int M, int N, int K,
int stride_bps
) {
const int wave_id = threadIdx.x / 64;
const int lane = threadIdx.x % 64;
int num_tile_n = (N + TILE_N - 1) / TILE_N;
int tile_m = blockIdx.x / num_tile_n;
int tile_n = blockIdx.x % num_tile_n;
int m_base = tile_m * TILE_M_GRID;
int m_wave = m_base + wave_id * TILE_M_PER_WAVE;
int n_base = tile_n * TILE_N;
int stride_a = K;
int num_scale_k = (K + 31) / 32;
floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};
int row_in_tile = lane & 15;
int k_group = lane >> 4;
const uint8_t* b_tile_row = B_preshuffle + (n_base / 16) * stride_bps;
intx4 buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3;
intx4 buf0_b_mfma;
uint8_t buf0_b_e8m0;
int buf0_a_path;
intx4 buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3;
intx4 buf1_b_mfma;
uint8_t buf1_b_e8m0;
int buf1_a_path;
int num_k_tiles = (K + K_STEP - 1) / K_STEP;
if (num_k_tiles == 0) goto store_output;
{
int k_base = 0;
int a_row_global = m_wave + row_in_tile;
int a_k_start = k_base + k_group * SCALE_GROUP;
LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(buf0_b_mfma, b_tile_row, k_base, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = k_base / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
}
if (num_k_tiles == 1) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
goto store_output;
}
{
int k_tile_idx = 0;
for (; k_tile_idx < num_k_tiles - 1; k_tile_idx += 2) {
{
int next_k_base = (k_tile_idx + 1) * K_STEP;
int a_row_global = m_wave + row_in_tile;
int a_k_start = next_k_base + k_group * SCALE_GROUP;
LOAD_A_DATA(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
buf1_a_path, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(buf1_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(buf1_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
}
{
int next_k_tile = k_tile_idx + 2;
if (next_k_tile < num_k_tiles) {
int next_k_base = next_k_tile * K_STEP;
int a_row_global = m_wave + row_in_tile;
int a_k_start = next_k_base + k_group * SCALE_GROUP;
LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(buf0_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
} else {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
}
PROCESS_TILE(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
buf1_b_mfma, buf1_b_e8m0, buf1_a_path, acc);
}
}
if (k_tile_idx == num_k_tiles - 1) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
}
}
store_output:
{
int col = lane & 15;
int row_base = (lane >> 4) * 4;
int n_global = n_base + col;
if (n_global < N) {
for (int v = 0; v < 4; v++) {
int m_global = m_wave + row_base + v;
if (m_global < M) {
C[m_global * N + n_global] = f32_to_bf16(acc[v]);
}
}
}
}
}
// ═══════════════════════════════════════════════════════════════
// SPLIT-K kernel (from v91, for splitk shapes)
// ═══════════════════════════════════════════════════════════════
__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void mxfp4_fused_gemm_splitk(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_preshuffle,
float* __restrict__ C_f32,
const uint8_t* __restrict__ B_scale_sh,
int M, int N, int K,
int stride_bps,
int num_spatial_blocks,
int k_per_split
) {
const int wave_id = threadIdx.x / 64;
const int lane = threadIdx.x % 64;
int split_id = blockIdx.x / num_spatial_blocks;
int spatial_idx = blockIdx.x % num_spatial_blocks;
int num_tile_n = (N + TILE_N - 1) / TILE_N;
int tile_m = spatial_idx / num_tile_n;
int tile_n = spatial_idx % num_tile_n;
int m_base = tile_m * TILE_M_GRID;
int m_wave = m_base + wave_id * TILE_M_PER_WAVE;
int n_base = tile_n * TILE_N;
int k_start = split_id * k_per_split;
int k_end = k_start + k_per_split;
if (k_end > K) k_end = K;
if (k_start >= K) return;
int stride_a = K;
int num_scale_k = (K + 31) / 32;
floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};
int row_in_tile = lane & 15;
int k_group = lane >> 4;
const uint8_t* b_tile_row = B_preshuffle + (n_base / 16) * stride_bps;
intx4 buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3;
intx4 buf0_b_mfma;
uint8_t buf0_b_e8m0;
int buf0_a_path;
intx4 buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3;
intx4 buf1_b_mfma;
uint8_t buf1_b_e8m0;
int buf1_a_path;
int first_k_tile = k_start / K_STEP;
int last_k_tile = (k_end + K_STEP - 1) / K_STEP;
int num_k_tiles = last_k_tile - first_k_tile;
if (num_k_tiles == 0) return;
{
int k_base = first_k_tile * K_STEP;
int a_row_global = m_wave + row_in_tile;
int a_k_start = k_base + k_group * SCALE_GROUP;
LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(buf0_b_mfma, b_tile_row, k_base, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = k_base / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
}
if (num_k_tiles == 1) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
goto store_splitk_output;
}
{
int k_tile_offset = 0;
for (; k_tile_offset < num_k_tiles - 1; k_tile_offset += 2) {
{
int next_abs_k_tile = first_k_tile + k_tile_offset + 1;
int next_k_base = next_abs_k_tile * K_STEP;
int a_row_global = m_wave + row_in_tile;
int a_k_start = next_k_base + k_group * SCALE_GROUP;
LOAD_A_DATA(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
buf1_a_path, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(buf1_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(buf1_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
}
{
int next_k_tile_offset = k_tile_offset + 2;
if (next_k_tile_offset < num_k_tiles) {
int next_abs_k_tile = first_k_tile + next_k_tile_offset;
int next_k_base = next_abs_k_tile * K_STEP;
int a_row_global = m_wave + row_in_tile;
int a_k_start = next_k_base + k_group * SCALE_GROUP;
LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);
LOAD_B_DATA(buf0_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
row_in_tile, k_group);
int b_row_global = n_base + row_in_tile;
int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
N, num_scale_k);
asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
} else {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
}
PROCESS_TILE(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
buf1_b_mfma, buf1_b_e8m0, buf1_a_path, acc);
}
}
if (k_tile_offset == num_k_tiles - 1) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
}
}
store_splitk_output:
{
int col = lane & 15;
int row_base = (lane >> 4) * 4;
int n_global = n_base + col;
if (n_global < N) {
for (int v = 0; v < 4; v++) {
int m_global = m_wave + row_base + v;
if (m_global < M) {
atomicAdd(&C_f32[m_global * N + n_global], acc[v]);
}
}
}
}
}
// ═══════════════════════════════════════════════════════════════
// PyTorch wrappers
// ═══════════════════════════════════════════════════════════════
void launch_mxfp4_fused_gemm(
torch::Tensor A,
torch::Tensor B_preshuffle,
torch::Tensor C,
torch::Tensor B_scale_sh,
int M, int N, int K,
int stride_bps
) {
TORCH_CHECK(A.is_cuda() && B_preshuffle.is_cuda() && C.is_cuda() && B_scale_sh.is_cuda());
int num_tile_m = (M + TILE_M_GRID - 1) / TILE_M_GRID;
int num_tile_n = (N + TILE_N - 1) / TILE_N;
int num_blocks = num_tile_m * num_tile_n;
dim3 grid(num_blocks);
dim3 block(BLOCK_SIZE);
mxfp4_fused_gemm_crossbuf<<<grid, block>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
B_preshuffle.data_ptr<uint8_t>(),
reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
B_scale_sh.data_ptr<uint8_t>(),
M, N, K,
stride_bps
);
}
void launch_mxfp4_fused_gemm_splitk(
torch::Tensor A,
torch::Tensor B_preshuffle,
torch::Tensor C_f32,
torch::Tensor B_scale_sh,
int M, int N, int K,
int stride_bps,
int split_k
) {
TORCH_CHECK(A.is_cuda() && B_preshuffle.is_cuda() && C_f32.is_cuda() && B_scale_sh.is_cuda());
TORCH_CHECK(C_f32.dtype() == torch::kFloat32, "C_f32 must be float32 for atomicAdd");
int num_tile_m = (M + TILE_M_GRID - 1) / TILE_M_GRID;
int num_tile_n = (N + TILE_N - 1) / TILE_N;
int num_spatial_blocks = num_tile_m * num_tile_n;
int k_tiles_total = (K + K_STEP - 1) / K_STEP;
int k_tiles_per_split = (k_tiles_total + split_k - 1) / split_k;
int k_per_split = k_tiles_per_split * K_STEP;
int total_blocks = num_spatial_blocks * split_k;
dim3 grid(total_blocks);
dim3 block(BLOCK_SIZE);
mxfp4_fused_gemm_splitk<<<grid, block>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
B_preshuffle.data_ptr<uint8_t>(),
C_f32.data_ptr<float>(),
B_scale_sh.data_ptr<uint8_t>(),
M, N, K,
stride_bps,
num_spatial_blocks,
k_per_split
);
}
// ═══════════════════════════════════════════════════════════════
// PyTorch wrapper for K=512 2-wave unrolled kernel
// ═══════════════════════════════════════════════════════════════
void launch_mxfp4_fused_gemm_k512_2wave(
torch::Tensor A,
torch::Tensor B_preshuffle,
torch::Tensor C,
torch::Tensor B_scale_sh,
int M, int N, int K,
int stride_bps
) {
TORCH_CHECK(A.is_cuda() && B_preshuffle.is_cuda() && C.is_cuda() && B_scale_sh.is_cuda());
// 2-wave kernel: TILE_M_GRID=32 per WG (same as crossbuf)
int num_tile_m = (M + TILE_M_GRID - 1) / TILE_M_GRID;
int num_tile_n = (N + TILE_N - 1) / TILE_N;
int num_blocks = num_tile_m * num_tile_n;
dim3 grid(num_blocks);
dim3 block(BLOCK_SIZE); // 128 threads = 2 wavefronts
mxfp4_fused_gemm_k512_2wave<<<grid, block>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
B_preshuffle.data_ptr<uint8_t>(),
reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
B_scale_sh.data_ptr<uint8_t>(),
M, N, K,
stride_bps
);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("launch_mxfp4_fused_gemm", &launch_mxfp4_fused_gemm,
"Fused MXFP4 quant+GEMM");
m.def("launch_mxfp4_fused_gemm_splitk", &launch_mxfp4_fused_gemm_splitk,
"Fused MXFP4 quant+GEMM Split-K");
m.def("launch_mxfp4_quant", &launch_mxfp4_quant,
"Standalone MXFP4 quantization with shuffled scale output");
m.def("launch_mxfp4_fused_gemm_m16_splitk", &launch_mxfp4_fused_gemm_m16_splitk,
"Split-K fused MXFP4 quant+GEMM for M=16 with 1-wave 16x16 tile");
m.def("launch_mxfp4_fused_gemm_m16_multiwave", &launch_mxfp4_fused_gemm_m16_multiwave,
"Multi-wave fused MXFP4 quant+GEMM for M=16 — single kernel, no workspace");
m.def("launch_mxfp4_fused_gemm_k512_2wave", &launch_mxfp4_fused_gemm_k512_2wave,
"K=512 specialized 2-wave unrolled MXFP4 quant+GEMM");
}
"""
# ═══════════════════════════════════════════════════════════════
# Module-level HIP compilation (outside timed region)
# ═══════════════════════════════════════════════════════════════
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
import torch.utils.cpp_extension as cpp_ext
_hip_module = cpp_ext.load_inline(
name="v177_k512_2wave_only_gfx950",
cpp_sources="",
cuda_sources=_HIP_SOURCE,
extra_cuda_cflags=[
"--offload-arch=gfx950",
"-O3",
"-std=c++17",
],
with_cuda=True,
verbose=True,
)
# ═══════════════════════════════════════════════════════════════
# AITER ASM GEMM import (outside timed region)
# ═══════════════════════════════════════════════════════════════
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter import dtypes as aiter_dtypes
# ═══════════════════════════════════════════════════════════════
# Shape-specific AITER ASM tile/splitK configuration (from v91)
# ═══════════════════════════════════════════════════════════════
SHAPE_CONFIG = {
# Shape 5: M=64, N=7168, K=2048
(64, 7168, 2048): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
1
),
# Shape 6: M=256, N=3072, K=1536 — log2_k_split=1 from v119 (12.7→12.4µs)
(256, 3072, 1536): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
1
),
}
DEFAULT_CONFIG = (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
3
)
# ═══════════════════════════════════════════════════════════════
# NEW: Fused M=16 path for Shape 2
# ═══════════════════════════════════════════════════════════════
def _fused_m16_path(A, B_shuffle, B_scale_sh, m, k, n):
"""Multi-wave fused quant+GEMM for M=16 with 8-wave 16x16 MFMA tile.
8 wavefronts per WG (512 threads). Each wave processes K/8 K-elements.
LDS-based reduction across waves. Direct BF16 output.
NO workspace, NO reduction kernel — single kernel launch!
"""
B_sh_u8 = B_shuffle if B_shuffle.dtype == torch.uint8 else B_shuffle.view(torch.uint8)
# B_shuffle is [N, K/2] but preshuffled. We need [N//16, K/2*16] view
n_tile_groups = n >> 4
k_half_times_16 = (k >> 1) << 4
B_preshuffle = B_sh_u8.reshape(n_tile_groups, k_half_times_16)
B_scale_flat = B_scale_sh if B_scale_sh.dtype == torch.uint8 else B_scale_sh.view(torch.uint8)
# Output BF16 tensor (direct — no workspace needed!)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_hip_module.launch_mxfp4_fused_gemm_m16_multiwave(
A, B_preshuffle, B_scale_flat, C, m, n, k
)
return C
# ═══════════════════════════════════════════════════════════════
# HIP kernel dispatch — K=512 path
# v177: K=512 with M>16 uses specialized 2-wave unrolled kernel
# ═══════════════════════════════════════════════════════════════
def _hip_path(A, B_shuffle, B_scale_sh, m, k, n):
"""HIP fused quant+GEMM path for K=512 shapes."""
B_sh_u8 = B_shuffle if B_shuffle.dtype == torch.uint8 else B_shuffle.view(torch.uint8)
n_tile_groups = n >> 4
k_half_times_16 = (k >> 1) << 4
B_preshuffle = B_sh_u8.reshape(n_tile_groups, k_half_times_16)
stride_bps = k_half_times_16
B_scale_flat = B_scale_sh if B_scale_sh.dtype == torch.uint8 else B_scale_sh.view(torch.uint8)
# K=512 with M>16: use specialized unrolled 2-wave kernel (S3/S4)
if k == 512 and m > 16:
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_hip_module.launch_mxfp4_fused_gemm_k512_2wave(
A, B_preshuffle, C, B_scale_flat, m, n, k, stride_bps
)
return C
# Default: original crossbuf (handles S1 M=4, and any other K=512 with small M)
split_k = 4 if (m <= 16 and k >= 4096) else 1
if split_k > 1:
C_f32 = torch.zeros((m, n), dtype=torch.float32, device=A.device)
_hip_module.launch_mxfp4_fused_gemm_splitk(
A, B_preshuffle, C_f32, B_scale_flat, m, n, k, stride_bps, split_k
)
C = C_f32.to(torch.bfloat16)
return C
else:
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_hip_module.launch_mxfp4_fused_gemm(
A, B_preshuffle, C, B_scale_flat, m, n, k, stride_bps
)
return C
# ═══════════════════════════════════════════════════════════════
# AITER ASM path — HIP quant + ASM GEMM with tuned tile/splitK
# ═══════════════════════════════════════════════════════════════
def _aiter_asm_path(A, B_shuffle, B_scale_sh, m, k, n):
"""HIP quant kernel + AITER ASM GEMM with per-shape tuning."""
A_q = torch.empty((m, k // 2), dtype=torch.uint8, device=A.device)
scaleN_valid = (k + 31) // 32
scaleN = ((scaleN_valid + 7) // 8) * 8
M_pad256 = ((m + 255) // 256) * 256
A_scale_sh = torch.empty((M_pad256, scaleN), dtype=torch.uint8, device=A.device)
_hip_module.launch_mxfp4_quant(A, A_q, A_scale_sh, m, k, scaleN, scaleN_valid)
config = SHAPE_CONFIG.get((m, n, k), DEFAULT_CONFIG)
kernelName, log2_k_split = config
M_pad32 = ((m + 31) // 32) * 32
out = torch.empty((M_pad32, n), dtype=torch.bfloat16, device=A.device)
A_q_fp4 = A_q.view(aiter_dtypes.fp4x2)
B_shuffle_fp4 = B_shuffle.view(aiter_dtypes.fp4x2)
gemm_a4w4_asm(
A_q_fp4, B_shuffle_fp4, A_scale_sh, B_scale_sh,
out, kernelName, None, 1.0, 0.0, log2_k_split
)
return out[:m]
# ═══════════════════════════════════════════════════════════════
# Main dispatch
# ═══════════════════════════════════════════════════════════════
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B_shuffle.shape[0]
if m == 16 and k >= 1024:
# NEW: Fused M=16 path for Shape 2 (M=16, N=2112, K=7168)
return _fused_m16_path(A, B_shuffle, B_scale_sh, m, k, n)
elif k >= 1024:
# AITER ASM path with shape-specific tile/splitK
return _aiter_asm_path(A, B_shuffle, B_scale_sh, m, k, n)
else:
# Custom HIP fused quant+GEMM for K=512 shapes
return _hip_path(A, B_shuffle, B_scale_sh, m, k, n)
scrolls · 2034 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