submission 675976
phoenixdna · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 396 lines, June 9 Researcher Reciprocity License v1.0.
submission_v41_hybrid_direct_bshuf.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-675976?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:579bb3e619433df4d5287589c56c7af561d59e947897440ddae0d4fe7a172136
license declaredunknown
license concludedunknown
authorsphoenixdna
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float s_abs[32 * WARP_SLOTS];split-k
f.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")Kernel source
submission_v41_hybrid_direct_bshuf.py396 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, 16, 2112, 7168, 21, 1, 12.90, "_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"),
]
_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_F32 = 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_F32 - MBITS_FP4) + 1) << MBITS_F32);
constexpr int32_t VAL_TO_ADD =
((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + 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_F32 - 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_F32 - MBITS_FP4));
return static_cast<uint8_t>(sign_lp | (e2m1 & MAX_INT));
}
__device__ __forceinline__ int scale_shuffle_offset(int row, int col, int scale_n_pad) {
int offs_0 = row / 32;
int row_in_32 = row % 32;
int offs_1 = row_in_32 / 16;
int offs_2 = row_in_32 % 16;
int offs_3 = col / 8;
int col_in_8 = col % 8;
int offs_4 = col_in_8 / 4;
int offs_5 = col_in_8 % 4;
return offs_1 + offs_4 * 2 + offs_2 * 4 + offs_5 * 64 + offs_3 * 256 + offs_0 * 32 * scale_n_pad;
}
template <int ROWS_PER_BLOCK, int GROUPS_PER_BLOCK>
__global__ void quant_mxfp4_a_kernel(
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
) {
constexpr int WARP_SLOTS = ROWS_PER_BLOCK * GROUPS_PER_BLOCK;
__shared__ float s_abs[32 * WARP_SLOTS];
__shared__ uint8_t s_scale[WARP_SLOTS];
constexpr int GROUP_SIZE = 32;
constexpr int PAIRS_PER_GROUP = GROUP_SIZE / 2;
int tid = static_cast<int>(threadIdx.x);
int lane = tid & 31;
int warp = tid >> 5;
int row_in_block = warp / GROUPS_PER_BLOCK;
int group_in_block = warp % GROUPS_PER_BLOCK;
int warp_slot = row_in_block * GROUPS_PER_BLOCK + group_in_block;
int row = static_cast<int>(blockIdx.y) * ROWS_PER_BLOCK + row_in_block;
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]);
}
int group_base = warp_slot * GROUP_SIZE;
s_abs[group_base + lane] = active_group ? fabsf(value) : 0.0f;
__syncthreads();
if (lane < 16) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 16];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane < 8) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 8];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane < 4) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 4];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane < 2) {
float a = s_abs[group_base + lane];
float b = s_abs[group_base + lane + 2];
s_abs[group_base + lane] = a > b ? a : b;
}
__syncthreads();
if (lane == 0) {
float a = s_abs[group_base];
float b = s_abs[group_base + 1];
float max_abs = a > b ? a : b;
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[warp_slot] = scale_byte;
if (active_group) {
int scale_idx = scale_shuffle_offset(row, scale_col, scale_n_pad);
scale_sh[scale_idx] = scale_byte;
}
}
__syncthreads();
if (lane < PAIRS_PER_GROUP && active_group) {
uint8_t scale_byte = s_scale[warp_slot];
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 pair_base = row * k + scale_col * GROUP_SIZE + lane * 2;
float x0 = bf16_to_f32(x[pair_base]) * inv_scale;
float x1 = bf16_to_f32(x[pair_base + 1]) * inv_scale;
uint8_t q0 = quantize_e2m1(x0);
uint8_t q1 = quantize_e2m1(x1);
int out_idx = row * (k / 2) + 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 / 32;
int scale_n_pad = static_cast<int>(scale_sh.size(1));
TORCH_CHECK(k % 64 == 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) / 2, "x_fp4 col mismatch");
if (k >= 1536) {
dim3 grid((scale_n_valid + 1) / 2, (m + 3) / 4);
dim3 block(256);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(quant_mxfp4_a_kernel<4, 2>),
grid,
block,
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) / 2, m);
dim3 block(64);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(quant_mxfp4_a_kernel<1, 2>),
grid,
block,
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_v41_hybrid_direct_bshuf",
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.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 _prepare_triton_weight_from_task(B_shuffle: torch.Tensor, B_scale_sh: torch.Tensor):
if B_shuffle.dtype != torch.uint8:
b_shuf_u8 = B_shuffle.view(torch.uint8)
else:
b_shuf_u8 = B_shuffle
if B_scale_sh.dtype != torch.uint8:
b_scale_u8 = B_scale_sh.view(torch.uint8)
else:
b_scale_u8 = B_scale_sh
n = b_shuf_u8.shape[0]
k_pack = b_shuf_u8.shape[1]
k_scale = k_pack // 16
# The task already provides B_shuffle = shuffle_weight(B_q, (16, 16)).
# Triton consumes the same bytes but with a different 2D view.
w_triton = b_shuf_u8.contiguous().view(n // 16, k_pack * 16)
# fp4_utils.e8m0_shuffle() and Triton's shuffle_scales() use the same
# underlying byte ordering; only the final 2D view differs.
w_scales_triton = b_scale_u8[:n, :k_scale].contiguous().view(n // 32, k_scale * 32)
return w_triton, w_scales_triton
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()
m, k = A.shape
n = B_shuffle.shape[0]
triton_shapes = {
(4, 2880, 512),
(16, 2112, 7168),
(32, 4096, 512),
(32, 2880, 512),
}
if (m, n, k) in triton_shapes:
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
w_triton, w_scales_triton = _prepare_triton_weight_from_task(
B_shuffle.contiguous(), B_scale_sh.contiguous()
)
return gemm_a16wfp4_preshuffle(
A,
w_triton,
w_scales_triton,
prequant=True,
dtype=torch.bfloat16,
)
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 · 396 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