Skip to content
KernelIndex
Search⌘K

submission 629672

j1ang6566 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b70d12810cdaaaa2e784fd2043491318f2f49fa773ffb2cffb75d934931d2e1b
license declaredunknown
license concludedunknown
authorsj1ang6566
imported2026-08-26

Kernel source

submission_v10.py551 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
V10 is a leaderboard-oriented branch from V4: keep the fused HIP quant path,
but use more aggressive launch settings keyed by the official small-M A-shapes.
"""
import os
from time import perf_counter_ns
from typing import Any

try:
    from task import input_t, output_t
except ImportError:
    input_t = Any
    output_t = Any

OFFICIAL_BENCHMARK_SHAPES = (
    (4, 2880, 512),
    (16, 2112, 7168),
    (32, 4096, 512),
    (32, 2880, 512),
    (64, 7168, 2048),
    (256, 3072, 1536),
)

_HIP_SMALL_M_SHAPES = {
    (4, 2880, 512),
    (16, 2112, 7168),
    (32, 4096, 512),
    (32, 2880, 512),
}

_AITER_CACHE = None
_HIP_QUANT_MODULE = None
_HIP_QUANT_ERROR = None
_LAST_PROFILE = None
_TORCH_CACHE = None

_HIP_CPP_SRC = r"""
#include <torch/extension.h>

