Skip to content
KernelIndex
Search⌘K

submission 665044

Foreverwonder · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v27_amd_fixed.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-665044?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
13.5µs
#437 of 1143
2026-03-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bb5d5e36d99ff643b0ccb0f1a95f4a50937d94954588e207481312bbcb78b23d
license declaredunknown
license concludedunknown
authorsForeverwonder
imported2026-08-26

Techniques

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

shared-memory__shared__ __attribute__((aligned(16))) float s_vals[32 * GROUPS_PER_BLOCK];
split-kf.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")

Kernel source

submission_v27_amd_fixed.py320 lines
import os
import tempfile

os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


_RUNTIME = None
_NATIVE_MOD = None
_A_BUFFER_CACHE = {}

_CFG_ROWS = [
    (256, 4, 2880, 512, 21, 0, 5.05, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 8, 2112, 7168, 21, 1, 8.50, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 16, 2112, 7168, 21, 1, 12.90, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 16, 3072, 1536, 21, 1, 8.80, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 32, 4096, 512, 21, 0, 5.40, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 32, 2880, 512, 21, 0, 5.10, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 64, 7168, 2048, 21, 0, 6.50, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 64, 3072, 1536, 21, 1, 9.20, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 256, 3072, 1536, 21, 1, 7.80, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 256, 2880, 512, 21, 0, 5.50, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
]

_INLINE_CPP = r"""
#include <torch/extension.h>
void quant_mxfp4_a(torch::Tensor x, torch::Tensor x_fp4, torch::Tensor scale_sh);
"""

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

namespace {

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

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

__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
    __hip_bfloat16 v;
    *reinterpret_cast<uint16_t*>(&v) = x;
    return static_cast<float>(v);
}

__device__ __forceinline__ uint8_t quantize_e2m1(float x) {
    constexpr int EXP_BIAS_FP32 = 127;
    constexpr int EXP_BIAS_FP4 = 1;
    constexpr int MBITS_FP32 = 23;
    constexpr int MBITS_FP4 = 1;
    constexpr uint8_t MAX_INT = 0x7;
    constexpr uint8_t SIGN_MASK = 0x8;
    constexpr uint32_t MAGIC_ADDER = (1u << 21) - 1u;
    constexpr float MAX_NORMAL = 6.0f;
    constexpr float MIN_NORMAL = 1.0f;
    constexpr uint32_t DENORM_MASK_INT =
        static_cast<uint32_t>(((EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_FP32 - MBITS_FP4) + 1) << MBITS_FP32);
    constexpr int32_t VAL_TO_ADD =
        ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_FP32) + static_cast<int32_t>(MAGIC_ADDER);

    uint32_t ux = f32_as_u32(x);
    uint8_t sign_lp = static_cast<uint8_t>((ux >> 28) & SIGN_MASK);
    ux &= 0x7FFFFFFFu;
    float ax = u32_as_f32(ux);

    if (ax >= MAX_NORMAL) {
        return static_cast<uint8_t>(sign_lp | MAX_INT);
    }

    if (ax < MIN_NORMAL) {
        float denormal_f = ax + u32_as_f32(DENORM_MASK_INT);
        uint32_t denormal_u = f32_as_u32(denormal_f);
        uint8_t denormal_x = static_cast<uint8_t>(denormal_u - DENORM_MASK_INT);
        return static_cast<uint8_t>(sign_lp | (denormal_x & MAX_INT));
    }

    uint32_t mant_odd = (ux >> (MBITS_FP32 - MBITS_FP4)) & 1u;
    uint32_t normal_x = ux + static_cast<uint32_t>(VAL_TO_ADD) + mant_odd;
    uint8_t e2m1 = static_cast<uint8_t>(normal_x >> (MBITS_FP32 - MBITS_FP4));
    return static_cast<uint8_t>(sign_lp | (e2m1 & MAX_INT));
}

__device__ __forceinline__ int scale_shuffle_offset_fast(int row, int col, int scale_n_pad) {
    int offs_0 = row >> 5;
    int row_in_32 = row & 31;
    int offs_1 = row_in_32 >> 4;
    int offs_2 = row_in_32 & 15;
    int offs_3 = col >> 3;
    int col_in_8 = col & 7;
    int offs_4 = col_in_8 >> 2;
    int offs_5 = col_in_8 & 3;
    return offs_1 + (offs_4 << 1) + (offs_2 << 2) + (offs_5 << 6) + (offs_3 << 8) + (offs_0 << 5) * scale_n_pad;
}

