Skip to content
KernelIndex
Search⌘K

submission 596924

Purple rain · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_amd_mxfp4_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-596924?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
24.2µs
#991 of 1143
2026-03-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4e49dbfab68f4675da0fa9e6cde3642b97848823ae46d60f7545a9d24985e8a5
license declaredunknown
license concludedunknown
authorsPurple rain
imported2026-08-26

Techniques

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

fp4TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");
shared-memory__shared__ __align__(16) uint8_t smem_A[GEMM_BLOCK_M * FP4_BYTES_PER_KBLOCK];
split-k_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"
tile-k = 32constexpr int BLOCK_K = 32;
tile-m = 64constexpr int GEMM_BLOCK_M = 64;
tile-n = 64constexpr int GEMM_BLOCK_N = 64;
vector-width = uint4uint4 v4 = *reinterpret_cast<const uint4*>(row_ptr + vec_idx * 16);

Kernel source

submission_amd_mxfp4_mm.py1084 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

import os
import threading
from typing import Dict, Tuple

import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline


_QUANT_BACKEND_ENV = "MXFP4_MM_QUANT_BACKEND"  # auto | triton | hip
_EXEC_BACKEND_ENV = "MXFP4_MM_EXEC_BACKEND"  # auto | aiter | hip
_B_LAYOUT_ENV = "MXFP4_MM_B_LAYOUT"  # raw | shuffle | auto
_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"
_DEFAULT_QUANT_BACKEND = "hip"
_DEFAULT_EXEC_BACKEND = "auto"
_DEFAULT_B_LAYOUT = "shuffle"


_HIP_LOCK = threading.Lock()
_HIP_MODULE = None
_HIP_BUILD_ERROR = None

_B_SCALE_RAW_LOCK = threading.Lock()
_B_SCALE_RAW_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_B_SCALE_RAW_CACHE_MAX = 8

_WORKSPACE_LOCK = threading.Lock()
_WORKSPACE_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_WORKSPACE_CACHE_MAX = 16

_B_SHUFFLE_INNER_LOCK = threading.Lock()
_B_SHUFFLE_INNER_CACHE: Dict[Tuple[int, ...], int] = {}


_STATIC_SPLITK_LOG2: Dict[Tuple[int, int, int], int] = {
    (4, 2880, 512): 2,
    (16, 2112, 7168): 3,
    (32, 4096, 512): 2,
    (32, 2880, 512): 2,
    (64, 7168, 2048): 1,
    (256, 3072, 1536): 0,
}


CPP_WRAPPER = r"""
#include <cstdint>
#include <vector>
std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x);
torch::Tensor hip_gemm_mxfp4(
    torch::Tensor a_fp4_u8,
    torch::Tensor b_u8,
    torch::Tensor a_scale_u8,
    torch::Tensor b_scale_u8,
    int64_t layout_mode,
    int64_t log2_k_split,
    torch::Tensor workspace,
    int64_t workspace_stride,
    int64_t b_shuffle_inner_mode);
"""


HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>

#include <algorithm>
#include <cmath>
#include <cstdint>
#include <vector>

namespace {

constexpr int BLOCK_K = 32;
constexpr int PAD_M_KERNEL = 64;
constexpr int PAD_M_SCALE = 256;
constexpr int PAD_SCALE_N = 8;

constexpr int FP4_BYTES_PER_KBLOCK = BLOCK_K / 2;

constexpr int GEMM_BLOCK_M = 64;
constexpr int GEMM_BLOCK_N = 64;
constexpr int GEMM_THREADS = 256;
constexpr int WAVE_SIZE = 64;
constexpr int WAVE_TILE_M = 32;
constexpr int WAVE_TILE_N = 32;
constexpr int MFMA_TILE_M = 16;
constexpr int MFMA_TILE_N = 16;

__device__ __constant__ uint16_t FP4_TO_BF16_LUT[16] = {
    0x0000, 0x3f00, 0x3f80, 0x3fc0, 0x4000, 0x4040, 0x4080, 0x40c0,
    0x8000, 0xbf00, 0xbf80, 0xbfc0, 0xc000, 0xc040, 0xc080, 0xc0c0
};

__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
    return __uint_as_float(static_cast<uint32_t>(x) << 16);
}

__device__ __forceinline__ uint8_t quantize_e2m1(float x_scaled) {
    constexpr uint32_t T0 = 0x3E800000u;  // 0.25
    constexpr uint32_t T1 = 0x3F400000u;  // 0.75
    constexpr uint32_t T2 = 0x3FA00000u;  // 1.25
    constexpr uint32_t T3 = 0x3FE00000u;  // 1.75
    constexpr uint32_t T4 = 0x40200000u;  // 2.5
    constexpr uint32_t T5 = 0x40600000u;  // 3.5
    constexpr uint32_t T6 = 0x40A00000u;  // 5.0

    uint32_t bits = __float_as_uint(x_scaled);
    uint32_t sign = (bits >> 31) & 0x1u;
    uint32_t abs_bits = bits & 0x7FFFFFFFu;

    uint32_t mag = 0;
    mag += static_cast<uint32_t>(abs_bits >= T0);
    mag += static_cast<uint32_t>(abs_bits >= T1);
    mag += static_cast<uint32_t>(abs_bits >= T2);
    mag += static_cast<uint32_t>(abs_bits >= T3);
    mag += static_cast<uint32_t>(abs_bits >= T4);
    mag += static_cast<uint32_t>(abs_bits >= T5);
    mag += static_cast<uint32_t>(abs_bits >= T6);

    return static_cast<uint8_t>(mag | (sign << 3));
}

