submission 720553
lgc0338 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 788 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-720553?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:e2b725dfe68317386495100512fa1de5ef4cce288901a5dd88578bdd953e3846
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float red[NW * 64 * 16];split-k
void fused_mfma_splitk_kernel(vector-width = uint4
uint4 b_pf = {};Kernel source
submission.py788 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os
os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')
os.environ.setdefault('CXX', 'clang++')
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# =============================================================================
# v35: Based on v22 code style (proven to compile), incremental additions:
# - 16×16 kernel for M=4
# - 32×32 8wf splitK for M=16
# - 32×96 wide for M=256
# =============================================================================
HIP_SRC = r"""
#include <hip/hip_runtime.h>
typedef uint8_t fp4x2_t;
typedef fp4x2_t fp4x64_t __attribute__((ext_vector_type(32)));
typedef float fp32x16_t __attribute__((ext_vector_type(16)));
typedef float fp32x4_t __attribute__((ext_vector_type(4)));
__device__ __forceinline__ uint32_t f2u(float f) {
uint32_t u; __builtin_memcpy(&u, &f, 4); return u;
}
__device__ __forceinline__ float u2f(uint32_t u) {
float f; __builtin_memcpy(&f, &u, 4); return f;
}
__device__ __forceinline__ uint8_t read_shuffled_scale(
const uint8_t* s, int n, int ks, int sn
) {
int br = n / 32, rh = (n & 31) / 16, rl = (n & 31) & 15;
int bc = ks / 8, ch = (ks & 7) / 4, cl = (ks & 7) & 3;
return s[br * (sn * 32) + bc * 256 + cl * 64 + rl * 4 + ch * 2 + rh];
}
__device__ __forceinline__ int b_shuffle_addr(int n, int k_byte, int K_half) {
int n_block = n / 16;
int n_local = n & 15;
int k_block = k_byte / 32;
int k_group = (k_byte & 31) / 16;
return n_block * (K_half * 16) + k_block * 512 + k_group * 256 + n_local * 16;
}
// ============================================================
// Kernel 0: 32×32 tile (template NW) — for M=32, M=64
// ============================================================
template <int NW, bool EXACT_M = false>
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(64 * NW)
void fused_mfma_kernel(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_sh,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
const int M, const int N, const int K,
const int K_half, const int sn_pad
) {
__shared__ float red[NW * 64 * 16];
const int warp_id = threadIdx.x / 64;
const int tid = threadIdx.x & 63;
const int wf_row = tid & 31;
const int k_half = tid >> 5;
const int m_base = blockIdx.y * 32;
const int n_base = blockIdx.x * 32;
const int my_m = m_base + wf_row;
const int my_n = n_base + wf_row;
const int K_per_wf = K / NW;
const int k_start = warp_id * K_per_wf;
const int k_end = k_start + K_per_wf;
fp32x16_t c_reg = {};
uint4 b_pf = {};
uint8_t bs_pf = 127;
if (my_n < N) {
int bk = k_start / 2 + k_half * 16;
b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
bs_pf = read_shuffled_scale(B_scale_sh, my_n, k_start / 32 + k_half, sn_pad);
}
uint4 a_pf[4] = {};
if (EXACT_M || my_m < M) {
const int a_k0 = k_start + k_half * 32;
if (a_k0 + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
}
}
for (int k = k_start; k < k_end; k += 64) {
// Schedule: loads first → VALU (A quant) → MFMA
__builtin_amdgcn_sched_group_barrier(0x020, 6, 0);
__builtin_amdgcn_sched_group_barrier(0x002, 120, 0);
__builtin_amdgcn_sched_group_barrier(0x008, 1, 0);
fp4x64_t b_reg = {};
__builtin_memcpy(&b_reg, &b_pf, 16);
uint8_t scale_b = bs_pf;
uint4 a_data[4];
__builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);
int nk = k + 64;
if (nk < k_end) {
if (my_n < N) {
int bk = nk / 2 + k_half * 16;
b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
bs_pf = read_shuffled_scale(B_scale_sh, my_n, nk / 32 + k_half, sn_pad);
}
if (EXACT_M || my_m < M) {
const int a_k_next = nk + k_half * 32;
if (a_k_next + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
}
}
}
fp4x64_t a_reg = {};
uint8_t scale_a = 127;
if (EXACT_M || my_m < M) {
const int a_k = k + k_half * 32;
if (a_k + 32 <= K) {
const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
float vals[32];
float amax = 0.0f;
#pragma unroll
for (int i = 0; i < 32; i++) {
vals[i] = u2f((uint32_t)a_u16[i] << 16);
amax = fmaxf(amax, fabsf(vals[i]));
}
uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
int exp_biased = (int)((au >> 23) & 0xFFu);
int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
scale_a = (uint8_t)(sub_i + 127);
float qs = u2f((uint32_t)(sub_i + 127) << 23);
{
uint32_t pk[4] = {0, 0, 0, 0};
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0], vals[1], qs, 0);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8], vals[9], qs, 0);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2], vals[3], qs, 1);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4], vals[5], qs, 2);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6], vals[7], qs, 3);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
__builtin_memcpy(&a_reg, pk, 16);
}
}
}
c_reg = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
a_reg, b_reg, c_reg, 4, 4,
0, (uint32_t)scale_a, 0, (uint32_t)scale_b);
}
const int lb = warp_id * 64 * 16 + tid * 16;
#pragma unroll
for (int i = 0; i < 16; i++) red[lb + i] = c_reg[i];
__syncthreads();
if (warp_id == 0) {
int c_col = n_base + wf_row;
if (c_col < N) {
#pragma unroll
for (int i = 0; i < 4; i++) {
#pragma unroll
for (int j = 0; j < 4; j++) {
int c_row = m_base + k_half * 4 + i * 8 + j;
if (EXACT_M || c_row < M) {
float sum = 0.0f;
#pragma unroll
for (int w = 0; w < NW; w++)
sum += red[w * 64 * 16 + tid * 16 + i * 4 + j];
uint32_t bits = f2u(sum);
bits += (0x7FFFu + ((bits >> 16) & 1u));
C[c_row * N + c_col] = (uint16_t)(bits >> 16);
}
}
}
}
}
}
// ============================================================
// Kernel 1: 32×32 8wf splitK → float32 partials (for M=16)
// ============================================================
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(512)
void fused_mfma_splitk_kernel(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_sh,
const uint8_t* __restrict__ B_scale_sh,
float* __restrict__ partial,
const int M, const int N, const int K,
const int K_half, const int sn_pad,
const int splitK
) {
const int NW = 8;
__shared__ float red[NW * 64 * 16];
const int warp_id = threadIdx.x / 64;
const int tid = threadIdx.x & 63;
const int wf_row = tid & 31;
const int k_half = tid >> 5;
const int m_base = blockIdx.y * 32;
const int n_base = blockIdx.x * 32;
const int split_idx = blockIdx.z;
const int my_m = m_base + wf_row;
const int my_n = n_base + wf_row;
const int K_per_split = K / splitK;
const int K_split_start = split_idx * K_per_split;
const int K_per_wf = K_per_split / NW;
const int k_start = K_split_start + warp_id * K_per_wf;
const int k_end = k_start + K_per_wf;
fp32x16_t c_reg = {};
uint4 b_pf = {};
uint8_t bs_pf = 127;
if (my_n < N) {
int bk = k_start / 2 + k_half * 16;
b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
bs_pf = read_shuffled_scale(B_scale_sh, my_n, k_start / 32 + k_half, sn_pad);
}
uint4 a_pf[4] = {};
if (my_m < M) {
const int a_k0 = k_start + k_half * 32;
if (a_k0 + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
}
}
for (int k = k_start; k < k_end; k += 64) {
// Schedule: loads first → VALU (A quant) → MFMA
__builtin_amdgcn_sched_group_barrier(0x020, 6, 0);
__builtin_amdgcn_sched_group_barrier(0x002, 120, 0);
__builtin_amdgcn_sched_group_barrier(0x008, 1, 0);
fp4x64_t b_reg = {};
__builtin_memcpy(&b_reg, &b_pf, 16);
uint8_t scale_b = bs_pf;
uint4 a_data[4];
__builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);
int nk = k + 64;
if (nk < k_end) {
if (my_n < N) {
int bk = nk / 2 + k_half * 16;
b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
bs_pf = read_shuffled_scale(B_scale_sh, my_n, nk / 32 + k_half, sn_pad);
}
if (my_m < M) {
const int a_k_next = nk + k_half * 32;
if (a_k_next + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
}
}
}
fp4x64_t a_reg = {};
uint8_t scale_a = 127;
if (my_m < M) {
const int a_k = k + k_half * 32;
if (a_k + 32 <= K) {
const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
float vals[32];
float amax = 0.0f;
#pragma unroll
for (int i = 0; i < 32; i++) {
vals[i] = u2f((uint32_t)a_u16[i] << 16);
amax = fmaxf(amax, fabsf(vals[i]));
}
uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
int exp_biased = (int)((au >> 23) & 0xFFu);
int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
scale_a = (uint8_t)(sub_i + 127);
float qs = u2f((uint32_t)(sub_i + 127) << 23);
{
uint32_t pk[4] = {0, 0, 0, 0};
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0], vals[1], qs, 0);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8], vals[9], qs, 0);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2], vals[3], qs, 1);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4], vals[5], qs, 2);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6], vals[7], qs, 3);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
__builtin_memcpy(&a_reg, pk, 16);
}
}
}
c_reg = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
a_reg, b_reg, c_reg, 4, 4,
0, (uint32_t)scale_a, 0, (uint32_t)scale_b);
}
const int lb = warp_id * 64 * 16 + tid * 16;
#pragma unroll
for (int i = 0; i < 16; i++) red[lb + i] = c_reg[i];
__syncthreads();
if (warp_id == 0) {
int c_col = n_base + wf_row;
if (c_col < N) {
#pragma unroll
for (int i = 0; i < 4; i++) {
#pragma unroll
for (int j = 0; j < 4; j++) {
int c_row = m_base + k_half * 4 + i * 8 + j;
if (c_row < M) {
float sum = 0.0f;
#pragma unroll
for (int w = 0; w < 8; w++)
sum += red[w * 64 * 16 + tid * 16 + i * 4 + j];
partial[(long long)split_idx * M * N + c_row * N + c_col] = sum;
}
}
}
}
}
}
// ============================================================
// Kernel 2: 16×16 tile (for M=4)
// ============================================================
#define SMALL_NW 4
template <bool EXACT_M = false>
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(64 * SMALL_NW)
void fused_mfma_16x16_kernel(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_sh,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
const int M, const int N, const int K,
const int K_half, const int sn_pad
) {
__shared__ float red[SMALL_NW * 64 * 4];
const int warp_id = threadIdx.x / 64;
const int tid = threadIdx.x & 63;
const int lane16 = tid & 15;
const int k_quarter = tid >> 4;
const int m_base = blockIdx.y * 16;
const int n_base = blockIdx.x * 16;
const int my_m = m_base + lane16;
const int my_n = n_base + lane16;
const int K_per_wf = K / SMALL_NW;
const int k_start = warp_id * K_per_wf;
const int k_end = k_start + K_per_wf;
fp32x4_t c_reg = {};
uint4 b_pf = {};
uint8_t bs_pf = 127;
if (my_n < N) {
int bk = k_start / 2 + k_quarter * 16;
b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
bs_pf = read_shuffled_scale(B_scale_sh, my_n, k_start / 32 + k_quarter, sn_pad);
}
uint4 a_pf[4] = {};
if (EXACT_M || my_m < M) {
const int a_k0 = k_start + k_quarter * 32;
if (a_k0 + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
}
}
for (int k = k_start; k < k_end; k += 128) {
__builtin_amdgcn_sched_group_barrier(0x020, 6, 0);
__builtin_amdgcn_sched_group_barrier(0x002, 120, 0);
__builtin_amdgcn_sched_group_barrier(0x008, 1, 0);
fp4x64_t b_reg = {};
__builtin_memcpy(&b_reg, &b_pf, 16);
uint8_t scale_b = bs_pf;
uint4 a_data[4];
__builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);
int nk = k + 128;
if (nk < k_end) {
if (my_n < N) {
int bk = nk / 2 + k_quarter * 16;
b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
bs_pf = read_shuffled_scale(B_scale_sh, my_n, nk / 32 + k_quarter, sn_pad);
}
if (EXACT_M || my_m < M) {
const int a_k_next = nk + k_quarter * 32;
if (a_k_next + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
}
}
}
fp4x64_t a_reg = {};
uint8_t scale_a = 127;
if (EXACT_M || my_m < M) {
const int a_k = k + k_quarter * 32;
if (a_k + 32 <= K) {
const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
float vals[32];
float amax = 0.0f;
#pragma unroll
for (int i = 0; i < 32; i++) {
vals[i] = u2f((uint32_t)a_u16[i] << 16);
amax = fmaxf(amax, fabsf(vals[i]));
}
uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
int exp_biased = (int)((au >> 23) & 0xFFu);
int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
scale_a = (uint8_t)(sub_i + 127);
float qs = u2f((uint32_t)(sub_i + 127) << 23);
{
uint32_t pk[4] = {0, 0, 0, 0};
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0], vals[1], qs, 0);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8], vals[9], qs, 0);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2], vals[3], qs, 1);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4], vals[5], qs, 2);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6], vals[7], qs, 3);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
__builtin_memcpy(&a_reg, pk, 16);
}
}
}
c_reg = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_reg, b_reg, c_reg, 4, 4,
0, (uint32_t)scale_a, 0, (uint32_t)scale_b);
}
const int lb = warp_id * 64 * 4 + tid * 4;
#pragma unroll
for (int i = 0; i < 4; i++) red[lb + i] = c_reg[i];
__syncthreads();
if (warp_id == 0) {
int c_col = n_base + lane16;
if (c_col < N) {
int group = k_quarter;
#pragma unroll
for (int i = 0; i < 4; i++) {
int c_row = m_base + group * 4 + i;
if (EXACT_M || c_row < M) {
float sum = 0.0f;
#pragma unroll
for (int w = 0; w < SMALL_NW; w++)
sum += red[w * 64 * 4 + tid * 4 + i];
uint32_t bits = f2u(sum);
bits += (0x7FFFu + ((bits >> 16) & 1u));
C[c_row * N + c_col] = (uint16_t)(bits >> 16);
}
}
}
}
}
// ============================================================
// Kernel 3: 32×96 wide tile (3 MFMA, for M=256)
// ============================================================
#define WIDE96_NW 4
template <bool EXACT_M = false>
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(64 * WIDE96_NW)
void fused_mfma_wide96_kernel(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_sh,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
const int M, const int N, const int K,
const int K_half, const int sn_pad
) {
__shared__ float red[WIDE96_NW * 64 * 48];
const int warp_id = threadIdx.x / 64;
const int tid = threadIdx.x & 63;
const int wf_row = tid & 31;
const int k_half = tid >> 5;
const int m_base = blockIdx.y * 32;
const int n_base = blockIdx.x * 96;
const int my_m = m_base + wf_row;
const int K_per_wf = K / WIDE96_NW;
const int k_start = warp_id * K_per_wf;
const int k_end = k_start + K_per_wf;
fp32x16_t c0 = {}, c1 = {}, c2 = {};
// Prefetch B for 3 N-subtiles
uint4 b_pf0 = {}, b_pf1 = {}, b_pf2 = {};
uint8_t bs_pf0 = 127, bs_pf1 = 127, bs_pf2 = 127;
{
int bk = k_start / 2 + k_half * 16;
int n0 = n_base + wf_row, n1 = n_base + 32 + wf_row, n2 = n_base + 64 + wf_row;
if (n0 < N) { b_pf0 = *(const uint4*)(B_sh + b_shuffle_addr(n0, bk, K_half)); bs_pf0 = read_shuffled_scale(B_scale_sh, n0, k_start/32+k_half, sn_pad); }
if (n1 < N) { b_pf1 = *(const uint4*)(B_sh + b_shuffle_addr(n1, bk, K_half)); bs_pf1 = read_shuffled_scale(B_scale_sh, n1, k_start/32+k_half, sn_pad); }
if (n2 < N) { b_pf2 = *(const uint4*)(B_sh + b_shuffle_addr(n2, bk, K_half)); bs_pf2 = read_shuffled_scale(B_scale_sh, n2, k_start/32+k_half, sn_pad); }
}
uint4 a_pf[4] = {};
if (EXACT_M || my_m < M) {
const int a_k0 = k_start + k_half * 32;
if (a_k0 + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
}
}
for (int k = k_start; k < k_end; k += 64) {
uint4 bc0 = b_pf0, bc1 = b_pf1, bc2 = b_pf2;
uint8_t bsc0 = bs_pf0, bsc1 = bs_pf1, bsc2 = bs_pf2;
uint4 a_data[4];
__builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);
int nk = k + 64;
if (nk < k_end) {
int bk = nk / 2 + k_half * 16;
int n0 = n_base + wf_row, n1 = n_base + 32 + wf_row, n2 = n_base + 64 + wf_row;
if (n0 < N) { b_pf0 = *(const uint4*)(B_sh + b_shuffle_addr(n0, bk, K_half)); bs_pf0 = read_shuffled_scale(B_scale_sh, n0, nk/32+k_half, sn_pad); }
if (n1 < N) { b_pf1 = *(const uint4*)(B_sh + b_shuffle_addr(n1, bk, K_half)); bs_pf1 = read_shuffled_scale(B_scale_sh, n1, nk/32+k_half, sn_pad); }
if (n2 < N) { b_pf2 = *(const uint4*)(B_sh + b_shuffle_addr(n2, bk, K_half)); bs_pf2 = read_shuffled_scale(B_scale_sh, n2, nk/32+k_half, sn_pad); }
if (EXACT_M || my_m < M) {
const int a_k_next = nk + k_half * 32;
if (a_k_next + 32 <= K) {
#pragma unroll
for (int i = 0; i < 4; i++)
a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
}
}
}
fp4x64_t a_reg = {};
uint8_t scale_a = 127;
if (EXACT_M || my_m < M) {
const int a_k = k + k_half * 32;
if (a_k + 32 <= K) {
const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
float vals[32];
float amax = 0.0f;
#pragma unroll
for (int i = 0; i < 32; i++) {
vals[i] = u2f((uint32_t)a_u16[i] << 16);
amax = fmaxf(amax, fabsf(vals[i]));
}
uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
int exp_biased = (int)((au >> 23) & 0xFFu);
int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
scale_a = (uint8_t)(sub_i + 127);
float qs = u2f((uint32_t)(sub_i + 127) << 23);
{
uint32_t pk[4] = {0, 0, 0, 0};
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0], vals[1], qs, 0);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8], vals[9], qs, 0);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2], vals[3], qs, 1);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4], vals[5], qs, 2);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6], vals[7], qs, 3);
pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
__builtin_memcpy(&a_reg, pk, 16);
}
}
}
// 3 MFMA sharing same A
{ fp4x64_t b_reg = {}; __builtin_memcpy(&b_reg, &bc0, 16);
c0 = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg, b_reg, c0, 4, 4, 0, (uint32_t)scale_a, 0, (uint32_t)bsc0); }
{ fp4x64_t b_reg = {}; __builtin_memcpy(&b_reg, &bc1, 16);
c1 = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg, b_reg, c1, 4, 4, 0, (uint32_t)scale_a, 0, (uint32_t)bsc1); }
{ fp4x64_t b_reg = {}; __builtin_memcpy(&b_reg, &bc2, 16);
c2 = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg, b_reg, c2, 4, 4, 0, (uint32_t)scale_a, 0, (uint32_t)bsc2); }
}
// LDS reduce + output all 3 subtiles
const int lb = warp_id * 64 * 48 + tid * 48;
#pragma unroll
for (int i = 0; i < 16; i++) { red[lb + i] = c0[i]; red[lb + 16 + i] = c1[i]; red[lb + 32 + i] = c2[i]; }
__syncthreads();
if (warp_id == 0) {
for (int s = 0; s < 3; s++) {
int c_col = n_base + s * 32 + wf_row;
if (c_col < N) {
#pragma unroll
for (int i = 0; i < 4; i++) {
#pragma unroll
for (int j = 0; j < 4; j++) {
int c_row = m_base + k_half * 4 + i * 8 + j;
if (EXACT_M || c_row < M) {
float sum = 0.0f;
#pragma unroll
for (int w = 0; w < WIDE96_NW; w++)
sum += red[w * 64 * 48 + tid * 48 + s * 16 + i * 4 + j];
uint32_t bits = f2u(sum);
bits += (0x7FFFu + ((bits >> 16) & 1u));
C[c_row * N + c_col] = (uint16_t)(bits >> 16);
}
}
}
}
}
}
}
// ============================================================
// Kernel 4: Reduction (sum splitK partials → bf16)
// ============================================================
__global__ void reduce_splitk_kernel(
const float* __restrict__ partial,
uint16_t* __restrict__ C,
const int M, const int N, const int splitK
) {
const int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= M * N) return;
const int m = idx / N, n = idx % N;
float sum = 0.0f;
for (int s = 0; s < splitK; s++)
sum += partial[(long long)s * M * N + m * N + n];
uint32_t bits = f2u(sum);
bits += (0x7FFFu + ((bits >> 16) & 1u));
C[m * N + n] = (uint16_t)(bits >> 16);
}
// ============================================================
// Dispatch
// ============================================================
void launch_v35(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
torch::Tensor C, int sn_pad, int kernel_id, int splitK,
torch::Tensor partial
) {
const int M = A.size(0), K = A.size(1), N = B_shuffle.size(0);
const uint16_t* Ap = (const uint16_t*)A.data_ptr();
const uint8_t* Bp = (const uint8_t*)B_shuffle.data_ptr();
const uint8_t* Sp = (const uint8_t*)B_scale_sh.data_ptr();
uint16_t* Cp = (uint16_t*)C.data_ptr();
if (kernel_id == 5) {
// 32×32 8wf splitK
dim3 grid_sk((N + 31) / 32, (M + 31) / 32, splitK);
dim3 block_sk(64 * 8);
fused_mfma_splitk_kernel<<<grid_sk, block_sk>>>(
Ap, Bp, Sp, partial.data_ptr<float>(), M, N, K, K/2, sn_pad, splitK);
int total = M * N;
dim3 grid_r((total + 255) / 256);
dim3 block_r(256);
reduce_splitk_kernel<<<grid_r, block_r>>>(partial.data_ptr<float>(), Cp, M, N, splitK);
} else if (kernel_id == 9) {
// 32×96 wide EXACT_M
dim3 grid_w((N + 95) / 96, (M + 31) / 32);
dim3 block_w(64 * WIDE96_NW);
fused_mfma_wide96_kernel<true><<<grid_w, block_w>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
} else if (kernel_id == 4) {
// 32×96 wide
dim3 grid_w((N + 95) / 96, (M + 31) / 32);
dim3 block_w(64 * WIDE96_NW);
fused_mfma_wide96_kernel<<<grid_w, block_w>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
} else if (kernel_id == 10) {
// 16×16 EXACT_M
dim3 grid_16((N + 15) / 16, (M + 15) / 16);
dim3 block_16(64 * SMALL_NW);
fused_mfma_16x16_kernel<true><<<grid_16, block_16>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
} else if (kernel_id == 3) {
// 16×16
dim3 grid_16((N + 15) / 16, (M + 15) / 16);
dim3 block_16(64 * SMALL_NW);
fused_mfma_16x16_kernel<<<grid_16, block_16>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
} else if (kernel_id == 1) {
// 32×32 8wf
dim3 grid_8((N + 31) / 32, (M + 31) / 32);
dim3 block_8(64 * 8);
fused_mfma_kernel<8><<<grid_8, block_8>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
} else if (kernel_id == 8) {
// 32×32 4wf EXACT_M (no M bounds check)
dim3 grid_4((N + 31) / 32, (M + 31) / 32);
dim3 block_4(64 * 4);
fused_mfma_kernel<4, true><<<grid_4, block_4>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
} else {
// 32×32 4wf
dim3 grid_4((N + 31) / 32, (M + 31) / 32);
dim3 block_4(64 * 4);
fused_mfma_kernel<4><<<grid_4, block_4>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
}
}
"""
CPP_SRC = "void launch_v35(torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C, int sn_pad, int kernel_id, int splitK, torch::Tensor partial);"
_CFLAGS = ["--offload-arch=gfx950", "-std=c++20", "-O3", "-ffast-math", "-funsafe-math-optimizations", "-mno-wavefrontsize64"]
import hashlib as _hl
_hip = load_inline(name='fm_'+_hl.md5(HIP_SRC.encode()).hexdigest()[:10], cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
functions=['launch_v35'], verbose=True, extra_cuda_cflags=_CFLAGS)
# Separate wide96 EXACT_MN module (isolated compilation → no icache/regalloc interference)
_W96_HELPERS = HIP_SRC[HIP_SRC.index('#include'):HIP_SRC.index('// Kernel 0')]
_W96_BODY = HIP_SRC[HIP_SRC.index('// Kernel 3:'):HIP_SRC.index('// Kernel 4:')]
# Replace all M/N bounds with true
_W96_BODY = _W96_BODY.replace('EXACT_M || my_m < M', 'true').replace('EXACT_M || c_row < M', 'true')
_W96_BODY = _W96_BODY.replace('if (n0 < N)', 'if (true)').replace('if (n1 < N)', 'if (true)').replace('if (n2 < N)', 'if (true)')
_W96_BODY = _W96_BODY.replace('if (c_col < N)', 'if (true)')
_W96_BODY = _W96_BODY.replace('template <bool EXACT_M = false>\n', '')
_W96_SRC = _W96_HELPERS + _W96_BODY + r"""
void launch_w96(torch::Tensor A, torch::Tensor B_sh, torch::Tensor B_sc, torch::Tensor C, int sn) {
const int M=A.size(0),K=A.size(1),N=B_sh.size(0);
dim3 g((N+95)/96,(M+31)/32);
fused_mfma_wide96_kernel<<<g,64*WIDE96_NW>>>((const uint16_t*)A.data_ptr(),(const uint8_t*)B_sh.data_ptr(),(const uint8_t*)B_sc.data_ptr(),(uint16_t*)C.data_ptr(),M,N,K,K/2,sn);
}
"""
_W96_CPP = "void launch_w96(torch::Tensor A, torch::Tensor B_sh, torch::Tensor B_sc, torch::Tensor C, int sn);"
_hip_w96 = load_inline(name='w9_'+_hl.md5(_W96_SRC.encode()).hexdigest()[:10], cpp_sources=[_W96_CPP], cuda_sources=[_W96_SRC],
functions=['launch_w96'], verbose=True, extra_cuda_cflags=_CFLAGS)
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
C = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
if m >= 128:
blocks_96 = ((n + 95) // 96) * ((m + 31) // 32)
blocks_32 = ((n + 31) // 32) * ((m + 31) // 32)
if (blocks_96 + 255) // 256 < (blocks_32 + 255) // 256 and m % 32 == 0:
# Use separate wide96 EXACT_MN module
_hip_w96.launch_w96(A, B_shuffle, B_scale_sh, C, B_scale_sh.shape[1])
return C
elif (blocks_96 + 255) // 256 < (blocks_32 + 255) // 256:
kernel_id, splitK = 4, 0
else:
kernel_id, splitK = (8 if m % 32 == 0 else 0), 0
elif m <= 8:
kernel_id, splitK = 3, 0
elif m <= 16 and k >= 2048:
splitK = 7
if (k // splitK) % (8 * 64) == 0:
kernel_id = 5
else:
kernel_id, splitK = 1, 0
elif m <= 16:
kernel_id, splitK = 3, 0
elif m <= 32 and k <= 512:
kernel_id, splitK = (10 if m % 16 == 0 else 3), 0
else:
# M=64: use EXACT_M version if M is multiple of 32
if m % 32 == 0:
kernel_id, splitK = 8, 0 # no bounds check
else:
kernel_id, splitK = 0, 0
partial = torch.empty(splitK * m * n, dtype=torch.float32, device=A.device) if splitK > 0 else torch.empty(0, device=A.device)
_hip.launch_v35(A, B_shuffle, B_scale_sh, C, B_scale_sh.shape[1], kernel_id, splitK, partial)
return C
scrolls · 788 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