submission 703548
Siuuuuuuu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 457 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-703548?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:b6aa236c44fa40927b9c06591a18c196579e2bcfdc56e5fa5b4a393e0790c599
license declaredunknown
license concludedunknown
authorsSiuuuuuuu
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
PYBIND11_MODULE(mxfp4, m) { m.def("run", &run, "HIP MXFP4 GEMM"); }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.py457 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
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);
}
// ── Tile / warp geometry — kept identical to original to preserve all
// scale-index and LDS-offset arithmetic that was already correct. ──────────
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); // 2 (was 4 — see note*)
// *Original had BN/MFMA_N/WARP_N = 128/16/4 = 2, not 4 as the comment said.
constexpr int FRAG_K = BK / MFMA_K; // 4
// ── FIX: full 32-element FP4 fragment ───────────────────────────────────────
// The original expand_fp4x64_from_pkt16() loaded one 16-byte packet (16 fp4x2
// = 32 FP4 nibbles) and left the upper 32 slots of fp4x64_t zeroed.
// __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4 reads 128 nibbles per
// lane from each operand (fp4x64_t holds exactly 64 fp4x2 = 128 nibbles).
// Zero-padding the upper half means every MFMA accumulates only the first
// 64 nibbles; the second 64 are treated as zero — silently halving output.
//
// Fix: load TWO consecutive 16-byte packets and fill all 32 fp4x2 slots.
// The two packets are adjacent in the shuffled layout:
// packet0: byte offset off
// packet1: byte offset off + 16
// (within the same n_in row and k_chunk — they are consecutive in memory
// because bshuf_packet_offset_bytes already places them that way).
__device__ inline fp4x64_t expand_fp4x64_full(const void* base, size_t off) {
fp4x64_t out{};
union { u32x4 u32; fp4x2_t x2[16]; } lo, hi;
lo.u32 = *reinterpret_cast<const u32x4*>(static_cast<const char*>(base) + off);
hi.u32 = *reinterpret_cast<const u32x4*>(static_cast<const char*>(base) + off + 16);
#pragma unroll
for (int i = 0; i < 16; ++i) { out[i] = lo.x2[i]; }
#pragma unroll
for (int i = 0; i < 16; ++i) { out[i+16] = hi.x2[i]; }
return out;
}
// Same for LDS: read two adjacent 16-byte chunks from A_s.
__device__ inline fp4x64_t expand_fp4x64_from_lds(const fp4x2_t* p) {
fp4x64_t out{};
union { u32x4 u32; fp4x2_t x2[16]; } lo, hi;
lo.u32 = *reinterpret_cast<const u32x4*>(p);
hi.u32 = *reinterpret_cast<const u32x4*>(p + 16);
#pragma unroll
for (int i = 0; i < 16; ++i) { out[i] = lo.x2[i]; }
#pragma unroll
for (int i = 0; i < 16; ++i) { out[i+16] = hi.x2[i]; }
return out;
}
struct BPrefetchBuf {
// Two base offsets per fragment so we can call expand_fp4x64_full() later.
// We store offsets rather than the data itself to keep register pressure
// manageable; the actual load happens inside mfma_compute via the pointer.
// Actually: store both halves as u32x4 pairs — avoids re-issuing globals
// from inside the compute lambda and keeps the original reg-prefetch idea.
u32x4 pkt_lo[FRAG_N_PER_WARP][FRAG_K]; // nibbles 0..31
u32x4 pkt_hi[FRAG_N_PER_WARP][FRAG_K]; // nibbles 32..63
uint8_t scale [FRAG_N_PER_WARP][FRAG_K];
};
struct AScaleBuf {
uint8_t scale[FRAG_M_PER_WARP][FRAG_K];
};
__device__ inline int e8m0_shuf_phys_offset(int r, int c, int padded_C) {
const int r_tile = r / 32, r_in = r % 32;
const int r_in_0 = r_in / 16, r_in_1 = r_in % 16;
const int c_tile = c / 8, c_in = c % 8;
const int c_in_0 = c_in / 4, 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}; }
// Returns byte offset of the FIRST 16-byte packet for (b_row, g_global).
// g_global is a "K/32 group" index (one group = 32 nibbles = 16 fp4x2_t).
// The second packet is always at offset + 16 bytes (adjacent in memory).
__device__ inline size_t bshuf_packet_offset_bytes(int b_row, int g_global, int K_logical) {
int n_tile = b_row / 16, n_in = b_row % 16;
int k_tile = g_global / 2, 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;
}
// ── Split-K: strided-slice epilogue (no atomicAdd) ──────────────────────────
// Each split z writes its partial FP32 result to slice z of the workspace
// (layout: [k_split, M, N]). A subsequent cheap reduction kernel sums them.
// This eliminates atomic contention entirely.
__global__ void reduce_splits_kernel(
const float* __restrict__ ws, bfp16* __restrict__ C,
int M, int N, int k_split)
{
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= M * N) return;
float acc = 0.f;
for (int s = 0; s < k_split; ++s)
acc += ws[(size_t)s * M * N + idx];
C[idx] = fast_f32tob16(acc);
}
__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, // [k_split, M, N] or nullptr
int M, int N, int K, int k_split)
{
int chunk_size = BK;
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 = min(my_chunk_start + chunks_per_split, total_chunks);
int k_start = my_chunk_start * chunk_size;
int k_end = min(my_chunk_end * chunk_size, K);
if (k_start >= k_end) return;
const int K32_global = K / 32;
const int PAD_SCALE_C = 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)));
// ── Async A tile → LDS (unchanged from original) ─────────────────────────
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, 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) {
#pragma unroll
for (int i = 0; i < FRAG_M_PER_WARP; ++i)
#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;
const int a_row = cur_m + warp_m*(BM/WARP_M) + i*MFMA_M + row_in_tile;
out.scale[i][kk] =
(a_row < M && g_global < K32_end)
? A_scale[e8m0_shuf_phys_offset(a_row, g_global, PAD_SCALE_C)]
: 0u;
}
};
// ── FIX: B prefetch loads BOTH 16-byte packets ────────────────────────────
// g_global is the K/32-group index. The shuffled B layout places consecutive
// groups contiguously within the same n_in row and k_chunk (verified by
// bshuf_packet_offset_bytes). Packet 0 is at offset off; packet 1 is at
// off + 16. Both must be loaded to fill fp4x64_t completely.
auto prefetch_B_regs = [&](int k0, BPrefetchBuf& 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 = cur_n + warp_n*(BN/WARP_N) + j*MFMA_N + row_in_tile;
if (b_row < N && g_global < K32_end) {
const char* base = reinterpret_cast<const char*>(B_shuf);
// g_global already in K/32 units — pass directly (unchanged).
// The second packet is +16 bytes within the same row/chunk.
const size_t off = bshuf_packet_offset_bytes(b_row, g_global, K);
out.pkt_lo[j][kk] = *reinterpret_cast<const u32x4*>(base + off);
out.pkt_hi[j][kk] = *reinterpret_cast<const u32x4*>(base + off + 16);
out.scale [j][kk] =
B_scale[e8m0_shuf_phys_offset(b_row, g_global, PAD_SCALE_C)];
} else {
out.pkt_lo[j][kk] = zero_u32x4();
out.pkt_hi[j][kk] = zero_u32x4();
out.scale [j][kk] = 0u;
}
}
}
};
// ── Compute: unchanged scale/LDS indexing; FIX fragment assembly ──────────
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) {
// Assemble full 128-nibble fragment from the two stored halves.
fp4x64_t out{};
union { u32x4 u; fp4x2_t x[16]; } lo, hi;
lo.u = bbuf.pkt_lo[j][kk];
hi.u = bbuf.pkt_hi[j][kk];
#pragma unroll
for (int i = 0; i < 16; ++i) { out[i] = lo.x[i]; }
#pragma unroll
for (int i = 0; i < 16; ++i) { out[i+16] = hi.x[i]; }
b_frag[j] = out;
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;
// LDS offset: g_local groups × 16 fp4x2_t per group,
// then load 32 fp4x2_t (two packets) — same base index as
// the original, but now we read 32 elements instead of 16.
const int a_lds_idx = g_local * 16; // in fp4x2_t units
const fp4x64_t a_frag = expand_fp4x64_from_lds(
&A_s[buf][a_row_local][a_lds_idx]);
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;
// ── Pipeline: issue A LDS async first, overlap with B reg prefetch ────────
int k0 = k_start;
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();
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;
}
mfma_compute(k0, b_cur, as_cur);
// ── Epilogue: strided-slice write (no atomicAdd) or direct bfp16 write ────
const size_t split_base = (size_t)blockIdx.z * M * N;
#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) continue;
#pragma unroll
for (int t = 0; t < 4; ++t) {
const int out_row = out_row_base + t;
if (out_row >= M) continue;
if (workspace) {
workspace[split_base + out_row * N + out_col] = c_reg[i][j][t];
} else {
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_ws = 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_ws, M, N, K, k_split);
if (k_split > 1) {
int total = M * N;
hipLaunchKernelGGL(reduce_splits_kernel,
dim3(ceil_div(total, 256)), dim3(256), 0, 0,
d_ws, d_C, M, N, k_split);
}
}
PYBIND11_MODULE(mxfp4, m) { m.def("run", &run, "HIP MXFP4 GEMM"); }
"""
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 dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
def _quant_mxfp4(x):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
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()
m, k = A.shape
n, _ = B.shape
A_q, A_scale_sh = _quant_mxfp4(A)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
target_blocks = 256
b = ((n + BN - 1) // BN) * ((m + BM - 1) // BM)
k_split = 1
if b < target_blocks:
k_split = min(16, (target_blocks + b - 1) // b)
k_chunks = k // BK
if k_chunks > 0:
k_split = min(k_split, k_chunks)
else:
k_split = 1
workspace = None
workspace_ptr = 0
if k_split > 1:
workspace = torch.zeros((k_split, m, n), dtype=torch.float32, device=A.device)
workspace_ptr = workspace.data_ptr()
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 C
BM, BN, BK = 32, 128, 512scrolls · 457 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