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
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