__device__ __forceinline__ int64_t shuffled_scale_offset(
    int64_t row,
    int64_t col,
    int64_t scale_n_pad) {
    int64_t bs_offs_0 = row / 32;
    int64_t bs_offs_1 = row % 32;
    int64_t bs_offs_2 = bs_offs_1 % 16;
    bs_offs_1 = bs_offs_1 / 16;

    int64_t bs_offs_3 = col / 8;
    int64_t bs_offs_4 = col % 8;
    int64_t bs_offs_5 = bs_offs_4 % 4;
    bs_offs_4 = bs_offs_4 / 4;

    return bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4 + bs_offs_5 * 64 +
           bs_offs_3 * 256 + bs_offs_0 * 32 * scale_n_pad;
}

__device__ __forceinline__ int64_t get_b_scale_shuffled_offset(
    int64_t n,
    int64_t k_blk,
    int64_t scale_k_pad) {
    int64_t n_outer = n / 32;
    int64_t n_inner = n % 32;
    int64_t k_outer = k_blk / 8;
    int64_t k_inner = k_blk % 8;

    int64_t n_16 = n_inner % 16;
    int64_t n_2 = n_inner / 16;
    int64_t k_4 = k_inner % 4;
    int64_t k_2 = k_inner / 4;

    return n_2 + (k_2 * 2) + (n_16 * 4) + (k_4 * 64) + (k_outer * 256) +
           (n_outer * 32 * scale_k_pad);
}

__device__ __forceinline__ int64_t get_b_fp4_shuffled_offset(
    int64_t n,
    int64_t k_fp4,
    int64_t k_fp4_pad,
    int64_t inner_mode) {
    int64_t n_blk = n / 16;
    int64_t k_blk = k_fp4 / 16;
    int64_t n_in = n % 16;
    int64_t k_in = k_fp4 % 16;

    int64_t blk_stride = k_fp4_pad / 16;
    int64_t blk_offset = (n_blk * blk_stride + k_blk) * 256;
    int64_t inner_offset = (inner_mode == 0) ? (k_in * 16 + n_in) : (n_in * 16 + k_in);
    return blk_offset + inner_offset;
}

__device__ __forceinline__ float e8m0_to_f32_fast(uint8_t e) {
    if (e == 0) {
        return __uint_as_float(0x00400000u);
    }
    if (e == 0xFF) {
        return __uint_as_float(0x7F800001u);
    }
    return __uint_as_float(static_cast<uint32_t>(e) << 23);
}

__device__ __forceinline__ uint16_t fp4_to_bf16_bits(uint8_t v) {
    return FP4_TO_BF16_LUT[v & 0x0Fu];
}

using floatx4 = __attribute__((__vector_size__(4 * sizeof(float)))) float;
using bit16x4 = __attribute__((__vector_size__(4 * sizeof(uint16_t)))) uint16_t;
using bit16x8 = __attribute__((__vector_size__(8 * sizeof(uint16_t)))) uint16_t;
struct B16x8 {
    bit16x4 xy[2];
};

__device__ __forceinline__ B16x8 unpack_fp4x8_to_b16x8(uint32_t pack) {
    B16x8 reg;
    const uint8_t* bytes = reinterpret_cast<const uint8_t*>(&pack);
    #pragma unroll
    for (int i = 0; i < 8; ++i) {
        uint8_t packed = bytes[i >> 1];
        uint8_t nib = (i & 1) ? static_cast<uint8_t>(packed >> 4)
                              : static_cast<uint8_t>(packed & 0x0F);
        uint16_t bits = fp4_to_bf16_bits(nib);
        if (i < 4) {
            reg.xy[0][i] = bits;
        } else {
            reg.xy[1][i - 4] = bits;
        }
    }
    return reg;
}

__device__ __forceinline__ uint32_t load_packed_fp4_word(
    const uint8_t* tile,
    int tile_idx,
    int word_idx) {
    const uint32_t* words = reinterpret_cast<const uint32_t*>(
        tile + tile_idx * FP4_BYTES_PER_KBLOCK);
    return words[word_idx];
}

__device__ __forceinline__ void accum_scaled_output_fragment(
    floatx4& dst,
    const floatx4& src,
    const floatx4& a_scale,
    float b_scale) {
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        dst[i] += src[i] * a_scale[i] * b_scale;
    }
}

__device__ __forceinline__ floatx4 decode_scale4(const uint8_t* scale_ptr) {
    floatx4 out;
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        out[i] = e8m0_to_f32_fast(scale_ptr[i]);
    }
    return out;
}

__device__ __forceinline__ void store_accum_tile(
    float* workspace,
    int64_t workspace_stride,
    int64_t split_idx,
    int64_t m,
    int64_t n,
    int64_t row_base,
    int64_t col,
    int row_group,
    const floatx4& acc) {
    if (col >= n) {
        return;
    }
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        int64_t row = row_base + row_group * 4 + i;
        if (row < m) {
            int64_t out_off = split_idx * workspace_stride + row * n + col;
            workspace[out_off] = acc[i];
        }
    }
}

