Skip to content
KernelIndex
Search⌘K

submission 628737

hongquant.17 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 365 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-628737?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
18.5µs
#697 of 1143
2026-03-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bb2993c634cdad2260a0d7ddf2367c4144d718a5f9d45e0b21a4a7497d161dae
license declaredunknown
license concludedunknown
authorshongquant.17
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4constexpr int kBlockElems = 32; // MXFP4 block size

Kernel source

submission.py365 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

import os
from functools import lru_cache

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

if "PYTORCH_ROCM_ARCH" not in os.environ:
    os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"

quant_src = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <stdint.h>
#include <vector>
#include <math.h>

constexpr int kWaveSize      = 64;
constexpr int kGroupLanes    = 16;   // one 16-lane subgroup per 1x32 block
constexpr int kBlocksPerWave = 4;    // 64 lanes / 16
constexpr int kBlockElems    = 32;   // MXFP4 block size
constexpr int kBytesPerBlock = 16;   // 32 fp4 elems / 2 per byte
constexpr int kScalePadRows  = 256;
constexpr int kScalePadCols  = 8;

using native_bf16   = __bf16;
using native_bf16x2 = __attribute__((ext_vector_type(2))) native_bf16;

template <typename T>
__host__ __device__ __forceinline__ T round_up(T x, T a) {
    return ((x + a - 1) / a) * a;
}

template <typename T>
__host__ __device__ __forceinline__ T ceil_div(T x, T y) {
    return (x + y - 1) / y;
}

__host__ __device__ __forceinline__ uint32_t bitcast_f32_to_u32(float x) {
    union { float f; uint32_t u; } v;
    v.f = x;
    return v.u;
}

__host__ __device__ __forceinline__ float bitcast_u32_to_f32(uint32_t x) {
    union { uint32_t u; float f; } v;
    v.u = x;
    return v.f;
}

__device__ __forceinline__ native_bf16 bitcast_u16_to_native_bf16(uint16_t x) {
    union { uint16_t u; native_bf16 b; } v;
    v.u = x;
    return v.b;
}

__host__ __device__ __forceinline__ float bf16_bits_to_f32(uint16_t x) {
    return bitcast_u32_to_f32(uint32_t(x) << 16);
}

// Same E8M0 reconstruction AITER uses.
__host__ __device__ __forceinline__ float e8m0_to_f32_quant(uint8_t e8m0) {
    uint32_t bits;
    if (e8m0 == 0x00u) {
        bits = 0x00400000u;
    } else if (e8m0 == 0xFFu) {
        bits = 0x7F800001u;
    } else {
        bits = uint32_t(e8m0) << 23;
    }
    return bitcast_u32_to_f32(bits);
}

// Match aiter.utility.fp4_utils.dynamic_mxfp4_quant scale selection:
//   amax_bits = (amax_bits + 0x200000) & 0xFF800000
//   scale_e8m0 = biased_exp(amax_rounded) - 2
__host__ __device__ __forceinline__ uint8_t choose_scale_e8m0_aiter(float amax) {
    if (amax == 0.0f) return 0u;

    uint32_t u = bitcast_f32_to_u32(amax);
    u = (u + 0x00200000u) & 0xFF800000u;

    int scale_biased = int((u >> 23) & 0xFFu) - 2;
    if (scale_biased < 0) scale_biased = 0;
    if (scale_biased > 0xFF) scale_biased = 0xFF;
    return static_cast<uint8_t>(scale_biased);
}

// Exact mapping for:
//   scale = scale.view(sm // 32, 2, 16, sn // 8, 2, 4)
//   scale = scale.permute(0, 3, 5, 2, 4, 1).contiguous()
//   scale = scale.view(sm, sn)
__host__ __device__ __forceinline__ uint64_t e8m0_shuffle_flat_index(
    int row, int col, int sn_padded)
{
    const int a = row / 32;
    const int b = (row % 32) / 16;
    const int c = row % 16;

    const int d = col / 8;
    const int e = (col % 8) / 4;
    const int f = col % 4;

    const int d_extent = sn_padded / 8;

    return ((((uint64_t(a) * uint64_t(d_extent) + uint64_t(d)) * 4ull
             + uint64_t(f)) * 16ull + uint64_t(c)) * 2ull
             + uint64_t(e)) * 2ull + uint64_t(b);
}

