Skip to content
KernelIndex
Search⌘K

submission 648988

j1atng · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v6b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-648988?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
14.9µs
#541 of 1143
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bffa874b7e7a07e6c243ca5f373c1bf55b62009b2ce6f2adcf175ffc0bf99880
license declaredunknown
license concludedunknown
authorsj1atng
imported2026-08-26

Techniques

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

fp4TORCH_CHECK(out_fp4.size(0)==m && out_fp4.size(1)==k_half, "fp4 shape mismatch");

Kernel source

submission_v6b.py265 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
V6b fixes the V6 timeout:
  - HIP compilation and aiter symbol loading now happen at MODULE IMPORT TIME,
    not inside the first custom_kernel call. This keeps compilation out of the
    timing window entirely.
  - try/except restored around load_inline so a compiler hang/failure falls back
    gracefully to the aiter Triton quant path (same as V5).
  - Hot path remains flat: one dict lookup + HIP call + GEMM call.
  - Pre-computed fp4x2/fp8_e8m0 views (share storage with raw buffers).
  - All 6 benchmark shapes use the HIP fast path.
"""
import os
from typing import Any

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

# ---------------------------------------------------------------------------
# HIP kernel source – generic K support via groups-based dispatch
# ---------------------------------------------------------------------------
_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 v, int off) { return __shfl_down(v, off, 32); }
template <typename T>
__device__ inline T shfl_32(T v, int lane)     { return __shfl(v, lane, 32); }

__device__ inline uint32_t float_as_uint(float x) {
    union { float f; uint32_t u; } b; b.f = x; return b.u;
}
__device__ inline float uint_as_float(uint32_t x) {
    union { float f; uint32_t u; } b; b.u = x; return b.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   dx = abs_x + uint_as_float(FP4_DENORM_MASK_INT);
        int32_t di = static_cast<int32_t>(float_as_uint(dx))
                   - static_cast<int32_t>(FP4_DENORM_MASK_INT);
        code = static_cast<uint8_t>(di);
    } else {
        int32_t ni   = static_cast<int32_t>(abs_bits);
        int32_t modd = (ni >> 22) & 1;
        ni += ((1 - 127) << 23) + static_cast<int32_t>(FP4_MAGIC_ADDER);
        ni += modd;
        ni >>= 22;
        code = static_cast<uint8_t>(ni);
    }
    return static_cast<uint8_t>(code | static_cast<uint8_t>((sign >> 28) & FP4_SIGN_MASK));
}

__device__ inline int64_t shuffled_scale_offset(int row, int group, int64_t sn8) {
    int64_t rb = row >> 5, rs = (row >> 4) & 1, r16 = row & 15;
    int64_t cb = group >> 3, ch = (group >> 2) & 1, cl = group & 3;
    return (((((rb * (sn8 >> 3) + cb) * 4 + cl) * 16 + r16) * 2 + ch) * 2 + rs);
}

template <int W>
__global__ void quant_mxfp4_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)
{
    int row   = static_cast<int>(blockIdx.y);
    int wid   = threadIdx.x >> 5;
    int lane  = threadIdx.x & 31;
    int group = static_cast<int>(blockIdx.x) * W + wid;
    if (group >= groups) return;

    int64_t ib = (static_cast<int64_t>(row) * groups + group) * 32;
    int64_t ob = static_cast<int64_t>(row) * k_half + static_cast<int64_t>(group) * 16;
    float   x  = bf16_to_float(input[ib + lane]);
    float   am = fabsf(x);
    for (int o = 16; o > 0; o >>= 1) am = fmaxf(am, shfl_down_32(am, o));

    int su = -127;
    if (lane == 0) {
        uint32_t rb = (float_as_uint(am) + MX_SCALE_ROUND_BIT) & MX_SCALE_MASK;
        if (rb) {
            su = static_cast<int32_t>((rb >> 23) & 0xffu) - 129;
            if (su < -127) su = -127;
            else if (su > 127) su = 127;
        }
        out_scale_sh[shuffled_scale_offset(row, group, sn8)] =
            static_cast<uint8_t>(su + 127);
    }
    su = shfl_32(su, 0);

    uint8_t  q   = float_to_e2m1(ldexpf(x, -su));
    uint32_t qhi = static_cast<uint32_t>(shfl_down_32(static_cast<uint32_t>(q), 1));
    if ((lane & 1) == 0)
        out_fp4[ob + (lane >> 1)] = static_cast<uint8_t>((qhi << 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.scalar_type() == at::kBFloat16
                && input.dim() == 2 && input.is_contiguous(), "bad input");
    TORCH_CHECK(out_fp4.is_cuda() && out_fp4.scalar_type() == at::kByte
                && out_fp4.is_contiguous(), "bad out_fp4");
    TORCH_CHECK(out_scale_sh.is_cuda() && out_scale_sh.scalar_type() == at::kByte
                && out_scale_sh.is_contiguous(), "bad out_scale_sh");

    int64_t m = input.size(0), k = input.size(1);
    TORCH_CHECK(k % 32 == 0, "K must be divisible by 32");
    int64_t groups   = k / 32;
    int64_t k_half   = k / 2;
    int64_t padded_m = ((m + 255) / 256) * 256;
    int64_t sn8      = ((groups + 7) / 8) * 8;
    TORCH_CHECK(out_fp4.size(0)==m && out_fp4.size(1)==k_half, "fp4 shape mismatch");
    TORCH_CHECK(out_scale_sh.size(0)==padded_m && out_scale_sh.size(1)==sn8,
                "scale shape mismatch");

    const auto* ip = reinterpret_cast<const uint16_t*>(input.data_ptr<at::BFloat16>());
    auto*       fp = out_fp4.data_ptr<uint8_t>();
    auto*       sp = out_scale_sh.data_ptr<uint8_t>();

    // warps=4 for K=512 (groups=16), warps=8 for larger K
    if (groups <= 16) {
        constexpr int W = 4;
        dim3 blk(static_cast<unsigned>((groups+W-1)/W), static_cast<unsigned>(m));
        hipLaunchKernelGGL(HIP_KERNEL_NAME(quant_mxfp4_kernel<W>),
                           blk, dim3(32*W), 0, 0, ip, fp, sp, groups, k_half, sn8);
    } else {
        constexpr int W = 8;
        dim3 blk(static_cast<unsigned>((groups+W-1)/W), static_cast<unsigned>(m));
        hipLaunchKernelGGL(HIP_KERNEL_NAME(quant_mxfp4_kernel<W>),
                           blk, dim3(32*W), 0, 0, ip, fp, sp, groups, k_half, sn8);
    }

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

# ---------------------------------------------------------------------------
# Module-level initialisation – runs at import time, NOT inside custom_kernel
# This keeps compilation and symbol loading out of the benchmark timing window.
# ---------------------------------------------------------------------------

import aiter as _aiter
from aiter import dtypes as _dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant as _dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle as _e8m0_shuffle

_GEMM_FN = _aiter.gemm_a4w4
_DTYPES  = _dtypes

# Compile HIP quant kernel (cached to disk by PyTorch after first compile).
# On a cache hit this is nearly instant; on a miss it takes ~30-60s but that
# happens during module load, before any timing starts.
_HIP_MODULE = None
try:
    os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
    os.environ.setdefault("CXX", "clang++")
    from torch.utils.cpp_extension import load_inline as _load_inline
    _arch = os.getenv("PYTORCH_ROCM_ARCH", "gfx950")
    _HIP_MODULE = _load_inline(
        name=f"mxfp4_v6b_quant_{_arch}",
        cpp_sources=[_HIP_CPP_SRC],
        cuda_sources=[_HIP_CUDA_SRC],
        functions=["quant_mxfp4_hip"],
        verbose=False,
        extra_cflags=["-O3"],
        extra_cuda_cflags=[f"--offload-arch={_arch}", "-O3", "-std=c++20"],
    )
except Exception:
    _HIP_MODULE = None  # falls back to aiter Triton quant inside custom_kernel

# Buffer cache: (m, n, k) -> (x_fp4_raw, scale_sh_raw, x_fp4_view, scale_sh_view)
# Pre-computed views share storage with raw tensors – no extra copy needed.
_BUFFER_CACHE: dict = {}


def _alloc_and_cache(m: int, n: int, k: int, device):
    import torch
    groups    = k // 32
    padded_m  = (m + 255) // 256 * 256
    sn8       = (groups + 7) // 8 * 8
    fp4_raw   = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
    scale_raw = torch.zeros((padded_m, sn8), dtype=torch.uint8, device=device)
    fp4_v     = fp4_raw.view(_DTYPES.fp4x2)
    scale_v   = scale_raw.view(_DTYPES.fp8_e8m0)
    entry = (fp4_raw, scale_raw, fp4_v, scale_v)
    _BUFFER_CACHE[(m, n, k)] = entry
    return entry


# ---------------------------------------------------------------------------
# Hot path – as flat as possible
# ---------------------------------------------------------------------------

def custom_kernel(data: input_t) -> output_t:
    A, B, _B_q, B_shuffle, B_scale_sh = data

    m = A.shape[0]
    k = A.shape[1]
    n = B.shape[0]

    if not A.is_contiguous():
        A = A.contiguous()

    if _HIP_MODULE is not None:
        # Fast path: zero Python allocation, pre-computed views
        entry = _BUFFER_CACHE.get((m, n, k))
        if entry is None:
            entry = _alloc_and_cache(m, n, k, A.device)
        fp4_raw, scale_raw, fp4_v, scale_v = entry
        _HIP_MODULE.quant_mxfp4_hip(A, fp4_raw, scale_raw)
    else:
        # Fallback: aiter Triton quant (correct for any K)
        fp4_raw, scale_raw = _dynamic_mxfp4_quant(A)
        scale_raw = _e8m0_shuffle(scale_raw)
        fp4_v   = fp4_raw.view(_DTYPES.fp4x2)
        scale_v = scale_raw.view(_DTYPES.fp8_e8m0)

    return _GEMM_FN(
        fp4_v,
        B_shuffle,
        scale_v,
        B_scale_sh,
        dtype=_DTYPES.bf16,
        bpreshuffle=True,
    )
scrolls · 265 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