__device__ __forceinline__ floatx4 gcn_mfma16x16x32_bf16(
    const B16x8& a,
    const B16x8& b,
    const floatx4& c) {
#if defined(__gfx950__)
    bit16x8 ta = __builtin_shufflevector(a.xy[0], a.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
    bit16x8 tb = __builtin_shufflevector(b.xy[0], b.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
    return __builtin_amdgcn_mfma_f32_16x16x32_bf16(ta, tb, c, 0, 0, 0);
#else
    return c;
#endif
}

__global__ void quant_mxfp4_kernel(
    const __hip_bfloat16* x,
    uint8_t* out_fp4,
    uint8_t* out_scale,
    int64_t m,
    int64_t m_pad,
    int64_t k,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad) {
    int64_t linear = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    int64_t total = m_pad * k_blocks_pad;
    if (linear >= total) {
        return;
    }

    int64_t row = linear / k_blocks_pad;
    int64_t kb = linear % k_blocks_pad;
    int64_t scale_off = shuffled_scale_offset(row, kb, k_blocks_pad);

    if (row >= m || kb >= k_blocks_valid) {
        out_scale[scale_off] = 127;

        if (row < m_pad && kb < k_blocks_valid) {
            int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);
            #pragma unroll
            for (int i = 0; i < BLOCK_K / 2; ++i) {
                out_fp4[out_base + i] = 0;
            }
        }
        return;
    }

    int64_t in_base = row * k + kb * BLOCK_K;

    float vals[BLOCK_K];
    float amax = 0.0f;

    const uint8_t* x_bytes = reinterpret_cast<const uint8_t*>(x);
    const uint8_t* row_ptr = x_bytes + in_base * static_cast<int64_t>(sizeof(__hip_bfloat16));

    #pragma unroll
    for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
        uint4 v4 = *reinterpret_cast<const uint4*>(row_ptr + vec_idx * 16);
        uint32_t w0 = v4.x;
        uint32_t w1 = v4.y;
        uint32_t w2 = v4.z;
        uint32_t w3 = v4.w;

        uint16_t b0 = static_cast<uint16_t>(w0 & 0xFFFFu);
        uint16_t b1 = static_cast<uint16_t>((w0 >> 16) & 0xFFFFu);
        uint16_t b2 = static_cast<uint16_t>(w1 & 0xFFFFu);
        uint16_t b3 = static_cast<uint16_t>((w1 >> 16) & 0xFFFFu);
        uint16_t b4 = static_cast<uint16_t>(w2 & 0xFFFFu);
        uint16_t b5 = static_cast<uint16_t>((w2 >> 16) & 0xFFFFu);
        uint16_t b6 = static_cast<uint16_t>(w3 & 0xFFFFu);
        uint16_t b7 = static_cast<uint16_t>((w3 >> 16) & 0xFFFFu);

        int base = vec_idx * 8;
        float f0 = bf16_to_f32(b0);
        float f1 = bf16_to_f32(b1);
        float f2 = bf16_to_f32(b2);
        float f3 = bf16_to_f32(b3);
        float f4 = bf16_to_f32(b4);
        float f5 = bf16_to_f32(b5);
        float f6 = bf16_to_f32(b6);
        float f7 = bf16_to_f32(b7);

        vals[base + 0] = f0;
        vals[base + 1] = f1;
        vals[base + 2] = f2;
        vals[base + 3] = f3;
        vals[base + 4] = f4;
        vals[base + 5] = f5;
        vals[base + 6] = f6;
        vals[base + 7] = f7;

        amax = fmaxf(amax, fabsf(f0));
        amax = fmaxf(amax, fabsf(f1));
        amax = fmaxf(amax, fabsf(f2));
        amax = fmaxf(amax, fabsf(f3));
        amax = fmaxf(amax, fabsf(f4));
        amax = fmaxf(amax, fabsf(f5));
        amax = fmaxf(amax, fabsf(f6));
        amax = fmaxf(amax, fabsf(f7));
    }

    int exp_unbiased = 0;
    float scale = 1.0f;
    if (amax > 0.0f) {
        float target = amax * (1.0f / 6.0f);
        exp_unbiased = static_cast<int>(ceilf(log2f(target)));
        scale = exp2f(static_cast<float>(exp_unbiased));
    }

    int exp_biased = exp_unbiased + 127;
    exp_biased = exp_biased < 0 ? 0 : (exp_biased > 255 ? 255 : exp_biased);
    out_scale[scale_off] = static_cast<uint8_t>(exp_biased);

    float inv_scale = 1.0f / scale;
    int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);

    #pragma unroll
    for (int i = 0; i < BLOCK_K / 2; ++i) {
        uint8_t lo = quantize_e2m1(vals[2 * i] * inv_scale);
        uint8_t hi = quantize_e2m1(vals[2 * i + 1] * inv_scale);
        out_fp4[out_base + i] = static_cast<uint8_t>((hi << 4) | lo);
    }
}

