submission 595946
Borui Xu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 205 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-595946?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:046eea6b711d42fe1c07c7f986c9733dbe4361367ea847ce31192ad9bc9ddcc9
license declaredunknown
license concludedunknown
authorsBorui Xu
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission.py205 lines
"""
MXFP4 GEMM — Virtual M-padding for better GEMM occupancy.
For small M with large K (low block count), pad M to increase blocks.
The quant kernel handles virtual padding (zero rows beyond actual M).
GEMM gets more blocks → better memory latency hiding.
"""
from task import input_t, output_t
import torch
import aiter
from aiter import dtypes
_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_e8m0 = dtypes.fp8_e8m0
_K32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
# Config patch
try:
from aiter.ops import gemm_op_a4w4 as _gmod
_orig_cfg = _gmod.get_GEMM_config
def _p(M, N, K):
try:
c = _orig_cfg(M, N, K)
if c is not None: return c
except: pass
return {'kernelId': 21, 'splitK': 0, 'us': 10.0,
'kernelName': _K32, 'tflops': 0, 'bw': 0, 'errRatio': 0}
_gmod.get_GEMM_config = _p
except: pass
_hip = None
_asm_fn = None
_init_done = False
_cache = {}
def _build_hip():
global _hip
if _hip is not None: return _hip
from torch.utils.cpp_extension import load_inline
cuda_src = r'''
#include <torch/extension.h>
#include <math.h>
__device__ __forceinline__ unsigned f2u(float f) {
union { float ff; unsigned uu; } c; c.ff = f; return c.uu;
}
__device__ __forceinline__ float u2f(unsigned u) {
union { unsigned uu; float ff; } c; c.uu = u; return c.ff;
}
__device__ __forceinline__ uint8_t quant_fp4(float v) {
uint8_t s = (v < 0.f) ? 8u : 0u;
float a = fabsf(v);
uint8_t c;
if (a > 5.0f) c = 7;
else if (a == 5.0f) c = 6;
else if (a > 3.5f) c = 6;
else if (a == 3.5f) c = 6;
else if (a > 2.5f) c = 5;
else if (a == 2.5f) c = 4;
else if (a > 1.75f) c = 4;
else if (a == 1.75f) c = 4;
else if (a > 1.25f) c = 3;
else if (a == 1.25f) c = 2;
else if (a > 0.75f) c = 2;
else if (a == 0.75f) c = 2;
else if (a > 0.25f) c = 1;
else if (a == 0.25f) c = 0;
else c = 0;
return s | c;
}
// Quant kernel with virtual M padding
// Processes virtual_M rows, but only reads A for rows < actual_M (rest = 0)
__global__ void mxfp4_quant_shuffle_kernel(
const uint16_t* __restrict__ A,
uint8_t* __restrict__ A_fp4,
uint8_t* __restrict__ A_scale_sh,
int actual_M, int virtual_M, int K, int ngroups, int sn_pad) {
int gtid = blockIdx.x * blockDim.x + threadIdx.x;
int hw = gtid >> 5, lane = gtid & 31;
if (hw >= virtual_M * ngroups) return;
int row = hw / ngroups, gcol = hw % ngroups;
// Virtual padding: zero for rows beyond actual_M
float val = 0.f;
if (row < actual_M)
val = u2f((unsigned)A[row * K + gcol * 32 + lane] << 16);
float mx = fabsf(val);
mx = fmaxf(mx, __shfl_xor(mx, 16));
mx = fmaxf(mx, __shfl_xor(mx, 8));
mx = fmaxf(mx, __shfl_xor(mx, 4));
mx = fmaxf(mx, __shfl_xor(mx, 2));
mx = fmaxf(mx, __shfl_xor(mx, 1));
uint8_t e8, fp4;
if (mx == 0.f) { e8 = 0; fp4 = 0; }
else {
unsigned rounded = (f2u(mx) + 0x200000u) & 0xFF800000u;
float su = floorf(log2f(u2f(rounded))) - 2.0f;
su = fminf(fmaxf(su, -127.f), 127.f);
e8 = (uint8_t)((int)su + 127);
fp4 = quant_fp4(val * exp2f(-su));
}
uint8_t p = (uint8_t)__shfl_xor((int)fp4, 1);
if ((lane & 1) == 0)
A_fp4[row * (K >> 1) + (gcol << 4) + (lane >> 1)] =
(fp4 & 0xFu) | ((p & 0xFu) << 4);
if (lane == 0) {
int d0=row>>5, d1=(row>>4)&1, d2=row&15, d3=gcol>>3, d4=(gcol>>2)&1, d5=gcol&3;
A_scale_sh[d0*((sn_pad>>3)*256) + d3*256 + d5*64 + d2*4 + d4*2 + d1] = e8;
}
}
void mxfp4_quant_vpad(torch::Tensor A, torch::Tensor fp4, torch::Tensor scale_sh,
int actual_M, int virtual_M) {
int K = A.size(1), ngrp = K / 32;
int sn_pad = scale_sh.size(1);
int tot = virtual_M * ngrp * 32, BS = 256;
mxfp4_quant_shuffle_kernel<<<(tot+BS-1)/BS, BS>>>(
(const uint16_t*)A.data_ptr(),
fp4.data_ptr<uint8_t>(), scale_sh.data_ptr<uint8_t>(),
actual_M, virtual_M, K, ngrp, sn_pad);
}
'''
cpp_src = "void mxfp4_quant_vpad(torch::Tensor, torch::Tensor, torch::Tensor, int, int);\n"
_hip = load_inline(name="hip_vpad", cpp_sources=cpp_src,
cuda_sources=cuda_src, functions=["mxfp4_quant_vpad"],
extra_cuda_cflags=["-O3"], verbose=False)
return _hip
def _init():
global _asm_fn, _init_done
if _init_done: return
_init_done = True
try:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
_asm_fn = gemm_a4w4_asm
except: pass
def _compute_virtual_m(m, n, k):
"""Compute virtual M that gives good GEMM occupancy."""
n_tiles = (n + 127) // 128
# Current blocks = ceil(m/32) * n_tiles
cur_blocks = ((m + 31) // 32) * n_tiles
# Target: at least 64 blocks for decent occupancy
if cur_blocks >= 64:
return m
# Pad M to get more M-tiles (each adds n_tiles blocks)
target_m_tiles = max(2, (64 + n_tiles - 1) // n_tiles)
virtual_m = target_m_tiles * 32
return virtual_m
def _get_cache(m, n, k, virtual_m, device):
key = (m, n, k, virtual_m)
if key not in _cache:
ngrp = k // 32
sm_pad = ((virtual_m + 31) // 32) * 32
sn_pad = ((ngrp + 7) // 8) * 8
_cache[key] = {
'fp4': torch.empty((virtual_m, k // 2), dtype=torch.uint8, device=device),
'scale': torch.zeros((sm_pad, sn_pad), dtype=torch.uint8, device=device),
'out': torch.empty((virtual_m, n), dtype=torch.bfloat16, device=device),
}
return _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]
if not _init_done:
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
aq, asc = dynamic_mxfp4_quant(A)
ash = e8m0_shuffle(asc)
out = aiter.gemm_a4w4(aq.view(_fp4x2), B_shuffle,
ash.view(_e8m0), B_scale_sh,
dtype=_bf16, bpreshuffle=True)
_init()
return out
hip = _build_hip()
virtual_m = _compute_virtual_m(m, n, k)
c = _get_cache(m, n, k, virtual_m, A.device)
# Quant with virtual padding (no extra copies!)
hip.mxfp4_quant_vpad(A, c['fp4'], c['scale'], m, virtual_m)
# GEMM on virtual_m rows (more blocks → better occupancy)
if _asm_fn is not None:
_asm_fn(c['fp4'].view(_fp4x2), B_shuffle,
c['scale'].view(_e8m0), B_scale_sh,
c['out'], _K32, bpreshuffle=True)
else:
c['out'][:] = aiter.gemm_a4w4(
c['fp4'].view(_fp4x2), B_shuffle,
c['scale'].view(_e8m0), B_scale_sh,
dtype=_bf16, bpreshuffle=True)
return c['out'][:m, :n]
scrolls · 205 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