submission 609571
nicholaswilde_08140 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 117 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-609571?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:f291514ff502140ff95ab8c5e563aff47410b241d0a1eee157f5850e9190962f
license declaredunknown
license concludedunknown
authorsnicholaswilde_08140
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
int fp4;Kernel source
submission.py117 lines
import torch
from task import input_t, output_t
CPP_WRAPPER = """
void quant_bf16_mxfp4(torch::Tensor x, torch::Tensor x_q, torch::Tensor x_s, int m, int n);
"""
CUDA_SRC = """
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
__global__ void kernel_quant_bf16_mxfp4(__hip_bfloat16 *x, char *x_q, char *x_s, int m, int n) {
if (blockIdx.x * 512 + threadIdx.x * 8 >= m * n) return;
__hip_bfloat16 elements[8];
*(float4 *)elements = *((float4*)(x + blockIdx.x * 512 + threadIdx.x * 8));
float absmax = 0.0f;
for (int i = 0; i < 8; i++) {
float element = __bfloat162float(elements[i]);
absmax = fmaxf(absmax, fabsf(element));
}
for (int off = 1; off < 4; off <<= 1) {
absmax = fmaxf(absmax, __shfl_down(absmax, off));
}
absmax = __shfl(absmax, (threadIdx.x / 4) * 4);
float scale;
unsigned int amax_bits = (__float_as_uint(absmax) + 0x200000u) & 0x7F800000u;
scale = __uint_as_float(amax_bits) / 4.0f;
int exp = (amax_bits >> 23) - 2;
if (threadIdx.x % 4 == 0) {
// x_s[blockIdx.x * 16 + threadIdx.x / 4] = (char)exp;
int this_m = blockIdx.x * 512 / n;
int this_n = (blockIdx.x * 512 + threadIdx.x * 8) % n / 32;
x_s[this_m / 32 * 32 * (n / 32) + this_n / 8 * 256 + (this_n % 8) % 4 * 64 + (this_m % 32) % 16 * 4 + (this_n % 8) / 4 * 2 + (this_m % 32) / 16] = (char)exp;
}
int fp4x8 = 0;
for (int i = 0; i < 8; i++) {
float element = __bfloat162float(elements[i]);
int sign = (element >= 0) ? 0x0 : 0x8;
float q = element / scale;
int fp4;
if (fabs(q) > 5.0) fp4 = sign | 0x7;
else if (fabs(q) >= 3.5) fp4 = sign | 0x6;
else if (fabs(q) > 2.5) fp4 = sign | 0x5;
else if (fabs(q) >= 1.75) fp4 = sign | 0x4;
else if (fabs(q) > 1.25) fp4 = sign | 0x3;
else if (fabs(q) >= 0.75) fp4 = sign | 0x2;
else if (fabs(q) > 0.25) fp4 = sign | 0x1;
else fp4 = 0x0;
fp4x8 |= (fp4 << (i * 4));
}
*((int *)(x_q + blockIdx.x * 256 + threadIdx.x * 4)) = fp4x8;
}
void quant_bf16_mxfp4(torch::Tensor x, torch::Tensor x_q, torch::Tensor x_s, int m, int n) {
__hip_bfloat16 *x_ptr = (__hip_bfloat16 *)x.data_ptr();
char *x_q_ptr = (char *)x_q.data_ptr();
char *x_s_ptr = (char *)x_s.data_ptr();
int grid = (m * n + 511) / 512;
hipLaunchKernelGGL(kernel_quant_bf16_mxfp4, dim3(grid), dim3(64), 0, 0, x_ptr, x_q_ptr, x_s_ptr, m, n);
}
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ["CXX"] = "clang++"
from torch.utils.cpp_extension import load_inline
module = load_inline(
name='mxfp4_mm',
cpp_sources=[CPP_WRAPPER],
cuda_sources=[CUDA_SRC],
functions=['quant_bf16_mxfp4'],
verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
)
def custom_kernel(data: input_t) -> output_t:
import aiter
from aiter import QuantType, dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
def _quant_mxfp4(x, shuffle=True):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
if shuffle:
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
B = B.contiguous()
m, k = A.shape
n, _ = B.shape
# A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
A_q = torch.empty((m, k // 2), dtype=dtypes.fp4x2, device=A.device)
A_scale_sh = torch.empty((256, k // 32), dtype=dtypes.fp8_e8m0, device=A.device)
module.quant_bf16_mxfp4(A, A_q, A_scale_sh, m, k)
out_gemm = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out_gemmscrolls · 117 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