__global__ void gemm_mxfp4_splitk_kernel(
    const uint8_t* a_fp4,
    const uint8_t* b_u8,
    const uint8_t* a_scale_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t layout_mode,
    int64_t b_shuffle_inner_mode,
    int64_t split_k) {
    int tid = static_cast<int>(threadIdx.x);
    int wave_id = tid / WAVE_SIZE;
    int lane = tid & (WAVE_SIZE - 1);
    int lane16 = lane & 15;
    int row_group = lane >> 4;  // 0..3
    int wave_row = wave_id >> 1;
    int wave_col = wave_id & 1;
    int quad_row_base = wave_row * WAVE_TILE_M;
    int quad_col_base = wave_col * WAVE_TILE_N;
    int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
    int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
    int local_col0 = quad_col_base + lane16;
    int local_col1 = local_col0 + MFMA_TILE_N;
    int64_t row_base0 = tile_m + quad_row_base;
    int64_t row_base1 = row_base0 + MFMA_TILE_M;
    int64_t col0 = tile_n + local_col0;
    int64_t col1 = tile_n + local_col1;

    floatx4 c_acc00 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc01 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc10 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc11 = {0.0f, 0.0f, 0.0f, 0.0f};

    __shared__ __align__(16) uint8_t smem_A[GEMM_BLOCK_M * FP4_BYTES_PER_KBLOCK];
    __shared__ __align__(16) uint8_t smem_B[GEMM_BLOCK_N * FP4_BYTES_PER_KBLOCK];
    __shared__ uint8_t smem_A_scale[GEMM_BLOCK_M];
    __shared__ uint8_t smem_B_scale[GEMM_BLOCK_N];

    auto* smem_a_vec = reinterpret_cast<uint4*>(smem_A);
    auto* smem_b_vec = reinterpret_cast<uint4*>(smem_B);

    int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
    int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
    int64_t kb_end = kb_start + kb_per_split;
    if (kb_end > k_blocks_valid) {
        kb_end = k_blocks_valid;
    }

    for (int64_t kb = kb_start; kb < kb_end; ++kb) {
        if (tid < GEMM_BLOCK_M) {
            int local_row = tid;
            int64_t a_row = tile_m + local_row;
            uint4 a_vec = make_uint4(0u, 0u, 0u, 0u);
            uint8_t a_scale = 127;
            if (a_row < m) {
                a_vec = *reinterpret_cast<const uint4*>(
                    a_fp4 + a_row * k2 + kb * FP4_BYTES_PER_KBLOCK);
                a_scale = a_scale_u8[a_row * k_blocks_pad + kb];
            }
            smem_a_vec[local_row] = a_vec;
            smem_A_scale[local_row] = a_scale;
        }
        if (tid >= GEMM_BLOCK_M && tid < GEMM_BLOCK_M + GEMM_BLOCK_N) {
            int local_n = tid - GEMM_BLOCK_M;
            int64_t b_row = tile_n + local_n;
            uint4 b_vec = make_uint4(0u, 0u, 0u, 0u);
            uint8_t b_scale = 127;
            if (b_row < n) {
                if (layout_mode == 0) {
                    b_vec = *reinterpret_cast<const uint4*>(
                        b_u8 + b_row * k2 + kb * FP4_BYTES_PER_KBLOCK);
                    b_scale = b_scale_u8[b_row * k_blocks_pad + kb];
                } else {
                    if (b_shuffle_inner_mode == 1) {
                        int64_t boff = get_b_fp4_shuffled_offset(
                            b_row, kb * FP4_BYTES_PER_KBLOCK, k2, b_shuffle_inner_mode);
                        b_vec = *reinterpret_cast<const uint4*>(b_u8 + boff);
                    } else {
                        uint8_t* dst = reinterpret_cast<uint8_t*>(&b_vec);
                        #pragma unroll
                        for (int bi = 0; bi < FP4_BYTES_PER_KBLOCK; ++bi) {
                            int64_t k_fp4 = kb * FP4_BYTES_PER_KBLOCK + bi;
                            int64_t boff = get_b_fp4_shuffled_offset(
                                b_row, k_fp4, k2, b_shuffle_inner_mode);
                            dst[bi] = b_u8[boff];
                        }
                    }
                    int64_t soff = get_b_scale_shuffled_offset(b_row, kb, k_blocks_pad);
                    b_scale = b_scale_u8[soff];
                }
            }
            smem_b_vec[local_n] = b_vec;
            smem_B_scale[local_n] = b_scale;
        }
        __syncthreads();

        uint32_t a_pack0 = load_packed_fp4_word(
            smem_A, quad_row_base + lane16, row_group);
        uint32_t a_pack1 = load_packed_fp4_word(
            smem_A, quad_row_base + MFMA_TILE_M + lane16, row_group);
        uint32_t b_pack0 = load_packed_fp4_word(
            smem_B, quad_col_base + lane16, row_group);
        uint32_t b_pack1 = load_packed_fp4_word(
            smem_B, quad_col_base + MFMA_TILE_N + lane16, row_group);

        B16x8 a_reg0 = unpack_fp4x8_to_b16x8(a_pack0);
        B16x8 a_reg1 = unpack_fp4x8_to_b16x8(a_pack1);
        B16x8 b_reg0 = unpack_fp4x8_to_b16x8(b_pack0);
        B16x8 b_reg1 = unpack_fp4x8_to_b16x8(b_pack1);

        floatx4 t_acc00 = {0.0f, 0.0f, 0.0f, 0.0f};
        floatx4 t_acc01 = {0.0f, 0.0f, 0.0f, 0.0f};
        floatx4 t_acc10 = {0.0f, 0.0f, 0.0f, 0.0f};
        floatx4 t_acc11 = {0.0f, 0.0f, 0.0f, 0.0f};

        t_acc00 = gcn_mfma16x16x32_bf16(a_reg0, b_reg0, t_acc00);
        t_acc01 = gcn_mfma16x16x32_bf16(a_reg0, b_reg1, t_acc01);
        t_acc10 = gcn_mfma16x16x32_bf16(a_reg1, b_reg0, t_acc10);
        t_acc11 = gcn_mfma16x16x32_bf16(a_reg1, b_reg1, t_acc11);

        int row_scale_base0 = quad_row_base + row_group * 4;
        int row_scale_base1 = row_scale_base0 + MFMA_TILE_M;
        const uint8_t* a_scale_ptr0 = smem_A_scale + row_scale_base0;
        const uint8_t* a_scale_ptr1 = smem_A_scale + row_scale_base1;
        floatx4 a_scale0 = decode_scale4(a_scale_ptr0);
        floatx4 a_scale1 = decode_scale4(a_scale_ptr1);
        float b_scale0 = e8m0_to_f32_fast(smem_B_scale[local_col0]);
        float b_scale1 = e8m0_to_f32_fast(smem_B_scale[local_col1]);

        accum_scaled_output_fragment(c_acc00, t_acc00, a_scale0, b_scale0);
        accum_scaled_output_fragment(c_acc01, t_acc01, a_scale0, b_scale1);
        accum_scaled_output_fragment(c_acc10, t_acc10, a_scale1, b_scale0);
        accum_scaled_output_fragment(c_acc11, t_acc11, a_scale1, b_scale1);
        __syncthreads();
    }

    store_accum_tile(
        workspace,
        workspace_stride,
        static_cast<int64_t>(blockIdx.z),
        m,
        n,
        row_base0,
        col0,
        row_group,
        c_acc00);
    store_accum_tile(
        workspace,
        workspace_stride,
        static_cast<int64_t>(blockIdx.z),
        m,
        n,
        row_base0,
        col1,
        row_group,
        c_acc01);
    store_accum_tile(
        workspace,
        workspace_stride,
        static_cast<int64_t>(blockIdx.z),
        m,
        n,
        row_base1,
        col0,
        row_group,
        c_acc10);
    store_accum_tile(
        workspace,
        workspace_stride,
        static_cast<int64_t>(blockIdx.z),
        m,
        n,
        row_base1,
        col1,
        row_group,
        c_acc11);
}