void quant_mxfp4_hip(torch::Tensor input, torch::Tensor out_fp4, torch::Tensor out_scale_sh);
"""

_HIP_CUDA_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>

#include <cmath>
#include <cstdint>
#include <stdexcept>

namespace {

constexpr uint32_t F32_SIGN_MASK = 0x80000000u;
constexpr uint32_t MX_SCALE_ROUND_BIT = 0x00200000u;
constexpr uint32_t MX_SCALE_MASK = 0xFF800000u;
constexpr uint32_t FP4_SIGN_MASK = 0x8u;
constexpr uint32_t FP4_MAX_INT = 0x7u;
constexpr uint32_t FP4_MAGIC_ADDER = (1u << 21) - 1u;
constexpr uint32_t FP4_DENORM_MASK_INT = 149u << 23;
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;

template <typename T>
__device__ inline T shfl_down_32(T value, int offset) {
    return __shfl_down(value, offset, 32);
}

template <typename T>
__device__ inline T shfl_32(T value, int src_lane) {
    return __shfl(value, src_lane, 32);
}

__device__ inline uint32_t float_as_uint(float x) {
    union {
        float f;
        uint32_t u;
    } bits;
    bits.f = x;
    return bits.u;
}

__device__ inline float uint_as_float(uint32_t x) {
    union {
        float f;
        uint32_t u;
    } bits;
    bits.u = x;
    return bits.f;
}

__device__ inline float bf16_to_float(uint16_t x) {
    return uint_as_float(static_cast<uint32_t>(x) << 16);
}

__device__ inline uint8_t float_to_e2m1(float x) {
    uint32_t bits = float_as_uint(x);
    uint32_t sign = bits & F32_SIGN_MASK;
    uint32_t abs_bits = bits ^ sign;
    float abs_x = uint_as_float(abs_bits);

    uint8_t code;
    if (abs_x >= FP4_MAX_NORMAL) {
        code = static_cast<uint8_t>(FP4_MAX_INT);
    } else if (abs_x < FP4_MIN_NORMAL) {
        float denorm_x = abs_x + uint_as_float(FP4_DENORM_MASK_INT);
        int32_t denorm_i = static_cast<int32_t>(float_as_uint(denorm_x)) - static_cast<int32_t>(FP4_DENORM_MASK_INT);
        code = static_cast<uint8_t>(denorm_i);
    } else {
        int32_t normal_i = static_cast<int32_t>(abs_bits);
        int32_t mant_odd = (normal_i >> 22) & 1;
        int32_t val_to_add = ((1 - 127) << 23) + static_cast<int32_t>(FP4_MAGIC_ADDER);
        normal_i += val_to_add;
        normal_i += mant_odd;
        normal_i >>= 22;
        code = static_cast<uint8_t>(normal_i);
    }

    uint8_t sign_lp = static_cast<uint8_t>((sign >> 28) & FP4_SIGN_MASK);
    return static_cast<uint8_t>(code | sign_lp);
}

__device__ inline int64_t shuffled_scale_offset(int row, int group, int64_t sn8) {
    const int64_t row_block = row >> 5;
    const int64_t row_sub = (row >> 4) & 1;
    const int64_t row_in16 = row & 15;
    const int64_t col_block = group >> 3;
    const int64_t col_hi = (group >> 2) & 1;
    const int64_t col_lo = group & 3;
    return (((((row_block * (sn8 >> 3) + col_block) * 4 + col_lo) * 16 + row_in16) * 2 + col_hi) * 2 + row_sub);
}

template <int WARPS_PER_BLOCK>
__global__ void quant_mxfp4_small_m_kernel(
    const uint16_t* __restrict__ input,
    uint8_t* __restrict__ out_fp4,
    uint8_t* __restrict__ out_scale_sh,
    int64_t groups,
    int64_t k_half,
    int64_t sn8
) {
    const int row = static_cast<int>(blockIdx.y);
    const int warp_id = threadIdx.x >> 5;
    const int lane = threadIdx.x & 31;
    const int group = static_cast<int>(blockIdx.x) * WARPS_PER_BLOCK + warp_id;

    if (group >= groups) {
        return;
    }

    const int64_t input_base = (static_cast<int64_t>(row) * groups + group) * 32;
    const int64_t out_base = static_cast<int64_t>(row) * k_half + static_cast<int64_t>(group) * 16;
    const uint16_t raw = input[input_base + lane];
    const float x = bf16_to_float(raw);
    float amax = fabsf(x);

    for (int offset = 16; offset > 0; offset >>= 1) {
        amax = fmaxf(amax, shfl_down_32(amax, offset));
    }

    int scale_unbiased = -127;
    uint8_t scale = 0;
    if (lane == 0) {
        const uint32_t amax_bits = float_as_uint(amax);
        const uint32_t rounded_bits = (amax_bits + MX_SCALE_ROUND_BIT) & MX_SCALE_MASK;

        if (rounded_bits != 0u) {
            const int32_t exp_bits = static_cast<int32_t>((rounded_bits >> 23) & 0xffu);
            scale_unbiased = exp_bits - 129;
            if (scale_unbiased < -127) {
                scale_unbiased = -127;
            } else if (scale_unbiased > 127) {
                scale_unbiased = 127;
            }
        }

        scale = static_cast<uint8_t>(scale_unbiased + 127);
        out_scale_sh[shuffled_scale_offset(row, group, sn8)] = scale;
    }
    scale_unbiased = shfl_32(scale_unbiased, 0);

    const uint8_t q = float_to_e2m1(ldexpf(x, -scale_unbiased));
    const uint32_t q_hi = static_cast<uint32_t>(shfl_down_32(static_cast<uint32_t>(q), 1));

    if ((lane & 1) == 0) {
        out_fp4[out_base + (lane >> 1)] = static_cast<uint8_t>((q_hi << 4) | q);
    }
}

}  // namespace

void quant_mxfp4_hip(torch::Tensor input, torch::Tensor out_fp4, torch::Tensor out_scale_sh) {
    TORCH_CHECK(input.is_cuda(), "input must be a ROCm tensor");
    TORCH_CHECK(input.scalar_type() == at::kBFloat16, "input must be bfloat16");
    TORCH_CHECK(input.dim() == 2, "input must be 2D");
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");

    TORCH_CHECK(out_fp4.is_cuda(), "out_fp4 must be a ROCm tensor");
    TORCH_CHECK(out_scale_sh.is_cuda(), "out_scale_sh must be a ROCm tensor");
    TORCH_CHECK(out_fp4.scalar_type() == at::kByte, "out_fp4 must be uint8");
    TORCH_CHECK(out_scale_sh.scalar_type() == at::kByte, "out_scale_sh must be uint8");
    TORCH_CHECK(out_fp4.is_contiguous(), "out_fp4 must be contiguous");
    TORCH_CHECK(out_scale_sh.is_contiguous(), "out_scale_sh must be contiguous");

    const int64_t m = input.size(0);
    const int64_t k = input.size(1);
    TORCH_CHECK(k % 32 == 0, "K must be divisible by 32");

    const int64_t groups = k / 32;
    const int64_t k_half = k / 2;
    const int64_t padded_m = ((m + 255) / 256) * 256;
    const int64_t sn8 = ((groups + 7) / 8) * 8;
    TORCH_CHECK(out_fp4.size(0) == m && out_fp4.size(1) == k_half, "unexpected out_fp4 shape");
    TORCH_CHECK(out_scale_sh.size(0) == padded_m && out_scale_sh.size(1) == sn8, "unexpected out_scale_sh shape");

    const auto* input_ptr = reinterpret_cast<const uint16_t*>(input.data_ptr<at::BFloat16>());
    auto* out_fp4_ptr = out_fp4.data_ptr<uint8_t>();
    auto* out_scale_ptr = out_scale_sh.data_ptr<uint8_t>();

    if (k == 512) {
        if (m <= 4) {
            constexpr int warps_per_block = 2;
            dim3 blocks(static_cast<unsigned int>((groups + warps_per_block - 1) / warps_per_block), static_cast<unsigned int>(m));
            dim3 threads(32 * warps_per_block);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(quant_mxfp4_small_m_kernel<warps_per_block>),
                blocks,
                threads,
                0,
                0,
                input_ptr,
                out_fp4_ptr,
                out_scale_ptr,
                groups,
                k_half,
                sn8
            );
        } else if (m >= 32) {
            constexpr int warps_per_block = 8;
            dim3 blocks(static_cast<unsigned int>((groups + warps_per_block - 1) / warps_per_block), static_cast<unsigned int>(m));
            dim3 threads(32 * warps_per_block);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(quant_mxfp4_small_m_kernel<warps_per_block>),
                blocks,
                threads,
                0,
                0,
                input_ptr,
                out_fp4_ptr,
                out_scale_ptr,
                groups,
                k_half,
                sn8
            );
        } else {
            constexpr int warps_per_block = 4;
            dim3 blocks(static_cast<unsigned int>((groups + warps_per_block - 1) / warps_per_block), static_cast<unsigned int>(m));
            dim3 threads(32 * warps_per_block);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(quant_mxfp4_small_m_kernel<warps_per_block>),
                blocks,
                threads,
                0,
                0,
                input_ptr,
                out_fp4_ptr,
                out_scale_ptr,
                groups,
                k_half,
                sn8
            );
        }
    } else if (k == 7168) {
        if (m <= 16) {
            constexpr int warps_per_block = 12;
            dim3 blocks(static_cast<unsigned int>((groups + warps_per_block - 1) / warps_per_block), static_cast<unsigned int>(m));
            dim3 threads(32 * warps_per_block);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(quant_mxfp4_small_m_kernel<warps_per_block>),
                blocks,
                threads,
                0,
                0,
                input_ptr,
                out_fp4_ptr,
                out_scale_ptr,
                groups,
                k_half,
                sn8
            );
        } else {
            constexpr int warps_per_block = 8;
            dim3 blocks(static_cast<unsigned int>((groups + warps_per_block - 1) / warps_per_block), static_cast<unsigned int>(m));
            dim3 threads(32 * warps_per_block);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(quant_mxfp4_small_m_kernel<warps_per_block>),
                blocks,
                threads,
                0,
                0,
                input_ptr,
                out_fp4_ptr,
                out_scale_ptr,
                groups,
                k_half,
                sn8
            );
        }
    } else {
        TORCH_CHECK(false, "quant_mxfp4_hip only supports K=512 or K=7168");
    }

    hipError_t err = hipGetLastError();
    if (err != hipSuccess) {
        throw std::runtime_error(hipGetErrorString(err));
    }
}
"""