__device__ __forceinline__ float subgroup16_max(float v, int lane) {
    #pragma unroll
    for (int mask = kGroupLanes >> 1; mask > 0; mask >>= 1) {
        const int peer = (lane & ~(kGroupLanes - 1)) | ((lane & (kGroupLanes - 1)) ^ mask);
        const float other = __shfl(v, peer, kWaveSize);
        v = fmaxf(v, other);
    }
    return v;
}

__device__ __forceinline__ uint8_t pack_two_bf16_to_fp4x2(
    uint16_t x0_bits, uint16_t x1_bits, float scale_f32)
{
    native_bf16x2 src;
    src[0] = bitcast_u16_to_native_bf16(x0_bits);
    src[1] = bitcast_u16_to_native_bf16(x1_bits);

    // opsel=0 -> write packed fp4x2 into byte 0 of returned u32
    const unsigned packed =
        __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0u, src, scale_f32, 0);

    return static_cast<uint8_t>(packed & 0xFFu);
}

template <bool kShuffleScale>
__global__ __launch_bounds__(256)
void quant_bf16_to_mxfp4_kernel(
    const uint16_t* __restrict__ x_bf16,
    uint8_t* __restrict__ q,   // [M, K/2]
    uint8_t* __restrict__ s,   // [M, K/32] or shuffled [Mp, Sp]
    int M,
    int K,
    int64_t stride_xm,
    int64_t stride_xn,
    int64_t stride_qm,
    int64_t stride_qn,
    int64_t stride_sm,
    int64_t stride_sn)
{
    const int row = int(blockIdx.y);
    if (row >= M) return;

    const int tid             = int(threadIdx.x);
    const int lane            = tid & (kWaveSize - 1);
    const int wave_in_block   = tid >> 6;
    const int waves_per_block = int(blockDim.x) >> 6;

    const int subgroup = lane / kGroupLanes;
    const int lane16   = lane & (kGroupLanes - 1);

    const int blocks_per_row = K / kBlockElems;
    const int sn_padded = kShuffleScale ? round_up(blocks_per_row, kScalePadCols) : 0;
    const int block_id =
        ((int(blockIdx.x) * waves_per_block + wave_in_block) * kBlocksPerWave) + subgroup;

    if (block_id >= blocks_per_row) return;

    const int x_col0 = block_id * kBlockElems + lane16 * 2 + 0;
    const int x_col1 = block_id * kBlockElems + lane16 * 2 + 1;
    const int64_t x_row_offset = int64_t(row) * stride_xm;
    const int64_t q_row_offset = int64_t(row) * stride_qm;

    const uint16_t x0_bits =
        x_bf16[x_row_offset + int64_t(x_col0) * stride_xn];
    const uint16_t x1_bits =
        x_bf16[x_row_offset + int64_t(x_col1) * stride_xn];

    const float x0 = bf16_bits_to_f32(x0_bits);
    const float x1 = bf16_bits_to_f32(x1_bits);

    const float local_amax = fmaxf(fabsf(x0), fabsf(x1));
    const float block_amax = subgroup16_max(local_amax, lane);

    const uint8_t scale_e8m0 = choose_scale_e8m0_aiter(block_amax);

    uint8_t q_byte = 0;
    if (block_amax != 0.0f) {
        const float scale_f32 = e8m0_to_f32_quant(scale_e8m0);
        q_byte = pack_two_bf16_to_fp4x2(x0_bits, x1_bits, scale_f32);
    }

    const int q_col = block_id * kBytesPerBlock + lane16;
    q[q_row_offset + int64_t(q_col) * stride_qn] = q_byte;

    if (lane16 == 0) {
        if constexpr (!kShuffleScale) {
            s[int64_t(row) * stride_sm + int64_t(block_id) * stride_sn] = scale_e8m0;
        } else {
            const uint64_t flat = e8m0_shuffle_flat_index(row, block_id, sn_padded);
            const int srow = int(flat / uint64_t(sn_padded));
            const int scol = int(flat % uint64_t(sn_padded));
            s[int64_t(srow) * stride_sm + int64_t(scol) * stride_sn] = scale_e8m0;
        }
    }
}

