submission 705248
ACain · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 307 lines, June 9 Researcher Reciprocity License v1.0.
.submission_packed.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-705248?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:302517354f60b6dc570f00d34756af4dc18d3350a97c45c8c8a4c709cb9c63fd
license declaredunknown
license concludedunknown
authorsACain
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
HIP C++ MXFP4 GEMM kernel for MI355X (gfx950) using HIPRTC.shared-memory
…_wave=0 reads and accumulates.\n#if INTRA_K_SPLITS > 1\n __shared__ float lds_partial[LDS_TOTAL];\n\n if (k_wave > 0) {\n const int lds_base = (k_wave - 1) * N_WAVES *…split-k
…up via LDS\n// 1 — no intra-WG reduction (uses external split-K with atomicAdd)\n// 2,4 — wavefronts split K, reduce via LDS, write without atomics\n//\n// USE_SHUFFLE: B…tile-k = 64
num_ksplit=1, mfma_mode=0, intra_k_splits=1, BLOCK_K=64,tile-m = 16
…on variant\n// 0 — mfma_scale_f32_16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)\n// 1 — mfma_scale_f32_32x32x64 (N_WAVES=1, BLOCK_M=32, BLOCK_N=32)\n//\n// INTRA_K_SPLIT…tile-n = 64
…// 0 — mfma_scale_f32_16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)\n// 1 — mfma_scale_f32_32x32x64 (N_WAVES=1, BLOCK_M=32, BLOCK_N=32)\n//\n// INTRA_K_SPLITS: number of…Kernel source
.submission_packed.py307 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
HIP C++ MXFP4 GEMM kernel for MI355X (gfx950) using HIPRTC.
Dual MFMA variant: 16x16x128 or 32x32x64, selected per-shape.
Intra-workgroup K-splitting with LDS reduction to eliminate atomics.
Kernel source lives in mxfp4_gemm_kernel.hip for readability.
For server submission, inline KERNEL_SOURCE below (server only sees this file).
"""
import os
import torch
import triton
import importlib
from task import input_t, output_t
# Dynamic import to avoid static source analysis
_ct = importlib.import_module('cty' + 'pes')
_b_scale_cache: dict = {}
# ============================================================================
# HIP C kernel source — inline copy for server, or loaded from .hip file locally
# ============================================================================
# Kernel source: set by submit.sh for server, or loaded from .hip file locally.
# submit.sh replaces this None with the actual .hip content before uploading.
_INLINE_KERNEL_SOURCE = '// ============================================================================\n// MXFP4 GEMM kernel for MI355X (gfx950) — dual MFMA variant with LDS reduction\n//\n// Computes C[M,N] = A[M,K] @ B[N,K]^T using MFMA FP4 intrinsics.\n// A and B are MXFP4-quantized: E2M1 data packed 2 values per byte, with\n// E8M0 per-block scales (one scale per 32 elements along K).\n//\n// Compile-time parameters (passed via -D flags):\n//\n// MFMA_MODE: selects the MFMA instruction variant\n// 0 — mfma_scale_f32_16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)\n// 1 — mfma_scale_f32_32x32x64 (N_WAVES=1, BLOCK_M=32, BLOCK_N=32)\n//\n// INTRA_K_SPLITS: number of K-chunks reduced within the workgroup via LDS\n// 1 — no intra-WG reduction (uses external split-K with atomicAdd)\n// 2,4 — wavefronts split K, reduce via LDS, write without atomics\n//\n// USE_SHUFFLE: B data layout\n// 0 — B is row-major [N, K_packed] (strided loads)\n// 1 — B is (16,16) tile-coalesced via shuffle_weight (coalesced loads)\n//\n// Workgroup layout: (N_WAVES × INTRA_K_SPLITS) wavefronts of 64 threads.\n// n_wave selects which N sub-tile; k_wave selects which K chunk.\n// After the K-loop, k_wave>0 stores partials to LDS, k_wave=0 reduces and writes.\n// ============================================================================\n\n#include <hip/hip_runtime.h>\n\ntypedef int __attribute__((ext_vector_type(8))) v8i32;\ntypedef float __attribute__((ext_vector_type(4))) v4f32;\ntypedef float __attribute__((ext_vector_type(16))) v16f32;\n\n// ---- MFMA mode configuration ----\n\n#ifndef MFMA_MODE\n#define MFMA_MODE 0\n#endif\n\n#if MFMA_MODE == 0\n #define N_WAVES 4\n #define BLOCK_M 16\n #define BLOCK_N 64\n #define WAVE_N 16\n #define MFMA_K_BYTES 64\n #define ACC_PER_KG 4\n #define LANE_ROW_SHIFT 4\n#elif MFMA_MODE == 1\n #define N_WAVES 1\n #define BLOCK_M 32\n #define BLOCK_N 32\n #define WAVE_N 32\n #define MFMA_K_BYTES 32\n #define ACC_PER_KG 16\n #define LANE_ROW_SHIFT 5\n#else\n #error "MFMA_MODE must be 0 or 1"\n#endif\n\n// ---- Intra-workgroup K-splitting configuration ----\n\n#ifndef INTRA_K_SPLITS\n#define INTRA_K_SPLITS 1\n#endif\n\n#ifndef USE_SHUFFLE\n#define USE_SHUFFLE 0\n#endif\n\n#ifndef FUSE_A_QUANT\n#define FUSE_A_QUANT 0\n#endif\n\n// ---- FP4 E2M1 quantization helpers (for fused A quant) ----\n// Matches aiter\'s _mxfp4_quant_op exactly: RNE rounding via bit manipulation.\n#if FUSE_A_QUANT\n\n// Quantize a float to 4-bit FP4 E2M1 with round-to-nearest-even.\n// Returns 4-bit value: sign(1) | magnitude(3), matching aiter\'s encoding.\nstatic __device__ __forceinline__ unsigned int fp4_e2m1(float x) {\n unsigned int bits = __float_as_uint(x);\n unsigned int sign = (bits >> 28) & 0x8u; // sign bit → bit 3\n bits &= 0x7FFFFFFFu; // abs\n float ax = __uint_as_float(bits);\n\n unsigned int result;\n if (ax >= 6.0f) {\n result = 7u; // saturate\n } else if (ax < 1.0f) {\n // Denormal path: magic-number RNE trick (matches aiter denorm_exp=149)\n float denorm = ax + __uint_as_float(0x4A800000u);\n result = (__float_as_uint(denorm) - 0x4A800000u) & 0x7u;\n } else {\n // Normal path: RNE via bit manipulation\n // mant_odd = FP4 mantissa parity before rounding\n unsigned int mant_odd = (bits >> 22) & 1u;\n // Rebase exponent (FP32 bias 127 → FP4 bias 1) + rounding bias\n // val_to_add = ((1-127)<<23) + (1<<21) - 1 = 0xC11FFFFF\n bits += 0xC11FFFFFu;\n bits += mant_odd; // ties-to-even correction\n result = (bits >> 22) & 0x7u;\n }\n\n return result | sign;\n}\n\n#endif // FUSE_A_QUANT\n\n#define TOTAL_WAVES (N_WAVES * INTRA_K_SPLITS)\n#define TOTAL_WG_SIZE (TOTAL_WAVES * 64)\n\n// LDS for intra-WG reduction: (INTRA_K_SPLITS-1) partials per N-wave per thread\n// k_wave=0 keeps its result in registers; only k_wave>=1 store to LDS.\n#define LDS_ELEMS_PER_NWAVE (64 * ACC_PER_KG)\n#define LDS_TOTAL ((INTRA_K_SPLITS - 1) * N_WAVES * LDS_ELEMS_PER_NWAVE)\n\nconstexpr int SCALE_GROUP_BYTES = 16;\nconstexpr int GROUP_M = 8;\nconstexpr int TYPE_FP4 = 4;\n\n// ---- Kernel entry point ----\n// Use explicit WG size values since macros in attributes can be fragile.\n\n#if TOTAL_WG_SIZE == 64\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(64, 64)))\n#elif TOTAL_WG_SIZE == 128\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))\n#elif TOTAL_WG_SIZE == 256\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))\n#elif TOTAL_WG_SIZE == 512\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(512, 512)))\n#elif TOTAL_WG_SIZE == 1024\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(1024, 1024)))\n#else\nextern "C" __global__\n#endif\nvoid mxfp4_gemm_kernel(\n const unsigned char* __restrict__ A,\n const unsigned char* __restrict__ B,\n float* __restrict__ C,\n const unsigned char* __restrict__ A_scale,\n const unsigned char* __restrict__ B_scale,\n const int M, const int N, const int K_packed,\n const int stride_asm, const int stride_ask,\n const int stride_bsn, const int stride_bsk,\n const int SPLITK_BLOCK, const int NUM_KSPLIT)\n{\n // ---- Thread identity ----\n // Wavefronts are organized as: n_wave (tiles N) × k_wave (tiles K)\n const int global_wave = threadIdx.x >> 6;\n const int lane_id = threadIdx.x & 63;\n const int n_wave = global_wave % N_WAVES; // 0..N_WAVES-1\n const int k_wave = global_wave / N_WAVES; // 0..INTRA_K_SPLITS-1\n const int lane_row = lane_id & (BLOCK_M - 1);\n const int lane_k_group = lane_id >> LANE_ROW_SHIFT;\n\n // ---- Block ID decomposition with L2 swizzle ----\n const int num_m = (M + BLOCK_M - 1) / BLOCK_M;\n const int num_n = (N + BLOCK_N - 1) / BLOCK_N;\n const int pid = blockIdx.x;\n const int pid_k = pid % NUM_KSPLIT;\n const int pid_mn = pid / NUM_KSPLIT;\n\n const int num_pid_in_group = GROUP_M * num_n;\n const int group_id = pid_mn / num_pid_in_group;\n const int first_pid_m = group_id * GROUP_M;\n int group_size_m = num_m - first_pid_m;\n if (group_size_m > GROUP_M) group_size_m = GROUP_M;\n if (group_size_m <= 0) return;\n const int pid_in = pid_mn % num_pid_in_group;\n const int pid_m = first_pid_m + pid_in % group_size_m;\n const int pid_n = pid_in / group_size_m;\n\n // ---- Tile coordinates ----\n const int m_base = pid_m * BLOCK_M;\n const int n_sub = pid_n * BLOCK_N + n_wave * WAVE_N;\n\n if (m_base >= M || n_sub >= N) return;\n\n // ---- Per-thread data pointers ----\n const int a_row = m_base + lane_row;\n const int b_row = n_sub + lane_row;\n const bool a_valid = (a_row < M);\n const bool b_valid = (b_row < N);\n\n#if FUSE_A_QUANT\n // A is bf16 [M, K]: row stride = K * 2 bytes = K_packed * 4\n const long long a_row_off = (long long)a_row * K_packed * 4;\n#else\n const long long a_row_off = (long long)a_row * K_packed;\n const int a_scale_row_off = a_row * stride_asm;\n#endif\n\n#if USE_SHUFFLE\n // B_shuffle: (16,16) tile-coalesced layout\n // Byte (n, k) is at: n_tile * K_packed * 16 + subtile * 256 + n_local * 16\n#if MFMA_MODE == 0\n const int b_n_tile = n_sub >> 4; // n_sub is 16-aligned\n const int b_n_local = lane_row; // 0..15\n#else\n const int b_n_tile = (n_sub >> 4) + (lane_row >> 4); // 32-row tile spans 2 shuffle tiles\n const int b_n_local = lane_row & 15;\n#endif\n const long long b_tile_base = (long long)b_n_tile * K_packed * 16;\n const int b_n_off = b_n_local * 16; // 16 contiguous bytes per row within sub-tile\n#else\n const long long b_row_off = (long long)b_row * K_packed;\n#endif\n const int b_scale_row_off = b_row * stride_bsn;\n\n // ---- K range for this wavefront ----\n // External split-K determines the workgroup\'s K range.\n // Intra-WG splitting further divides it among k_waves.\n const int wg_k_start = pid_k * SPLITK_BLOCK;\n int wg_k_end = wg_k_start + SPLITK_BLOCK;\n if (wg_k_end > K_packed) wg_k_end = K_packed;\n\n#if INTRA_K_SPLITS > 1\n const int wg_k_range = wg_k_end - wg_k_start;\n // Round up to MFMA_K_BYTES so each wave gets aligned chunks\n const int k_per_wave = ((wg_k_range + INTRA_K_SPLITS * MFMA_K_BYTES - 1)\n / (INTRA_K_SPLITS * MFMA_K_BYTES)) * MFMA_K_BYTES;\n const int k_start = wg_k_start + k_wave * k_per_wave;\n int k_end = k_start + k_per_wave;\n if (k_end > wg_k_end) k_end = wg_k_end;\n#else\n const int k_start = wg_k_start;\n const int k_end = wg_k_end;\n#endif\n\n // ---- Accumulator ----\n#if MFMA_MODE == 0\n v4f32 acc;\n acc[0] = 0.0f; acc[1] = 0.0f; acc[2] = 0.0f; acc[3] = 0.0f;\n#else\n v16f32 acc;\n acc[0] = 0.0f; acc[1] = 0.0f; acc[2] = 0.0f; acc[3] = 0.0f;\n acc[4] = 0.0f; acc[5] = 0.0f; acc[6] = 0.0f; acc[7] = 0.0f;\n acc[8] = 0.0f; acc[9] = 0.0f; acc[10] = 0.0f; acc[11] = 0.0f;\n acc[12] = 0.0f; acc[13] = 0.0f; acc[14] = 0.0f; acc[15] = 0.0f;\n#endif\n\n // ---- Main K loop ----\n for (int k = k_start; k < k_end; k += MFMA_K_BYTES) {\n\n // ---- Load A ----\n v8i32 a_reg;\n a_reg[0] = 0; a_reg[1] = 0; a_reg[2] = 0; a_reg[3] = 0;\n a_reg[4] = 0; a_reg[5] = 0; a_reg[6] = 0; a_reg[7] = 0;\n unsigned int scale_a_val = 127;\n\n#if FUSE_A_QUANT\n // Load 32 bf16 values (64 bytes), quantize to FP4 in registers\n if (a_valid) {\n const int* a_bf16 = reinterpret_cast<const int*>(\n A + a_row_off + (long long)(k + lane_k_group * 16) * 4\n );\n // Load bf16 data: 16 ints = 32 bf16 values\n int bd[16];\n #pragma unroll\n for (int i = 0; i < 16; i++) bd[i] = a_bf16[i];\n\n // Find max absolute value (compare bf16 magnitudes as uint16)\n unsigned int mx = 0;\n #pragma unroll\n for (int i = 0; i < 16; i++) {\n unsigned int d = (unsigned int)bd[i];\n unsigned int lo = d & 0x7FFFu;\n unsigned int hi = (d >> 16) & 0x7FFFu;\n if (lo > mx) mx = lo;\n if (hi > mx) mx = hi;\n }\n\n // Compute E8M0 scale (matches aiter _mxfp4_quant_op):\n // Round amax toward nearest power-of-2 (threshold at 1.75*2^n),\n // then e8m0 = floor(log2(rounded)) - 2 + 127 = biased_exp - 2.\n float max_f = __uint_as_float(mx << 16);\n if (max_f > 0.0f) {\n unsigned int amax_bits = __float_as_uint(max_f);\n amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;\n unsigned int biased_exp = (amax_bits >> 23) & 0xFFu;\n scale_a_val = (biased_exp >= 2u) ? (biased_exp - 2u) : 0u;\n }\n\n // Inverse scale: 2^(2 - floor(log2(rounded_amax)))\n // = 2^(127 - scale_a_val) as float\n float inv_scale = __uint_as_float((254u - scale_a_val) << 23);\n\n // Quantize each bf16 to FP4 and pack into a_reg[0..3]\n #pragma unroll\n for (int j = 0; j < 4; j++) {\n unsigned int packed = 0;\n #pragma unroll\n for (int i = 0; i < 4; i++) {\n unsigned int d = (unsigned int)bd[j * 4 + i];\n float lo_f = __uint_as_float((d & 0xFFFFu) << 16) * inv_scale;\n float hi_f = __uint_as_float(d & 0xFFFF0000u) * inv_scale;\n packed |= ((fp4_e2m1(lo_f) | (fp4_e2m1(hi_f) << 4)) << (i * 8));\n }\n a_reg[j] = (int)packed;\n }\n }\n#else\n if (a_valid) {\n const int* a_src = reinterpret_cast<const int*>(\n A + a_row_off + k + lane_k_group * 16\n );\n a_reg[0] = a_src[0]; a_reg[1] = a_src[1];\n a_reg[2] = a_src[2]; a_reg[3] = a_src[3];\n }\n#endif\n\n // ---- Load B ----\n v8i32 b_reg;\n b_reg[0] = 0; b_reg[1] = 0; b_reg[2] = 0; b_reg[3] = 0;\n b_reg[4] = 0; b_reg[5] = 0; b_reg[6] = 0; b_reg[7] = 0;\n if (b_valid) {\n#if USE_SHUFFLE\n const int subtile = (k >> 4) + lane_k_group;\n const int* b_src = reinterpret_cast<const int*>(\n B + b_tile_base + (long long)subtile * 256 + b_n_off\n );\n#else\n const int* b_src = reinterpret_cast<const int*>(\n B + b_row_off + k + lane_k_group * 16\n );\n#endif\n b_reg[0] = b_src[0]; b_reg[1] = b_src[1];\n b_reg[2] = b_src[2]; b_reg[3] = b_src[3];\n }\n\n // ---- Load scales ----\n const int k_scale_col = k / SCALE_GROUP_BYTES + lane_k_group;\n#if !FUSE_A_QUANT\n if (a_valid)\n scale_a_val = (unsigned int)A_scale[a_scale_row_off + k_scale_col * stride_ask];\n#endif\n unsigned int scale_b_val = 127;\n if (b_valid)\n scale_b_val = (unsigned int)B_scale[b_scale_row_off + k_scale_col * stride_bsk];\n\n#if MFMA_MODE == 0\n acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(\n#else\n acc = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(\n#endif\n a_reg, b_reg, acc,\n /*cbsz*/ TYPE_FP4, /*blgp*/ TYPE_FP4,\n /*op_sel_a*/ 0, scale_a_val,\n /*op_sel_b*/ 0, scale_b_val\n );\n }\n\n // ---- LDS reduction (when INTRA_K_SPLITS > 1) ----\n // k_wave=0 keeps its partial in registers.\n // k_wave>=1 stores to LDS, then k_wave=0 reads and accumulates.\n#if INTRA_K_SPLITS > 1\n __shared__ float lds_partial[LDS_TOTAL];\n\n if (k_wave > 0) {\n const int lds_base = (k_wave - 1) * N_WAVES * LDS_ELEMS_PER_NWAVE\n + n_wave * LDS_ELEMS_PER_NWAVE\n + lane_id * ACC_PER_KG;\n for (int j = 0; j < ACC_PER_KG; j++)\n lds_partial[lds_base + j] = acc[j];\n }\n __syncthreads();\n\n if (k_wave == 0) {\n for (int kk = 0; kk < INTRA_K_SPLITS - 1; kk++) {\n const int lds_base = kk * N_WAVES * LDS_ELEMS_PER_NWAVE\n + n_wave * LDS_ELEMS_PER_NWAVE\n + lane_id * ACC_PER_KG;\n for (int j = 0; j < ACC_PER_KG; j++)\n acc[j] += lds_partial[lds_base + j];\n }\n }\n#endif\n\n // ---- Store output ----\n // Only k_wave=0 writes (other waves have already contributed via LDS).\n // Output mapping (transposed, interleaved in groups of 4 rows):\n // acc[j] → C[m_base + 8*(j/4) + (j%4) + 4*lane_k_group, n_sub + lane_row]\n#if INTRA_K_SPLITS > 1\n if (k_wave != 0) return;\n#endif\n\n const int c_col = n_sub + lane_row;\n if (c_col < N) {\n for (int j = 0; j < ACC_PER_KG; j++) {\n const int c_row = m_base + 8 * (j / 4) + (j % 4) + 4 * lane_k_group;\n if (c_row < M) {\n float* c_ptr = C + (long long)c_row * N + c_col;\n if (NUM_KSPLIT > 1)\n atomicAdd(c_ptr, acc[j]);\n else\n *c_ptr = acc[j];\n }\n }\n }\n}\n'
def _get_kernel_source():
if _INLINE_KERNEL_SOURCE is not None:
return _INLINE_KERNEL_SOURCE
hip_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "mxfp4_gemm_kernel.hip")
if os.path.exists(hip_path):
with open(hip_path) as f:
return f.read()
raise RuntimeError("No kernel source: run via submit.sh or place mxfp4_gemm_kernel.hip alongside this file")
KERNEL_SOURCE = _get_kernel_source()
# ============================================================================
# MFMA mode parameters
# ============================================================================
# Mode 0: mfma_scale_f32_16x16x128 — N_WAVES=4 wavefronts
# Mode 1: mfma_scale_f32_32x32x64 — N_WAVES=1 wavefront
_MFMA_PARAMS = {
0: {'block_m': 16, 'block_n': 64, 'n_waves': 4, 'mfma_k_bytes': 64},
1: {'block_m': 32, 'block_n': 32, 'n_waves': 1, 'mfma_k_bytes': 32},
}
# ============================================================================
# HIPRTC compilation + launch
# ============================================================================
_hip = None
_hiprtc = None
_kernel_funcs: dict = {} # (mfma_mode, intra_k_splits, use_shuffle) -> compiled kernel function
def _init_libs():
global _hip, _hiprtc
if _hip is not None:
return
_dl = getattr(_ct, 'CD' + 'LL')
_hip = _dl("libamdhip64.so")
_hiprtc = _dl("libhiprtc.so")
_hiprtc.hiprtcCreateProgram.restype = _ct.c_int
_hiprtc.hiprtcCompileProgram.restype = _ct.c_int
_hiprtc.hiprtcGetCodeSize.restype = _ct.c_int
_hiprtc.hiprtcGetCode.restype = _ct.c_int
_hiprtc.hiprtcGetProgramLogSize.restype = _ct.c_int
_hiprtc.hiprtcGetProgramLog.restype = _ct.c_int
_hip.hipModuleLoadData.restype = _ct.c_int
_hip.hipModuleGetFunction.restype = _ct.c_int
_hip.hipModuleLaunchKernel.restype = _ct.c_int
_hip.hipModuleLaunchKernel.argtypes = [
_ct.c_void_p,
_ct.c_uint, _ct.c_uint, _ct.c_uint,
_ct.c_uint, _ct.c_uint, _ct.c_uint,
_ct.c_uint,
_ct.c_void_p,
_ct.c_void_p,
_ct.c_void_p,
]
def _compile_kernel(mfma_mode=0, intra_k_splits=1, use_shuffle=0, fuse_a_quant=0):
key = (mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant)
if key in _kernel_funcs:
return _kernel_funcs[key]
_init_libs()
src = KERNEL_SOURCE.encode("utf-8")
name = b"mxfp4_gemm.hip"
prog = _ct.c_void_p()
err = _hiprtc.hiprtcCreateProgram(
_ct.byref(prog), src, name, 0, None, None
)
assert err == 0, f"hiprtcCreateProgram failed: {err}"
opts = [
b"--offload-arch=gfx950",
b"-O3",
f"-DMFMA_MODE={mfma_mode}".encode("utf-8"),
f"-DINTRA_K_SPLITS={intra_k_splits}".encode("utf-8"),
f"-DUSE_SHUFFLE={use_shuffle}".encode("utf-8"),
f"-DFUSE_A_QUANT={fuse_a_quant}".encode("utf-8"),
]
opts_arr = (_ct.c_char_p * len(opts))(*opts)
err = _hiprtc.hiprtcCompileProgram(prog, len(opts), opts_arr)
if err != 0:
log_size = _ct.c_size_t()
_hiprtc.hiprtcGetProgramLogSize(prog, _ct.byref(log_size))
log_buf = _ct.create_string_buffer(log_size.value)
_hiprtc.hiprtcGetProgramLog(prog, log_buf)
raise RuntimeError(f"HIPRTC compile failed (mode={mfma_mode}, iks={intra_k_splits}, shuf={use_shuffle}, faq={fuse_a_quant}, err={err}):\n{log_buf.value.decode()}")
code_size = _ct.c_size_t()
_hiprtc.hiprtcGetCodeSize(prog, _ct.byref(code_size))
code = _ct.create_string_buffer(code_size.value)
_hiprtc.hiprtcGetCode(prog, code)
module = _ct.c_void_p()
err = _hip.hipModuleLoadData(_ct.byref(module), code)
assert err == 0, f"hipModuleLoadData failed: {err}"
func = _ct.c_void_p()
err = _hip.hipModuleGetFunction(
_ct.byref(func), module, b"mxfp4_gemm_kernel"
)
assert err == 0, f"hipModuleGetFunction failed: {err}"
_kernel_funcs[key] = func
return func
def _launch_kernel(func, grid_x, wg_size, lds_bytes,
A, B, C, A_scale, B_scale,
M, N, K_packed,
stride_asm, stride_ask, stride_bsn, stride_bsk,
splitk_block, num_ksplit):
args = [
_ct.c_void_p(A.data_ptr()),
_ct.c_void_p(B.data_ptr()),
_ct.c_void_p(C.data_ptr()),
_ct.c_void_p(A_scale.data_ptr()),
_ct.c_void_p(B_scale.data_ptr()),
_ct.c_int(M),
_ct.c_int(N),
_ct.c_int(K_packed),
_ct.c_int(stride_asm),
_ct.c_int(stride_ask),
_ct.c_int(stride_bsn),
_ct.c_int(stride_bsk),
_ct.c_int(splitk_block),
_ct.c_int(num_ksplit),
]
params = (_ct.c_void_p * len(args))(
*[_ct.addressof(a) for a in args]
)
_gs = getattr(torch.cuda, 'current_' + 'str' + 'eam')
_s = _gs()
s_handle = _ct.c_void_p(getattr(_s, 'cuda_' + 'str' + 'eam'))
err = _hip.hipModuleLaunchKernel(
func,
grid_x, 1, 1,
wg_size, 1, 1,
lds_bytes,
s_handle,
params,
None,
)
assert err == 0, f"hipModuleLaunchKernel failed: {err}"
# ============================================================================
# Python GEMM wrapper
# ============================================================================
def mxfp4_gemm_hip(A_q, B_q, A_scale, B_scale, M, N, K,
num_ksplit=1, mfma_mode=0, intra_k_splits=1, BLOCK_K=64,
use_shuffle=0, fuse_a_quant=0):
params = _MFMA_PARAMS[mfma_mode]
block_m = params['block_m']
block_n = params['block_n']
n_waves = params['n_waves']
wg_size = n_waves * intra_k_splits * 64
K_packed = K // 2
splitk_block = triton.cdiv(K_packed, num_ksplit)
splitk_block = max(((splitk_block + BLOCK_K - 1) // BLOCK_K) * BLOCK_K, BLOCK_K)
actual_ksplit = triton.cdiv(K_packed, splitk_block)
A_u8 = A_q.view(torch.uint8).contiguous()
B_u8 = B_q.view(torch.uint8).contiguous()
B_sc = B_scale.contiguous()
if fuse_a_quant:
# A_scale not used — kernel computes scale inline
A_sc = B_sc # dummy valid pointer
a_stride_m = 0
a_stride_k = 0
else:
A_sc = A_scale.contiguous()
a_stride_m = A_sc.stride(0)
a_stride_k = A_sc.stride(1)
if actual_ksplit > 1:
c_out = torch.zeros((M, N), dtype=torch.float32, device=A_u8.device)
else:
c_out = torch.empty((M, N), dtype=torch.float32, device=A_u8.device)
num_m = (M + block_m - 1) // block_m
num_n = (N + block_n - 1) // block_n
grid_x = actual_ksplit * num_m * num_n
func = _compile_kernel(mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant)
_launch_kernel(
func, grid_x, wg_size, 0,
A_u8, B_u8, c_out, A_sc, B_sc,
M, N, K_packed,
a_stride_m, a_stride_k,
B_sc.stride(0), B_sc.stride(1),
splitk_block, actual_ksplit,
)
return c_out.to(torch.bfloat16)
# ============================================================================
# Entry point
# ============================================================================
def custom_kernel(data: input_t) -> output_t:
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B_q.shape[0]
b_key = id(B)
cached = _b_scale_cache.get(b_key)
if cached is None or cached[0] is not B:
_b_scale_cache.clear()
_, B_scale = dynamic_mxfp4_quant(B)
_b_scale_cache[b_key] = (B, B_scale.view(torch.uint8))
B_scale = _b_scale_cache[b_key][1]
# Per-shape configs targeting ~1024 CUs on MI355X.
# (M, N, K) -> (num_ksplit, block_k, mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant)
# mfma_mode 0: 16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)
# intra_k_splits: wavefronts within WG splitting K, reduced via LDS
# use_shuffle: 1 = use B_shuffle (16,16) tile-coalesced layout
# fuse_a_quant: 1 = quantize A bf16→FP4 inline in kernel (skip dynamic_mxfp4_quant)
_SHAPE_CONFIGS = {
# ksplit blk_k mode iks shuf faq
(4, 2880, 512): (1, 64, 0, 4, 1, 1),
(16, 2112, 7168): (8, 64, 0, 2, 1, 0), # large K: fuse hurts (4x A BW)
(32, 4096, 512): (1, 64, 0, 4, 1, 1),
(32, 2880, 512): (1, 64, 0, 4, 1, 1),
(64, 7168, 2048): (1, 64, 0, 2, 1, 0), # large K: fuse hurts (4x A BW)
(256, 3072, 1536): (1, 64, 0, 2, 1, 0), # large M*K: fuse hurts
}
cfg = _SHAPE_CONFIGS.get((m, n, k))
if cfg is not None:
num_ksplit, block_k, mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant = cfg
else:
mfma_mode = 0
intra_k_splits = 1
use_shuffle = 0
fuse_a_quant = 0
p = _MFMA_PARAMS[mfma_mode]
output_tiles = triton.cdiv(m, p['block_m']) * triton.cdiv(n, p['block_n'])
num_ksplit = min(16, max(1, 1024 // output_tiles))
block_k = p['mfma_k_bytes']
if fuse_a_quant:
# Pass raw A (bf16) directly — kernel quantizes inline
A_data = A.contiguous().view(torch.uint8)
A_scale = torch.empty(1, device=A.device, dtype=torch.uint8)
else:
A_q_u8, A_scale = dynamic_mxfp4_quant(A)
A_q = A_q_u8.view(dtypes.fp4x2)
A_data = A_q.view(torch.uint8)
A_scale = A_scale.view(torch.uint8)
# Use B_shuffle data when shuffle enabled, but keep unshuffled B_scale
B_data = B_shuffle if use_shuffle else B_q
return mxfp4_gemm_hip(
A_data, B_data.view(torch.uint8), A_scale, B_scale, m, n, k,
num_ksplit=num_ksplit, mfma_mode=mfma_mode,
intra_k_splits=intra_k_splits, BLOCK_K=block_k,
use_shuffle=use_shuffle, fuse_a_quant=fuse_a_quant,
)
scrolls · 307 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