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
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.
fp4
TORCH_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