__device__ __forceinline__ float warp_reduce_max_amd(float val) {
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        float other = __shfl_down(val, offset);
        val = (val > other) ? val : other;
    }
    return val;
}

template <int GROUPS_PER_BLOCK, int BLOCK_SIZE>
__global__ void __launch_bounds__(BLOCK_SIZE) quant_mxfp4_a_kernel_amd_fixed(
    const uint16_t* __restrict__ x,
    uint8_t* __restrict__ x_fp4,
    uint8_t* __restrict__ scale_sh,
    int m,
    int k,
    int scale_n_valid,
    int scale_n_pad
) {
    __shared__ __attribute__((aligned(16))) float s_vals[32 * GROUPS_PER_BLOCK];
    __shared__ __attribute__((aligned(4))) uint8_t s_scale[GROUPS_PER_BLOCK];

    constexpr int GROUP_SIZE = 32;
    constexpr int PAIRS_PER_GROUP = GROUP_SIZE >> 1;

    int tid = static_cast<int>(threadIdx.x);
    int lane = tid & 31;
    int group_in_block = tid >> 5;
    int row = static_cast<int>(blockIdx.y);
    int scale_col = (static_cast<int>(blockIdx.x) * GROUPS_PER_BLOCK) + group_in_block;
    bool active_group = (row < m) && (scale_col < scale_n_valid);

    float value = 0.0f;
    if (active_group) {
        int x_idx = (row * k) + (scale_col * GROUP_SIZE) + lane;
        value = bf16_to_f32(x[x_idx]);
    }

    float abs_val = __builtin_fabsf(value);
    float max_abs = warp_reduce_max_amd(abs_val);
    
    if (lane == 0) {
        uint32_t max_bits = f32_as_u32(max_abs);
        max_bits = (max_bits + 0x00200000u) & 0xFF800000u;
        uint8_t scale_byte = 0;
        if (max_bits != 0) {
            scale_byte = static_cast<uint8_t>(((max_bits >> 23) & 0xFFu) - 2u);
        }
        s_scale[group_in_block] = scale_byte;
        if (active_group) {
            int scale_idx = scale_shuffle_offset_fast(row, scale_col, scale_n_pad);
            scale_sh[scale_idx] = scale_byte;
        }
    }

    int group_base = group_in_block << 5;
    s_vals[group_base + lane] = value;
    __syncthreads();

    if ((lane < PAIRS_PER_GROUP) && active_group) {
        uint8_t scale_byte = s_scale[group_in_block];
        float inv_scale = 0.0f;
        if (scale_byte != 0) {
            float scale_f = u32_as_f32(static_cast<uint32_t>(scale_byte) << 23);
            inv_scale = 1.0f / scale_f;
        }
        
        int lane2 = lane << 1;
        float x0 = s_vals[group_base + lane2] * inv_scale;
        float x1 = s_vals[group_base + lane2 + 1] * inv_scale;
        
        uint8_t q0 = quantize_e2m1(x0);
        uint8_t q1 = quantize_e2m1(x1);
        
        int out_idx = (row * (k >> 1)) + (scale_col * PAIRS_PER_GROUP) + lane;
        x_fp4[out_idx] = static_cast<uint8_t>(q0 | (q1 << 4));
    }
}

}  // namespace