def _load_torch():
    global _TORCH_CACHE
    if _TORCH_CACHE is None:
        import torch

        _TORCH_CACHE = torch
    return _TORCH_CACHE


def _profiling_enabled():
    value = os.getenv("MXFP4_PROFILE", "")
    return value.lower() not in ("", "0", "false", "no", "off")


def _hip_quant_enabled():
    value = os.getenv("MXFP4_V10_QUANT_IMPL", "auto")
    return value.lower() in ("", "1", "auto", "hip", "inline", "native_hip")


def consume_last_profile():
    global _LAST_PROFILE
    profile = _LAST_PROFILE
    _LAST_PROFILE = None
    return profile


def _maybe_sync():
    if not _profiling_enabled():
        return
    try:
        torch = _load_torch()
    except ImportError:
        return
    if torch.cuda.is_available():
        torch.cuda.synchronize()


def _profile_start(profile):
    if profile is None:
        return None
    _maybe_sync()
    return perf_counter_ns()


def _profile_stop(profile, key, start_ns):
    if start_ns is None:
        return
    _maybe_sync()
    profile[key] = profile.get(key, 0.0) + (perf_counter_ns() - start_ns) / 1_000.0


def _load_aiter_symbols():
    global _AITER_CACHE
    if _AITER_CACHE is None:
        import aiter
        from aiter import dtypes
        from aiter.ops.triton.quant import dynamic_mxfp4_quant
        from aiter.utility.fp4_utils import e8m0_shuffle

        _AITER_CACHE = (
            aiter.gemm_a4w4,
            dtypes,
            dynamic_mxfp4_quant,
            e8m0_shuffle,
        )
    return _AITER_CACHE


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


