Skip to content
KernelIndex
Search⌘K

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

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.

fp4FP4 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 = 512constexpr int BK = 512;
tile-m = 32constexpr int BM = 32;
tile-n = 128constexpr 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 C
scrolls · 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