submission 701074
dungthai414 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 543 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-701074?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:9ecaebbf210ea0164ea53f7ee1a4c92a19d09830afd110d283a5e352639da1de
license declaredunknown
license concludedunknown
authorsdungthai414
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.shared-memory
__shared__ fp4x2_t A_s[2][BM][BK / 2];tile-k = 512
constexpr int BK = 512;tile-m = 32
constexpr int BM = 32;tile-n = 128
constexpr int BN = 128;Kernel source
submission.py543 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
"""
from task import input_t, output_t
import torch
from torch.utils.cpp_extension import load_inline
import os
if "PYTORCH_ROCM_ARCH" not in os.environ:
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"
kernel_cpp = r"""
#pragma once
#undef __HIP_NO_HALF_OPERATORS__
#undef __HIP_NO_HALF_CONVERSIONS__
#include <hip/hip_fp16.h>
#include <hip/hip_runtime.h>
#include <hip/hip_fp8.h>
#include <hip/hip_fp4.h>
#include <hip/hip_bf16.h>
#include <hip/hip_ext_ocp.h>
#include <pybind11/pybind11.h>
#include <math.h>
#include <stdint.h>
using fp4x2_t = __amd_fp4x2_storage_t;
using fp4x64_t = fp4x2_t __attribute__((ext_vector_type(32)));
using bfp16 = __hip_bfloat16;
using floatx4_t = float __attribute__((ext_vector_type(4)));
using i32x4 = int32_t __attribute__((ext_vector_type(4)));
using u32x4 = uint32_t __attribute__((ext_vector_type(4)));
using as3_uint32_ptr = uint32_t __attribute__((address_space(3)))*;
static constexpr inline int ceil_div(int x, int y) { return (x + y - 1) / y; }
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
i32x4 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset, int soffset, int offset, int aux)
__asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource {
uint64_t ptr;
uint32_t range;
uint32_t config;
};
__device__ inline i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
return *reinterpret_cast<const i32x4*>(&rsrc);
}
__device__ inline as3_uint32_ptr as_lds_u32_ptr(void* p) {
return reinterpret_cast<as3_uint32_ptr>(reinterpret_cast<uintptr_t>(p));
}
__host__ __device__ inline bfp16 fast_f32tob16(float f) {
union { float fp32; uint32_t u32; } u = {f};
u.u32 += 0x7fff + ((u.u32 >> 16) & 1);
union { uint16_t u16; bfp16 bf16; } out;
out.u16 = static_cast<uint16_t>(u.u32 >> 16);
return out.bf16;
}
#define WARP_SIZE 64
__device__ inline floatx4_t mfma_fp4_fp4(fp4x64_t a, fp4x64_t b, floatx4_t c,
uint8_t scale_a, uint8_t scale_b) {
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 4, 4, 0, scale_a, 0, scale_b
);
}
constexpr int BM = 32;
constexpr int BN = 128;
constexpr int BK = 512;
constexpr int WARP_M = 2;
constexpr int WARP_N = 4;
constexpr int BLOCK_SIZE = 512;
constexpr int MFMA_M = 16;
constexpr int MFMA_N = 16;
constexpr int MFMA_K = 128;
constexpr int FRAG_M_PER_WARP = BM / (MFMA_M * WARP_M); // 1
constexpr int FRAG_N_PER_WARP = BN / (MFMA_N * WARP_N); // 4
constexpr int FRAG_K = BK / MFMA_K; // 1 for BK=128
struct BPrefetchBuf {
u32x4 pkt[FRAG_N_PER_WARP][FRAG_K]; // compact 16B packets in regs
uint8_t scale[FRAG_N_PER_WARP][FRAG_K];
};
struct AScaleBuf {
uint8_t scale[FRAG_M_PER_WARP][FRAG_K];
};
// E8M0 1D Physical Offset Mapping
__device__ inline int e8m0_shuf_phys_offset(int r, int c, int padded_C) {
const int r_tile = r / 32;
const int r_in = r % 32;
const int r_in_0 = r_in / 16;
const int r_in_1 = r_in % 16;
const int c_tile = c / 8;
const int c_in = c % 8;
const int c_in_0 = c_in / 4;
const int c_in_1 = c_in % 4;
const int sn_tiles = padded_C / 8;
return r_tile * sn_tiles * 256 +
c_tile * 256 +
c_in_1 * 64 +
r_in_1 * 4 +
c_in_0 * 2 +
r_in_0;
}
__device__ inline u32x4 zero_u32x4() {
return u32x4{0u, 0u, 0u, 0u};
}
__device__ inline fp4x64_t expand_fp4x64_from_pkt16(u32x4 raw) {
fp4x64_t out{};
union {
u32x4 u32;
fp4x2_t x2[16];
} tmp;
tmp.u32 = raw;
#pragma unroll
for (int i = 0; i < 16; ++i) out[i] = tmp.x2[i];
#pragma unroll
for (int i = 16; i < 32; ++i) out[i] = 0;
return out;
}
__device__ inline u32x4 lds_load_u32x4(const fp4x2_t* p) {
return *reinterpret_cast<const u32x4*>(p);
}
__device__ inline size_t bshuf_packet_offset_bytes(
int b_row,
int g_global, // 16-byte logical chunk index
int K_logical
) {
int n_tile = b_row / 16;
int n_in = b_row % 16;
int k_tile = g_global / 2;
int k_chunk = g_global % 2;
int k_tiles_32 = K_logical / 64;
return (size_t)n_tile * k_tiles_32 * 512 +
k_tile * 512 +
k_chunk * 256 +
n_in * 16;
}
__device__ inline void zero_bprefetch(BPrefetchBuf& buf) {
#pragma unroll
for (int j = 0; j < FRAG_N_PER_WARP; ++j) {
#pragma unroll
for (int kk = 0; kk < FRAG_K; ++kk) {
buf.pkt[j][kk] = zero_u32x4();
buf.scale[j][kk] = 0;
}
}
}
__device__ inline void zero_ascale(AScaleBuf& buf) {
#pragma unroll
for (int i = 0; i < FRAG_M_PER_WARP; ++i) {
#pragma unroll
for (int kk = 0; kk < FRAG_K; ++kk) {
buf.scale[i][kk] = 0;
}
}
}
// ---------------------------------------------------------
// Split-K Workspace Cast Kernel
// ---------------------------------------------------------
__global__ void cast_kernel(const float* __restrict__ workspace, bfp16* __restrict__ C, int total_elems) {
int tid = blockIdx.x * blockDim.x + threadIdx.x;
if (tid < total_elems) {
C[tid] = fast_f32tob16(workspace[tid]);
}
}
__global__ __launch_bounds__(BLOCK_SIZE)
void gemm(
const fp4x2_t* __restrict__ A_q,
const fp4x2_t* __restrict__ B_shuf,
const uint8_t* __restrict__ A_scale,
const uint8_t* __restrict__ B_scale,
bfp16* __restrict__ C,
float* __restrict__ workspace,
int M, int N, int K
) {
// ---------------- Split-K Bounds Math ----------------
int chunk_size = BK; // 256
int total_chunks = ceil_div(K, chunk_size);
int chunks_per_split = ceil_div(total_chunks, gridDim.z);
int my_chunk_start = blockIdx.z * chunks_per_split;
int my_chunk_end = my_chunk_start + chunks_per_split;
if (my_chunk_end > total_chunks) my_chunk_end = total_chunks;
int k_start = my_chunk_start * chunk_size;
int k_end = my_chunk_end * chunk_size;
if (k_end > K) k_end = K;
// If this split block has no chunks, exit early
if (k_start >= k_end) return;
// ------------------------------------------------------
const int K32_global = K / 32;
const int PAD_SCALE_COLS = ceil_div(K32_global, 8) * 8;
const int K32_end = k_end / 32;
const int tid = threadIdx.x;
const int wid = tid / WARP_SIZE;
const int lane = tid % WARP_SIZE;
const int warp_m = wid / WARP_N;
const int warp_n = wid % WARP_N;
const int cur_m = blockIdx.y * BM;
const int cur_n = blockIdx.x * BN;
const int row_in_tile = lane & 15;
const int row_group = lane >> 4;
__shared__ fp4x2_t A_s[2][BM][BK / 2];
floatx4_t c_reg[FRAG_M_PER_WARP][FRAG_N_PER_WARP];
#pragma unroll
for (int i = 0; i < FRAG_M_PER_WARP; ++i) {
#pragma unroll
for (int j = 0; j < FRAG_N_PER_WARP; ++j) {
c_reg[i][j] = floatx4_t{0.f, 0.f, 0.f, 0.f};
}
}
i32x4 srcA = make_srsrc(A_q, M * (K / 2) * int(sizeof(fp4x2_t)));
auto prefetch_A_to_lds = [&](int k0) {
const int buf = (k0 / BK) & 1;
constexpr int VEC_BYTES = 16;
constexpr int ROW_BYTES = BK / 2;
constexpr int VECS_PER_ROW = ROW_BYTES / VEC_BYTES;
for (int x = tid; x < BM * VECS_PER_ROW; x += BLOCK_SIZE) {
const int row = x / VECS_PER_ROW;
const int vec = x % VECS_PER_ROW;
const int gm = cur_m + row;
if (gm < M && k0 < k_end) {
const int g_byte_off = gm * (K / 2) + (k0 / 2) + vec * VEC_BYTES;
llvm_amdgcn_raw_buffer_load_lds(
srcA,
as_lds_u32_ptr((void*)&A_s[buf][row][vec * VEC_BYTES]),
16,
g_byte_off,
0, 0, 0
);
} else {
*reinterpret_cast<u32x4*>(&A_s[buf][row][vec * VEC_BYTES]) = zero_u32x4();
}
}
};
auto prefetch_A_scales = [&](int k0, AScaleBuf& out) {
zero_ascale(out);
#pragma unroll
for (int kk = 0; kk < FRAG_K; ++kk) {
const int g_local = kk * (MFMA_K / 32) + row_group;
const int g_global = (k0 / 32) + g_local;
#pragma unroll
for (int i = 0; i < FRAG_M_PER_WARP; ++i) {
const int a_row_local = warp_m * (BM / WARP_M) + i * MFMA_M + row_in_tile;
const int a_row_global = cur_m + a_row_local;
out.scale[i][kk] =
(a_row_global < M && g_global < K32_end)
? A_scale[e8m0_shuf_phys_offset(a_row_global, g_global, PAD_SCALE_COLS)]
: 0;
}
}
};
auto prefetch_B_regs = [&](int k0, BPrefetchBuf& out) {
zero_bprefetch(out);
#pragma unroll
for (int kk = 0; kk < FRAG_K; ++kk) {
const int g_local = kk * (MFMA_K / 32) + row_group;
const int g_global = (k0 / 32) + g_local;
#pragma unroll
for (int j = 0; j < FRAG_N_PER_WARP; ++j) {
const int b_row_local = warp_n * (BN / WARP_N) + j * MFMA_N + row_in_tile;
const int b_row_global = cur_n + b_row_local;
if (b_row_global < N && g_global < K32_end) {
const size_t off_bytes = bshuf_packet_offset_bytes(b_row_global, g_global, K); // Global K dictates physical memory offsets
const char* base = reinterpret_cast<const char*>(B_shuf);
out.pkt[j][kk] = *reinterpret_cast<const u32x4*>(base + off_bytes);
} else {
out.pkt[j][kk] = zero_u32x4();
}
out.scale[j][kk] =
(b_row_global < N && g_global < K32_end)
? B_scale[e8m0_shuf_phys_offset(b_row_global, g_global, PAD_SCALE_COLS)]
: 0;
}
}
};
auto mfma_compute = [&](int k0, const BPrefetchBuf& bbuf, const AScaleBuf& asbuf) {
const int buf = (k0 / BK) & 1;
#pragma unroll
for (int kk = 0; kk < FRAG_K; ++kk) {
fp4x64_t b_frag[FRAG_N_PER_WARP];
uint8_t b_scl[FRAG_N_PER_WARP];
#pragma unroll
for (int j = 0; j < FRAG_N_PER_WARP; ++j) {
b_frag[j] = expand_fp4x64_from_pkt16(bbuf.pkt[j][kk]);
b_scl[j] = bbuf.scale[j][kk];
}
#pragma unroll
for (int i = 0; i < FRAG_M_PER_WARP; ++i) {
const int a_row_local = warp_m * (BM / WARP_M) + i * MFMA_M + row_in_tile;
const int g_local = kk * (MFMA_K / 32) + row_group;
const int a_byte_local = g_local * 16;
const u32x4 a_pkt = lds_load_u32x4(&A_s[buf][a_row_local][a_byte_local]);
const fp4x64_t a_frag = expand_fp4x64_from_pkt16(a_pkt);
const uint8_t sa = asbuf.scale[i][kk];
#pragma unroll
for (int j = 0; j < FRAG_N_PER_WARP; ++j) {
c_reg[i][j] = mfma_fp4_fp4(a_frag, b_frag[j], c_reg[i][j], sa, b_scl[j]);
}
}
}
};
BPrefetchBuf b_cur, b_next;
AScaleBuf as_cur, as_next;
int k0 = k_start;
// Prologue
prefetch_A_to_lds(k0);
prefetch_B_regs(k0, b_cur);
prefetch_A_scales(k0, as_cur);
asm volatile("s_waitcnt vmcnt(0)");
__builtin_amdgcn_s_barrier();
// Steady state
for (; k0 + BK < k_end; k0 += BK) {
prefetch_A_to_lds(k0 + BK);
prefetch_B_regs(k0 + BK, b_next);
prefetch_A_scales(k0 + BK, as_next);
mfma_compute(k0, b_cur, as_cur);
asm volatile("s_waitcnt vmcnt(0)");
__builtin_amdgcn_s_barrier();
b_cur = b_next;
as_cur = as_next;
}
// Final tile for this split
mfma_compute(k0, b_cur, as_cur);
// ==========================================
// Epilogue Selection
// ==========================================
if (workspace != nullptr) {
// Atomic FP32 writes for Split-K merging
#pragma unroll
for (int i = 0; i < FRAG_M_PER_WARP; ++i) {
const int out_row_base = cur_m + warp_m * (BM / WARP_M) + i * MFMA_M + row_group * 4;
#pragma unroll
for (int j = 0; j < FRAG_N_PER_WARP; ++j) {
const int out_col = cur_n + warp_n * (BN / WARP_N) + j * MFMA_N + row_in_tile;
if (out_col < N) {
#pragma unroll
for (int t = 0; t < 4; ++t) {
const int out_row = out_row_base + t;
if (out_row < M) {
atomicAdd(&workspace[out_row * N + out_col], c_reg[i][j][t]);
}
}
}
}
}
} else {
// Standard Scalar bfp16 writes
#pragma unroll
for (int i = 0; i < FRAG_M_PER_WARP; ++i) {
const int out_row_base = cur_m + warp_m * (BM / WARP_M) + i * MFMA_M + row_group * 4;
#pragma unroll
for (int j = 0; j < FRAG_N_PER_WARP; ++j) {
const int out_col = cur_n + warp_n * (BN / WARP_N) + j * MFMA_N + row_in_tile;
if (out_col < N) {
#pragma unroll
for (int t = 0; t < 4; ++t) {
const int out_row = out_row_base + t;
if (out_row < M) {
C[out_row * N + out_col] = fast_f32tob16(c_reg[i][j][t]);
}
}
}
}
}
}
}
void run(
uintptr_t a_ptr,
uintptr_t b_shuf_ptr,
uintptr_t a_scale_ptr,
uintptr_t b_scale_ptr,
uintptr_t c_ptr,
uintptr_t workspace_ptr,
int M, int N, int K, int k_split
) {
const auto* d_A = reinterpret_cast<const fp4x2_t*>(a_ptr);
const auto* d_B_shuf = reinterpret_cast<const fp4x2_t*>(b_shuf_ptr);
const auto* d_A_scale = reinterpret_cast<const uint8_t*>(a_scale_ptr);
const auto* d_B_scale = reinterpret_cast<const uint8_t*>(b_scale_ptr);
auto* d_C = reinterpret_cast<bfp16*>(c_ptr);
float* d_workspace = reinterpret_cast<float*>(workspace_ptr);
dim3 threads(BLOCK_SIZE);
dim3 blocks(ceil_div(N, BN), ceil_div(M, BM), k_split);
hipLaunchKernelGGL(gemm, blocks, threads, 0, 0,
d_A, d_B_shuf, d_A_scale, d_B_scale, d_C, d_workspace, M, N, K);
if (k_split > 1) {
int total_elems = M * N;
int threads_cast = 256;
int blocks_cast = ceil_div(total_elems, threads_cast);
hipLaunchKernelGGL(cast_kernel, dim3(blocks_cast), dim3(threads_cast), 0, 0,
d_workspace, d_C, total_elems);
}
}
PYBIND11_MODULE(mxfp4, m) {
m.def("run", &run, "HIP MXFP4 GEMM kernel");
}
"""
hip_module = load_inline(
name="mxfp4",
cpp_sources="",
cuda_sources=kernel_cpp,
with_cuda=True,
verbose=False,
extra_cuda_cflags=["-std=c++20", "-O3"],
no_implicit_headers=True,
)
def custom_kernel(data: input_t) -> output_t:
import aiter
from aiter import QuantType, dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
def _quant_mxfp4(x, shuffle=True):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
if shuffle:
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
B = B.contiguous()
m, k = A.shape
n, _ = B.shape
# hip module run
A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
# ---------------------------------------------------------
# Split-K Occupancy Tuning for CDNA4 (MI355X)
# ---------------------------------------------------------
# MI355X has 256 CUs. We want to guarantee at least 1 thread block
# per CU to prevent grid starvation on small shapes.
target_blocks = 256
b = ((n + 511) // 512) * ((m + 31) // 32)
k_split = 1
if b < target_blocks:
k_split = min(16, (target_blocks + b - 1) // b)
# Cap split by available 256-element chunks in K (BK=256)
k_chunks = k // 256
if k_chunks > 0:
k_split = min(k_split, k_chunks)
else:
k_split = 1
if k_split > 1:
workspace = torch.zeros((m, n), dtype=torch.float32, device=A.device)
workspace_ptr = workspace.data_ptr()
else:
workspace_ptr = 0
hip_module.run(
A_q.data_ptr(),
B_shuffle.data_ptr(),
A_scale_sh.data_ptr(),
B_scale_sh.data_ptr(),
C.data_ptr(),
workspace_ptr,
m, n, k, k_split
)
return Cscrolls · 543 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