def _load_hip_quant_module():
    global _HIP_QUANT_MODULE, _HIP_QUANT_ERROR
    if _HIP_QUANT_MODULE is not None:
        return _HIP_QUANT_MODULE
    if _HIP_QUANT_ERROR is not None:
        raise RuntimeError(_HIP_QUANT_ERROR)

    _prepare_rocm_env()
    try:
        _load_torch()
        from torch.utils.cpp_extension import load_inline

        arch = os.getenv("PYTORCH_ROCM_ARCH", "gfx950")
        verbose = os.getenv("MXFP4_V10_VERBOSE_BUILD", "").lower() not in ("", "0", "false", "no", "off")
        _HIP_QUANT_MODULE = load_inline(
            name=f"mxfp4_v10_quant_{arch}",
            cpp_sources=[_HIP_CPP_SRC],
            cuda_sources=[_HIP_CUDA_SRC],
            functions=["quant_mxfp4_hip"],
            verbose=verbose,
            extra_cflags=["-O3"],
            extra_cuda_cflags=[f"--offload-arch={arch}", "-O3", "-std=c++20"],
        )
        return _HIP_QUANT_MODULE
    except Exception as exc:
        _HIP_QUANT_ERROR = f"{type(exc).__name__}: {exc}"
        raise RuntimeError(_HIP_QUANT_ERROR) from exc


def _view_quant_outputs(x_fp4, scale_sh, dtypes):
    return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)


def _ensure_contiguous(x, profile):
    start_ns = _profile_start(profile)
    if not x.is_contiguous():
        x = x.contiguous()
    _profile_stop(profile, "layout_us", start_ns)
    return x