void quant_mxfp4_a(torch::Tensor x, torch::Tensor x_fp4, torch::Tensor scale_sh) {
    TORCH_CHECK(x.is_cuda(), "x must be a CUDA tensor");
    TORCH_CHECK(x_fp4.is_cuda(), "x_fp4 must be a CUDA tensor");
    TORCH_CHECK(scale_sh.is_cuda(), "scale_sh must be a CUDA tensor");
    TORCH_CHECK(x.scalar_type() == at::ScalarType::BFloat16, "x must be bf16");
    TORCH_CHECK(x_fp4.scalar_type() == at::ScalarType::Byte, "x_fp4 must be uint8");
    TORCH_CHECK(scale_sh.scalar_type() == at::ScalarType::Byte, "scale_sh must be uint8");
    TORCH_CHECK(x.dim() == 2, "x must be 2D");
    TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
    TORCH_CHECK(x_fp4.is_contiguous(), "x_fp4 must be contiguous");
    TORCH_CHECK(scale_sh.is_contiguous(), "scale_sh must be contiguous");

    int m = static_cast<int>(x.size(0));
    int k = static_cast<int>(x.size(1));
    int scale_n_valid = k >> 5;
    int scale_n_pad = static_cast<int>(scale_sh.size(1));

    TORCH_CHECK((k & 63) == 0, "k must be divisible by 64");
    TORCH_CHECK(x_fp4.size(0) == x.size(0), "x_fp4 row mismatch");
    TORCH_CHECK(x_fp4.size(1) == (x.size(1) >> 1), "x_fp4 col mismatch");

    if (k >= 2048) {
        dim3 grid((scale_n_valid + 3) >> 2, m);
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME(quant_mxfp4_a_kernel_amd_fixed<4, 128>),
            grid, dim3(128), 0, 0,
            reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>()),
            x_fp4.data_ptr<uint8_t>(),
            scale_sh.data_ptr<uint8_t>(),
            m, k, scale_n_valid, scale_n_pad
        );
    } else {
        dim3 grid((scale_n_valid + 1) >> 1, m);
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME(quant_mxfp4_a_kernel_amd_fixed<2, 64>),
            grid, dim3(64), 0, 0,
            reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>()),
            x_fp4.data_ptr<uint8_t>(),
            scale_sh.data_ptr<uint8_t>(),
            m, k, scale_n_valid, scale_n_pad
        );
    }
}
"""


def _get_native_mod():
    global _NATIVE_MOD
    if _NATIVE_MOD is None:
        _NATIVE_MOD = load_inline(
            name="mxfp4_mm_native_hip_quant_v27_amd_fixed",
            cpp_sources=[_INLINE_CPP],
            cuda_sources=[_INLINE_HIP],
            functions=["quant_mxfp4_a"],
            extra_cuda_cflags=[
                "--offload-arch=gfx950",
                "-O3",
                "-std=c++20",
            ],
            verbose=False,
        )
    return _NATIVE_MOD


def _ensure_runtime():
    global _RUNTIME
    if _RUNTIME is not None:
        return _RUNTIME

    cfg_path = os.path.join(tempfile.gettempdir(), "aiter_mxfp4_mm_submission_v27.csv")
    if not os.path.exists(cfg_path):
        with open(cfg_path, "w", encoding="ascii", newline="") as f:
            f.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")
            for cu_num, m, n, k, kernel_id, split_k, us, kernel_name in _CFG_ROWS:
                f.write(
                    f"{cu_num},{m},{n},{k},{kernel_id},{split_k},{us},{kernel_name},0.0,0.0,0.0\n"
                )

    default_cfg = "/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
    os.environ["AITER_CONFIG_GEMM_A4W4"] = cfg_path + os.pathsep + default_cfg

    import aiter
    from aiter import dtypes

    _RUNTIME = (aiter, dtypes)
    return _RUNTIME


def _native_quant_a(A: torch.Tensor, dtypes):
    mod = _get_native_mod()
    m, k = A.shape
    scale_m_pad = ((m + 255) // 256) * 256
    scale_n_valid = k // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8

    cache_key = (A.device.index, m, k)
    cached = _A_BUFFER_CACHE.get(cache_key)
    if cached is None:
        A_q_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=A.device)
        A_scale_u8 = torch.zeros((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=A.device)
        cached = (
            A_q_u8,
            A_scale_u8,
            A_q_u8.view(dtypes.fp4x2),
            A_scale_u8.view(dtypes.fp8_e8m0),
        )
        _A_BUFFER_CACHE[cache_key] = cached

    A_q_u8, A_scale_u8, A_q_view, A_scale_view = cached
    mod.quant_mxfp4_a(A, A_q_u8, A_scale_u8)
    return A_q_view, A_scale_view


def custom_kernel(data: input_t) -> output_t:
    aiter, dtypes = _ensure_runtime()
    A, _, _, B_shuffle, B_scale_sh = data
    if not A.is_contiguous():
        A = A.contiguous()

    A_q, A_scale_sh = _native_quant_a(A, dtypes)

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