__global__ void reduce_splitk_kernel(
    const float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t split_k) {
    int64_t row = static_cast<int64_t>(blockIdx.y) * blockDim.y + threadIdx.y;
    int64_t col = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (row >= m || col >= n) {
        return;
    }

    float sum = 0.0f;
    int64_t base = row * n + col;
    for (int64_t s = 0; s < split_k; ++s) {
        sum += workspace[s * workspace_stride + base];
    }
    out[base] = static_cast<__hip_bfloat16>(sum);
}

inline void check_hip_error(const char* where) {
    hipError_t err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}

}  // namespace

std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x) {
    TORCH_CHECK(x.is_cuda(), "x must be CUDA tensor");
    TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bfloat16");
    TORCH_CHECK(x.dim() == 2, "x must be 2D [M, K]");

    auto x_contig = x.contiguous();
    int64_t m = x_contig.size(0);
    int64_t k = x_contig.size(1);
    TORCH_CHECK(k % 64 == 0, "k must be divisible by 64");

    int64_t k_blocks_valid = k / BLOCK_K;
    int64_t k_blocks_pad = ((k_blocks_valid + (PAD_SCALE_N - 1)) / PAD_SCALE_N) * PAD_SCALE_N;
    int64_t m_pad_kernel = ((m + (PAD_M_KERNEL - 1)) / PAD_M_KERNEL) * PAD_M_KERNEL;
    int64_t m_pad_scale = ((m + (PAD_M_SCALE - 1)) / PAD_M_SCALE) * PAD_M_SCALE;

    auto u8_opts = x_contig.options().dtype(torch::kUInt8);
    auto out_fp4 = torch::empty({m_pad_kernel, k / 2}, u8_opts);
    auto out_scale = torch::full({m_pad_scale, k_blocks_pad}, 127, u8_opts);

    int64_t total = m_pad_kernel * k_blocks_pad;
    int threads = 256;
    int blocks = static_cast<int>((total + threads - 1) / threads);
    if (blocks > 0) {
        hipLaunchKernelGGL(
            quant_mxfp4_kernel,
            dim3(blocks),
            dim3(threads),
            0,
            0,
            reinterpret_cast<const __hip_bfloat16*>(x_contig.data_ptr()),
            reinterpret_cast<uint8_t*>(out_fp4.data_ptr()),
            reinterpret_cast<uint8_t*>(out_scale.data_ptr()),
            m,
            m_pad_kernel,
            k,
            k_blocks_valid,
            k_blocks_pad);
        check_hip_error("quant_mxfp4_kernel");
    }
    return {out_fp4, out_scale};
}