template <bool kShuffleScale>
void launch_quantize(
    const uint16_t* x_ptr,
    uint8_t* q_ptr,
    uint8_t* s_ptr,
    int M,
    int K,
    int64_t stride_xm,
    int64_t stride_xn,
    int64_t stride_qm,
    int64_t stride_qn,
    int64_t stride_sm,
    int64_t stride_sn)
{
    constexpr int kThreads = 256;
    constexpr int kBlocksPerCta = (kThreads / kWaveSize) * kBlocksPerWave;

    dim3 threads(kThreads);
    dim3 blocks(ceil_div(K / kBlockElems, kBlocksPerCta), M);

    hipLaunchKernelGGL(
        HIP_KERNEL_NAME(quant_bf16_to_mxfp4_kernel<kShuffleScale>),
        blocks,
        threads,
        0,
        0,
        x_ptr,
        q_ptr,
        s_ptr,
        M,
        K,
        stride_xm,
        stride_xn,
        stride_qm,
        stride_qn,
        stride_sm,
        stride_sn);
}

std::vector<torch::Tensor> quantize(torch::Tensor x, bool shuffle) {
    TORCH_CHECK(x.is_cuda(), "x must be a CUDA/HIP tensor");
    TORCH_CHECK(x.scalar_type() == at::kBFloat16, "x must be bf16");
    TORCH_CHECK(x.dim() == 2, "x must be 2D [M, K]");

    x = x.contiguous();

    const auto M = x.size(0);
    const auto K = x.size(1);

    TORCH_CHECK(K % 32 == 0, "K must be a multiple of 32");
    TORCH_CHECK(M <= INT_MAX, "M exceeds kernel int indexing range");
    TORCH_CHECK(K <= INT_MAX, "K exceeds kernel int indexing range");

    const c10::cuda::CUDAGuard device_guard(x.device());

    auto u8_opts = x.options().dtype(at::kByte);
    auto q = torch::empty({M, K / 2}, u8_opts);

    torch::Tensor s;
    if (shuffle) {
        const auto Mp = round_up<int64_t>(M, kScalePadRows);
        const auto Sp = round_up<int64_t>(K / 32, kScalePadCols);
        s = torch::full({Mp, Sp}, 127, u8_opts);   // pad value must be 127
    } else {
        s = torch::empty({M, K / 32}, u8_opts);
    }

    const auto* x_ptr = reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>());
    auto* q_ptr = q.data_ptr<uint8_t>();
    auto* s_ptr = s.data_ptr<uint8_t>();
    const int M_i = static_cast<int>(M);
    const int K_i = static_cast<int>(K);

    if (shuffle) {
        launch_quantize<true>(
            x_ptr,
            q_ptr,
            s_ptr,
            M_i,
            K_i,
            x.stride(0),
            x.stride(1),
            q.stride(0),
            q.stride(1),
            s.stride(0),
            s.stride(1));
    } else {
        launch_quantize<false>(
            x_ptr,
            q_ptr,
            s_ptr,
            M_i,
            K_i,
            x.stride(0),
            x.stride(1),
            q.stride(0),
            q.stride(1),
            s.stride(0),
            s.stride(1));
    }

    auto err = hipGetLastError();
    TORCH_CHECK(
        err == hipSuccess,
        "quant_bf16_to_mxfp4_kernel launch failed: ",
        hipGetErrorString(err));

    return {q, s};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("quantize", &quantize, "HIP MXFP4 quant kernel");
}
"""

@lru_cache(maxsize=1)
def _get_quant_module():
    return load_inline(
        name="mxfp4_quant_ext_v2",   # bump this name if you want to force a rebuild
        cpp_sources="",
        cuda_sources=quant_src,
        with_cuda=True,
        verbose=False,
        extra_cuda_cflags=["-std=c++20", "-O3"],
        no_implicit_headers=True,
    )

def _quant_mxfp4_custom(x, shuffle=True):
    from aiter import dtypes

    q_u8, s_u8 = _get_quant_module().quantize(x, shuffle)
    return q_u8.view(dtypes.fp4x2), s_u8.view(dtypes.fp8_e8m0)


def custom_kernel(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes

    A, _, _, B_shuffle, B_scale_sh = data
    A = A.contiguous()

    A_q, A_scale_sh = _quant_mxfp4_custom(A, shuffle=True)

    out_gemm = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out_gemm
scrolls · 365 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