def _quant_mxfp4_native(x, dynamic_mxfp4_quant, e8m0_shuffle, dtypes, profile):
    start_ns = _profile_start(profile)
    x_fp4, scale = dynamic_mxfp4_quant(x)
    scale_sh = e8m0_shuffle(scale)
    _profile_stop(profile, "quant_us", start_ns)
    if profile is not None:
        profile["quant_impl"] = "aiter_dynamic_mxfp4_quant"
    return _view_quant_outputs(x_fp4, scale_sh, dtypes)


def _padded_scale_shape(m, groups):
    return ((m + 255) // 256 * 256, (groups + 7) // 8 * 8)


def _small_m_kernel_kind(shape_key):
    k = shape_key[2]
    if k == 512:
        if shape_key[0] <= 4:
            return "k512_m4_w2"
        if shape_key[0] >= 32:
            return "k512_m32_w8"
        return "k512_w4"
    if k == 7168:
        if shape_key[0] <= 16:
            return "k7168_m16_w12"
        return "k7168_w8"
    return None


def _quant_mxfp4_hip_small_m(x, shape_key, dtypes, profile):
    torch = _load_torch()
    module = _load_hip_quant_module()
    kernel_kind = _small_m_kernel_kind(shape_key)
    if kernel_kind is None:
        raise RuntimeError(f"unsupported small-M shape: {shape_key}")

    start_ns = _profile_start(profile)
    m, k = x.shape
    x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=x.device)
    scale_shape = _padded_scale_shape(m, k // 32)
    scale_sh = torch.empty(scale_shape, dtype=torch.uint8, device=x.device)
    module.quant_mxfp4_hip(x, x_fp4, scale_sh)
    _profile_stop(profile, "quant_us", start_ns)

    if profile is not None:
        profile["quant_impl"] = f"inline_hip_fused_shuffle_{kernel_kind}"
    return _view_quant_outputs(x_fp4, scale_sh, dtypes)


def _quant_mxfp4_v10(x, shape_key, dynamic_mxfp4_quant, e8m0_shuffle, dtypes, profile):
    if _hip_quant_enabled() and shape_key in _HIP_SMALL_M_SHAPES:
        try:
            return _quant_mxfp4_hip_small_m(x, shape_key, dtypes, profile)
        except Exception as exc:
            if profile is not None:
                profile["quant_fallback"] = type(exc).__name__

    return _quant_mxfp4_native(x, dynamic_mxfp4_quant, e8m0_shuffle, dtypes, profile)


def _run_quant_gemm(
    A,
    shape_key,
    B_shuffle,
    B_scale_sh,
    gemm_a4w4,
    dtypes,
    dynamic_mxfp4_quant,
    e8m0_shuffle,
    profile,
):
    A = _ensure_contiguous(A, profile)
    A_q, A_scale_sh = _quant_mxfp4_v10(A, shape_key, dynamic_mxfp4_quant, e8m0_shuffle, dtypes, profile)

    gemm_start_ns = _profile_start(profile)
    out_gemm = gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    _profile_stop(profile, "gemm_us", gemm_start_ns)
    return out_gemm


def custom_kernel(data: input_t) -> output_t:
    global _LAST_PROFILE

    gemm_a4w4, dtypes, dynamic_mxfp4_quant, e8m0_shuffle = _load_aiter_symbols()

    A, B, _B_q, B_shuffle, B_scale_sh = data
    shape_key = (A.shape[0], B.shape[0], A.shape[1])
    profile = {"shape": shape_key} if _profiling_enabled() else None
    total_start_ns = _profile_start(profile)

    out_gemm = _run_quant_gemm(
        A,
        shape_key,
        B_shuffle,
        B_scale_sh,
        gemm_a4w4,
        dtypes,
        dynamic_mxfp4_quant,
        e8m0_shuffle,
        profile,
    )

    _profile_stop(profile, "total_us", total_start_ns)
    if profile is not None:
        _LAST_PROFILE = profile

    return out_gemm
scrolls · 551 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