torch::Tensor hip_gemm_mxfp4(
    torch::Tensor a_fp4_u8,
    torch::Tensor b_u8,
    torch::Tensor a_scale_u8,
    torch::Tensor b_scale_u8,
    int64_t layout_mode,
    int64_t log2_k_split,
    torch::Tensor workspace,
    int64_t workspace_stride,
    int64_t b_shuffle_inner_mode) {
    TORCH_CHECK(a_fp4_u8.is_cuda(), "a_fp4_u8 must be CUDA");
    TORCH_CHECK(b_u8.is_cuda(), "b_u8 must be CUDA");
    TORCH_CHECK(a_scale_u8.is_cuda(), "a_scale_u8 must be CUDA");
    TORCH_CHECK(b_scale_u8.is_cuda(), "b_scale_u8 must be CUDA");
    TORCH_CHECK(a_fp4_u8.scalar_type() == torch::kUInt8, "a_fp4_u8 must be uint8");
    TORCH_CHECK(b_u8.scalar_type() == torch::kUInt8, "b_u8 must be uint8");
    TORCH_CHECK(a_scale_u8.scalar_type() == torch::kUInt8, "a_scale_u8 must be uint8");
    TORCH_CHECK(b_scale_u8.scalar_type() == torch::kUInt8, "b_scale_u8 must be uint8");
    TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");
    TORCH_CHECK(a_scale_u8.dim() == 2 && b_scale_u8.dim() == 2, "A/B scale must be 2D");
    TORCH_CHECK(layout_mode == 0 || layout_mode == 1, "layout_mode must be 0 or 1");

    auto a_fp4 = a_fp4_u8.contiguous();
    auto b_fp4 = b_u8.contiguous();
    auto a_scale = a_scale_u8.contiguous();
    auto b_scale = b_scale_u8.contiguous();

    int64_t m = a_fp4.size(0);
    int64_t n = b_fp4.size(0);
    int64_t k2 = a_fp4.size(1);
    TORCH_CHECK(b_fp4.size(1) == k2, "A/B K/2 mismatch");

    int64_t k = k2 * 2;
    TORCH_CHECK(k % BLOCK_K == 0, "K must be divisible by 32");
    int64_t k_blocks_valid = k / BLOCK_K;
    int64_t k_blocks_pad = a_scale.size(1);
    TORCH_CHECK(b_scale.size(1) == k_blocks_pad, "A/B scale padded K-block mismatch");
    TORCH_CHECK(a_scale.size(0) >= m, "a_scale rows must cover m");
    TORCH_CHECK(b_scale.size(0) >= n, "b_scale rows must cover n");

    int64_t split_k = 1;
    if (log2_k_split > 0) {
        split_k = static_cast<int64_t>(1) << log2_k_split;
    }
    if (split_k < 1) {
        split_k = 1;
    }

    int64_t mn = m * n;
    if (workspace_stride <= 0) {
        workspace_stride = mn;
    }
    TORCH_CHECK(workspace_stride >= mn, "workspace_stride must be >= m*n");

    auto ws = workspace;
    auto ws_opts = a_fp4.options().dtype(torch::kFloat);
    int64_t need = split_k * workspace_stride;
    if (!ws.defined() || !ws.is_cuda() || ws.scalar_type() != torch::kFloat || ws.numel() < need) {
        ws = torch::empty({split_k, workspace_stride}, ws_opts);
    } else {
        ws = ws.contiguous().view({split_k, workspace_stride});
    }

    auto out = torch::empty({m, n}, a_fp4.options().dtype(torch::kBFloat16));

    dim3 block(GEMM_THREADS);
    dim3 grid(
        static_cast<unsigned int>((n + GEMM_BLOCK_N - 1) / GEMM_BLOCK_N),
        static_cast<unsigned int>((m + GEMM_BLOCK_M - 1) / GEMM_BLOCK_M),
        static_cast<unsigned int>(split_k));

    hipLaunchKernelGGL(
        gemm_mxfp4_splitk_kernel,
        grid,
        block,
        0,
        0,
        reinterpret_cast<const uint8_t*>(a_fp4.data_ptr()),
        reinterpret_cast<const uint8_t*>(b_fp4.data_ptr()),
        reinterpret_cast<const uint8_t*>(a_scale.data_ptr()),
        reinterpret_cast<const uint8_t*>(b_scale.data_ptr()),
        reinterpret_cast<float*>(ws.data_ptr()),
        workspace_stride,
        m,
        n,
        k2,
        k_blocks_valid,
        k_blocks_pad,
        layout_mode,
        b_shuffle_inner_mode,
        split_k);
    check_hip_error("gemm_mxfp4_splitk_kernel");

    dim3 rblock(16, 16);
    dim3 rgrid(
        static_cast<unsigned int>((n + 15) / 16),
        static_cast<unsigned int>((m + 15) / 16));
    hipLaunchKernelGGL(
        reduce_splitk_kernel,
        rgrid,
        rblock,
        0,
        0,
        reinterpret_cast<const float*>(ws.data_ptr()),
        reinterpret_cast<__hip_bfloat16*>(out.data_ptr()),
        workspace_stride,
        m,
        n,
        split_k);
    check_hip_error("reduce_splitk_kernel");
    return out;
}
"""


def _sanitize_quant_backend(mode: str) -> str:
    mode = (mode or _DEFAULT_QUANT_BACKEND).strip().lower()
    if mode in {"auto", "triton", "hip"}:
        return mode
    return _DEFAULT_QUANT_BACKEND


def _get_quant_backend() -> str:
    return _sanitize_quant_backend(os.getenv(_QUANT_BACKEND_ENV, _DEFAULT_QUANT_BACKEND))


def _sanitize_exec_backend(mode: str) -> str:
    mode = (mode or _DEFAULT_EXEC_BACKEND).strip().lower()
    if mode in {"auto", "aiter", "hip"}:
        return mode
    return _DEFAULT_EXEC_BACKEND


def _get_exec_backend() -> str:
    mode = _sanitize_exec_backend(os.getenv(_EXEC_BACKEND_ENV, _DEFAULT_EXEC_BACKEND))
    if mode == "auto":
        return "aiter"
    return mode


def _sanitize_b_layout(mode: str) -> str:
    mode = (mode or _DEFAULT_B_LAYOUT).strip().lower()
    if mode in {"raw", "shuffle", "auto"}:
        return mode
    return _DEFAULT_B_LAYOUT


def _get_b_layout() -> str:
    mode = _sanitize_b_layout(os.getenv(_B_LAYOUT_ENV, _DEFAULT_B_LAYOUT))
    if mode == "auto":
        return "shuffle"
    return mode


def _get_splitk_override() -> int | None:
    raw = os.getenv(_SPLITK_ENV)
    if raw is None:
        return None
    try:
        return max(0, int(raw))
    except ValueError:
        return None


def _e8m0_unshuffle(scale_sh: torch.Tensor) -> torch.Tensor:
    if scale_sh.ndim != 2:
        raise RuntimeError(f"scale_sh must be 2D, got {tuple(scale_sh.shape)}")
    sm, sn = scale_sh.shape
    if sm % 32 != 0 or sn % 8 != 0:
        raise RuntimeError(f"scale_sh shape must be divisible by (32,8), got {(sm, sn)}")

    s = scale_sh.view(torch.uint8)
    s = s.view(sm // 32, sn // 8, 4, 16, 2, 2)
    s = s.permute(0, 5, 3, 1, 4, 2).contiguous()
    s = s.view(sm, sn)
    return s.view(scale_sh.dtype)


def _get_b_scale_raw_cached(b_scale_sh: torch.Tensor) -> torch.Tensor:
    dev = int(b_scale_sh.device.index) if b_scale_sh.device.index is not None else -1
    key = (
        int(b_scale_sh.data_ptr()),
        int(b_scale_sh.shape[0]),
        int(b_scale_sh.shape[1]),
        dev,
    )
    cached = _B_SCALE_RAW_CACHE.get(key)
    if cached is not None:
        return cached

    with _B_SCALE_RAW_LOCK:
        cached = _B_SCALE_RAW_CACHE.get(key)
        if cached is not None:
            return cached
        raw = _e8m0_unshuffle(b_scale_sh).contiguous()
        if len(_B_SCALE_RAW_CACHE) >= _B_SCALE_RAW_CACHE_MAX:
            _B_SCALE_RAW_CACHE.pop(next(iter(_B_SCALE_RAW_CACHE)))
        _B_SCALE_RAW_CACHE[key] = raw
        return raw


def _get_workspace(device: torch.device, m: int, n: int, split_k: int) -> torch.Tensor:
    dev = int(device.index) if device.index is not None else -1
    key = (dev, int(m), int(n), int(split_k))
    cached = _WORKSPACE_CACHE.get(key)
    need = split_k * m * n
    if cached is not None and cached.numel() >= need:
        return cached

    with _WORKSPACE_LOCK:
        cached = _WORKSPACE_CACHE.get(key)
        if cached is not None and cached.numel() >= need:
            return cached
        ws = torch.empty((split_k, m * n), dtype=torch.float32, device=device)
        if len(_WORKSPACE_CACHE) >= _WORKSPACE_CACHE_MAX:
            _WORKSPACE_CACHE.pop(next(iter(_WORKSPACE_CACHE)))
        _WORKSPACE_CACHE[key] = ws
        return ws


def _pick_splitk_log2(m: int, n: int, k: int) -> int:
    override = _get_splitk_override()
    if override is not None:
        return override

    key = (int(m), int(n), int(k))
    if key in _STATIC_SPLITK_LOG2:
        return _STATIC_SPLITK_LOG2[key]

    if m <= 16 and k >= 4096:
        return 3
    if m <= 32:
        return 2
    if m <= 64 and k >= 2048:
        return 1
    return 0


def _get_hip_module():
    global _HIP_MODULE, _HIP_BUILD_ERROR
    if _HIP_MODULE is not None:
        return _HIP_MODULE
    if _HIP_BUILD_ERROR is not None:
        raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")

    with _HIP_LOCK:
        if _HIP_MODULE is not None:
            return _HIP_MODULE
        if _HIP_BUILD_ERROR is not None:
            raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")

        try:
            os.environ.setdefault("CXX", "clang++")
            _HIP_MODULE = load_inline(
                name="mxfp4_mm_inline_quant_gemm_v4",
                cpp_sources=[CPP_WRAPPER],
                cuda_sources=[HIP_SRC],
                functions=["hip_quant_mxfp4", "hip_gemm_mxfp4"],
                verbose=False,
                extra_cuda_cflags=["-O3", "-std=c++20"],
            )
        except Exception as e:  # pragma: no cover - runtime dependent
            _HIP_BUILD_ERROR = e
            raise RuntimeError(f"HIP inline build failed: {e}") from e

    return _HIP_MODULE


def _quant_triton_mxfp4(x: torch.Tensor, shuffle: bool = True):
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    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)


def _run_aiter_pipeline(
    a_bf16: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    import aiter
    from aiter import dtypes

    a_q_sh, a_scale_sh = _quant_triton_mxfp4(a_bf16, shuffle=True)
    return aiter.gemm_a4w4(
        a_q_sh,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def _detect_b_shuffle_inner_mode(
    module,
    a_bf16: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> int:
    dev = int(a_bf16.device.index) if a_bf16.device.index is not None else -1
    key = (int(b_shuffle.shape[0]), int(b_shuffle.shape[1]), dev)
    cached = _B_SHUFFLE_INNER_CACHE.get(key)
    if cached is not None:
        return cached

    with _B_SHUFFLE_INNER_LOCK:
        cached = _B_SHUFFLE_INNER_CACHE.get(key)
        if cached is not None:
            return cached

        best_mode = 0
        try:
            import aiter
            from aiter import dtypes

            m_probe = min(8, int(a_bf16.shape[0]))
            a_probe = a_bf16[:m_probe, :].contiguous()

            a_q_sh, a_scale_sh = _quant_triton_mxfp4(a_probe, shuffle=True)
            ref = aiter.gemm_a4w4(
                a_q_sh,
                b_shuffle,
                a_scale_sh,
                b_scale_sh,
                dtype=dtypes.bf16,
                bpreshuffle=True,
            )

            a_q_raw, a_scale_raw = _quant_triton_mxfp4(a_probe, shuffle=False)
            a_fp4_u8 = a_q_raw.view(torch.uint8).contiguous()
            a_scale_u8 = a_scale_raw.view(torch.uint8).contiguous()
            b_u8 = b_shuffle.view(torch.uint8).contiguous()
            b_scale_u8 = b_scale_sh.view(torch.uint8).contiguous()
            ws = _get_workspace(a_probe.device, m_probe, int(b_shuffle.shape[0]), 1)

            errs = []
            for mode in (0, 1):
                out = module.hip_gemm_mxfp4(
                    a_fp4_u8,
                    b_u8,
                    a_scale_u8,
                    b_scale_u8,
                    1,
                    0,
                    ws,
                    int(m_probe * int(b_shuffle.shape[0])),
                    mode,
                )
                err = (out.float() - ref.float()).abs().max().item()
                errs.append(err)
            best_mode = 0 if errs[0] <= errs[1] else 1
        except Exception:
            best_mode = 0

        _B_SHUFFLE_INNER_CACHE[key] = best_mode
        return best_mode


def _run_hip_full_pipeline(
    a_bf16: torch.Tensor,
    b_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    module = _get_hip_module()
    quant_backend = _get_quant_backend()
    b_layout = _get_b_layout()

    if quant_backend == "hip":
        a_fp4_u8, a_scale_sh_u8 = module.hip_quant_mxfp4(a_bf16)
        a_fp4_u8 = a_fp4_u8[: a_bf16.shape[0], :].contiguous()
        a_scale_u8 = _e8m0_unshuffle(a_scale_sh_u8.view(torch.uint8)).contiguous()
    else:
        a_q, a_scale_raw = _quant_triton_mxfp4(a_bf16, shuffle=False)
        a_fp4_u8 = a_q.view(torch.uint8).contiguous()
        a_scale_u8 = a_scale_raw.view(torch.uint8).contiguous()

    if b_layout == "shuffle":
        b_u8 = b_shuffle.view(torch.uint8).contiguous()
        b_scale_u8 = b_scale_sh.view(torch.uint8).contiguous()
        layout_mode = 1
        b_inner_mode = _detect_b_shuffle_inner_mode(module, a_bf16, b_shuffle, b_scale_sh)
    else:
        b_u8 = b_q.view(torch.uint8).contiguous()
        b_scale_u8 = _get_b_scale_raw_cached(b_scale_sh).view(torch.uint8).contiguous()
        layout_mode = 0
        b_inner_mode = 0

    m = int(a_bf16.shape[0])
    k = int(a_bf16.shape[1])
    n = int(b_q.shape[0])
    log2_k_split = _pick_splitk_log2(m, n, k)
    split_k = 1 << log2_k_split
    ws = _get_workspace(a_bf16.device, m, n, split_k)

    return module.hip_gemm_mxfp4(
        a_fp4_u8,
        b_u8,
        a_scale_u8,
        b_scale_u8,
        int(layout_mode),
        int(log2_k_split),
        ws,
        int(m * n),
        int(b_inner_mode),
    )


def custom_kernel(data: input_t) -> output_t:
    """
    Default path uses aiter's preshuffled GEMM for leaderboard throughput.
    Set MXFP4_MM_EXEC_BACKEND=hip to force the custom HIP pipeline.
    """
    A, B, B_q, B_shuffle, B_scale_sh = data
    del B

    A = A.contiguous()
    B_q = B_q.contiguous()
    B_shuffle = B_shuffle.contiguous()
    B_scale_sh = B_scale_sh.contiguous()

    if _get_exec_backend() == "aiter":
        return _run_aiter_pipeline(
            a_bf16=A,
            b_shuffle=B_shuffle,
            b_scale_sh=B_scale_sh,
        )

    return _run_hip_full_pipeline(
        a_bf16=A,
        b_q=B_q,
        b_shuffle=B_shuffle,
        b_scale_sh=B_scale_sh,
    )
scrolls · 1084 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