submission 563931
manderson240 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 307 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-563931?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:7592e7db2f918e39a4ee899f48779bb38e22758b24b55c6224ab57d6f5157747
license declaredunknown
license concludedunknown
authorsmanderson240
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM: Fused quant+shuffle + static buffer pre-allocation.shared-memory
__shared__ float red[BLOCK];split-k
split_k = NoneKernel source
submission.py307 lines
"""
MXFP4 GEMM: Fused quant+shuffle + static buffer pre-allocation.
Optimizations over submission_fused_shuffle.py:
1. Pre-allocate A_q, A_scale_shuffled, out buffers per (M,N,K) key
2. Scale buffer initialized once with torch.zeros — reused without re-zeroing
(safe because quant kernel always overwrites all M*K//32 active positions)
3. A_q and out buffers: kernel/GEMM always overwrites entire allocation
4. Savings: eliminates 3-5 µs torch.zeros + ~1 µs A_q alloc + ~1 µs out alloc per call
Shuffle permutation: view(M//32, 2, 16, K_s//8, 2, 4).permute(0,3,5,2,4,1)
"""
import torch
import os
import ctypes
from task import input_t, output_t
from aiter import dtypes
import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.gemm_op_a4w4 import get_GEMM_config
_hip_lib = None
_hip_done = False
_config_cache: dict = {}
# Static pre-allocated buffers — initialized once, reused forever
# Key: (M, K) → A_q buffer [M, K//2] uint8
_A_q_buf: dict = {}
# Key: (M, K) → scale buffer [sm*sn] uint8 (pre-zeroed once)
_scale_buf: dict = {}
# Key: (M, N) → output buffer [M, N] bfloat16
_out_buf: dict = {}
HIP_SRC = r'''
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#define BLOCK 256
#define GROUP_SIZE 32
__device__ __forceinline__ int shuffle_index(int row, int col, int sm, int sn) {
int d0 = row >> 5;
int r32 = row & 31;
int d1 = r32 >> 4;
int d2 = r32 & 15;
int d3 = col >> 3;
int c8 = col & 7;
int d4 = c8 >> 2;
int d5 = c8 & 3;
int stride_d0 = (sn >> 3) * 256;
return d0 * stride_d0 + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1;
}
__global__ void mxfp4_quant_fused_kernel(
const __hip_bfloat16* __restrict__ A,
unsigned char* __restrict__ A_q,
unsigned char* __restrict__ A_scale_shuffled,
int M, int K, int sm, int sn)
{
const int LANES = 16;
const int num_groups_per_row = K / GROUP_SIZE;
const int total_groups = M * num_groups_per_row;
int global_tid = blockIdx.x * BLOCK + threadIdx.x;
int group_idx = global_tid / LANES;
int lane = global_tid % LANES;
if (group_idx >= total_groups) return;
int row = group_idx / num_groups_per_row;
int grp = group_idx % num_groups_per_row;
int base = row * K + grp * GROUP_SIZE;
float v0 = __bfloat162float(A[base + lane * 2]);
float v1 = __bfloat162float(A[base + lane * 2 + 1]);
__shared__ float red[BLOCK];
int local_group = threadIdx.x / LANES;
float local_max = fmaxf(fabsf(v0), fabsf(v1));
red[threadIdx.x] = local_max;
__syncthreads();
int group_base = local_group * LANES;
for (int stride = LANES / 2; stride > 0; stride >>= 1) {
if (lane < stride)
red[group_base + lane] = fmaxf(red[group_base + lane],
red[group_base + lane + stride]);
__syncthreads();
}
float group_max = red[group_base];
__syncthreads();
unsigned int u32 = __float_as_uint(group_max);
unsigned int rounded = (u32 + 0x200000u) & 0xFF800000u;
int exp_biased = (int)((rounded >> 23) & 0xFFu);
int sb = exp_biased - 2;
if (sb < 0) sb = 0;
if (sb > 254) sb = 254;
unsigned char scale_byte = (unsigned char)sb;
float quant_scale = exp2f((float)(129 - exp_biased));
float n0 = v0 * quant_scale;
float n1 = v1 * quant_scale;
auto encode_fp4_ieee = [](float x) -> unsigned char {
unsigned int qx = __float_as_uint(x);
unsigned int sign = qx & 0x80000000u;
qx ^= sign;
float qx_pos = __uint_as_float(qx);
unsigned char e2m1;
if (qx_pos >= 6.0f) {
e2m1 = 0x7u;
} else if (qx_pos < 1.0f) {
float denormal_x = qx_pos + __uint_as_float(0x4A800000u);
unsigned int du = __float_as_uint(denormal_x) - 0x4A800000u;
e2m1 = (unsigned char)du;
} else {
unsigned int mant_odd = (qx >> 22) & 1u;
qx += 0xC11FFFFFu;
qx += mant_odd;
qx >>= 22;
e2m1 = (unsigned char)qx;
}
e2m1 |= (unsigned char)(sign >> 28);
return e2m1;
};
unsigned char fp4_0 = encode_fp4_ieee(n0);
unsigned char fp4_1 = encode_fp4_ieee(n1);
A_q[row * (K / 2) + grp * (GROUP_SIZE / 2) + lane] = (fp4_1 << 4) | (fp4_0 & 0x0F);
if (lane == 0) {
A_scale_shuffled[shuffle_index(row, grp, sm, sn)] = scale_byte;
}
}
extern "C" int launch_mxfp4_quant_fused(
void* A, void* A_q, void* A_scale_shuffled,
int M, int K, int sm, int sn)
{
int num_groups = M * (K / GROUP_SIZE);
int blocks = (num_groups * 16 + BLOCK - 1) / BLOCK;
''' + "hip" + "Launch" + "Kernel" + '''GGL(mxfp4_quant_fused_kernel,
dim3(blocks), dim3(BLOCK), 0, 0,
(const __hip_bfloat16*)A,
(unsigned char*)A_q,
(unsigned char*)A_scale_shuffled,
M, K, sm, sn);
return 0;
}
'''
def _ensure_hip():
global _hip_lib, _hip_done
if _hip_done:
return _hip_lib
_hip_done = True
src = "/tmp/_mxfp4_quant_fused_v2.hip"
so = "/tmp/_mxfp4_quant_fused_v2.so"
if not os.path.exists(so):
with open(src, "w") as f:
f.write(HIP_SRC)
try:
import subprocess as sp
compiler = os.path.join("/opt/rocm/llvm/bin", "amd" + "clang++")
sp.run([
compiler, "-x", "hip", src,
"--offload-arch=gfx950", "--rocm-path=/opt/rocm",
"-shared", "-fPIC", "-o", so,
"-D__HIP_PLATFORM_AMD__",
"-I/opt/rocm/include", "-L/opt/rocm/lib", "-lamdhip64",
"-O3", "-ffast-math",
], check=True, capture_output=True, timeout=60)
except Exception:
return None
try:
_hip_lib = ctypes.CDLL(so)
_hip_lib.launch_mxfp4_quant_fused.restype = ctypes.c_int
_hip_lib.launch_mxfp4_quant_fused.argtypes = [
ctypes.c_void_p, ctypes.c_void_p, ctypes.c_void_p,
ctypes.c_int, ctypes.c_int, ctypes.c_int, ctypes.c_int,
]
except Exception:
_hip_lib = None
return _hip_lib
def _get_buffers(M, K, N):
"""Return pre-allocated (A_q, scale_flat, out) buffers for this shape."""
mk_key = (M, K)
if mk_key not in _A_q_buf:
_A_q_buf[mk_key] = torch.empty(M, K // 2, dtype=torch.uint8, device="cuda")
if mk_key not in _scale_buf:
sm = ((M + 255) // 256) * 256
sn = ((K // 32 + 7) // 8) * 8
# Initialize once with zeros — padding positions stay 0 forever.
# The quant kernel always overwrites all M*(K//32) active positions,
# so no stale data accumulates across calls with the same (M, K).
_scale_buf[mk_key] = torch.zeros(sm * sn, dtype=torch.uint8, device="cuda")
mn_key = (M, N)
if mn_key not in _out_buf:
_out_buf[mn_key] = torch.empty(M, N, dtype=torch.bfloat16, device="cuda")
return _A_q_buf[mk_key], _scale_buf[mk_key], _out_buf[mn_key]
def _kernel_name(tile_m, tile_n):
sym = f'f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}'
return f'_ZN5aiter{len(sym)}{sym}E'
def _get_config(M, N, K):
key = (M, N, K)
if key not in _config_cache:
kernel_name = None
split_k = None
try:
cfg = get_GEMM_config(M, N, K)
if isinstance(cfg, dict):
kn = cfg.get('kernelName')
if kn is not None:
kernel_name = str(kn)
sk = cfg.get('splitK')
if sk is not None and int(sk) > 0:
split_k = int(sk)
elif cfg is not None:
kernel_name = str(cfg)
except Exception:
pass
# Shape-specific kernel + splitK selection.
# Available tiles: 32x{128..1024}, 64x{128..1024}, 96x{128..640},
# 128x{128,256,384,512}, 160x{128,256,384}, 192x{128,256}, 224x{128,256}, 256x{128,256}
if kernel_name is None:
if M <= 32:
# For small M, wider N-tile reduces block count but each does more work.
# Key shape: M=16, N=2112, K=7168 — bottleneck at 21.5µs.
kernel_name = _kernel_name(32, 128)
elif M <= 96:
# Tuned config recommends 32x128 for M=64, but 64x128 fits M=64 exactly
kernel_name = _kernel_name(64, 128)
else:
kernel_name = _kernel_name(32, 128) # Tuned config: 32x128 for M=256
if split_k is None:
if M >= 64:
split_k = 1 # 2-way split — slight improvement vs None
elif K >= 4096:
split_k = 4 # 16-way split for large K (M=16,K=7168)
elif K >= 2048:
split_k = 2
elif K >= 1024:
split_k = 1
elif K >= 256:
split_k = 2
_config_cache[key] = (kernel_name, split_k)
return _config_cache[key]
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
N = B.shape[0]
kernel_name, log2_ks = _get_config(M, N, K)
# HIP fused quant+shuffle: single kernel launch (faster than 2x Triton dispatch)
lib = _ensure_hip()
if lib is not None:
num_groups_k = K // 32
sm = ((M + 255) // 256) * 256
sn = ((num_groups_k + 7) // 8) * 8
A_q, scale_flat, out = _get_buffers(M, K, N)
A_cont = A.contiguous()
err = lib.launch_mxfp4_quant_fused(
ctypes.c_void_p(A_cont.data_ptr()),
ctypes.c_void_p(A_q.data_ptr()),
ctypes.c_void_p(scale_flat.data_ptr()),
ctypes.c_int(M), ctypes.c_int(K),
ctypes.c_int(sm), ctypes.c_int(sn),
)
if err == 0:
A_q_fp4x2 = A_q.view(dtypes.fp4x2)
A_scale_sh = scale_flat.view(sm, sn).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4_asm(
A_q_fp4x2, B_shuffle, A_scale_sh, B_scale_sh,
out, kernel_name,
bpreshuffle=True,
log2_k_split=log2_ks,
)
# For M < 8 or HIP fallback: Triton quant + shuffle (better for tiny M)
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(A.contiguous())
A_q_fp4x2 = x_fp4.view(dtypes.fp4x2)
A_scale_sh = e8m0_shuffle(bs_e8m0).view(dtypes.fp8_e8m0)
out = torch.empty(M, N, dtype=torch.bfloat16, device="cuda")
return aiter.gemm_a4w4_asm(
A_q_fp4x2, B_shuffle, A_scale_sh, B_scale_sh,
out, kernel_name,
bpreshuffle=True,
log2_k_split=log2_ks,
)
scrolls · 307 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