Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
25.8µs
#1107 of 1143
2026-04-02

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.

fp4PYBIND11_MODULE(mxfp4, m) { m.def("run", &run, "HIP MXFP4 GEMM"); }
shared-memory__shared__ fp4x2_t A_s[2][BM][BK / 2];
tile-k = 512constexpr int BK = 512;
tile-m = 32constexpr int BM = 32;
tile-n = 128constexpr 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, 512
scrolls · 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