submission 720763
mingkai_37292 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 202 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-720763?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:56b4697fcbdad7c711fd3942a97698a83fa2e19dacf49d3ceca40d89f14e6d55
license declaredunknown
license concludedunknown
authorsmingkai_37292
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
the e8m0_shuffle step. BLOCK_M is autotuned per (M, K) shape to minimize wasted workfp4
FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.num-warps = 1
triton.Config({"BLOCK_M": 4}, num_warps=1),tile-m = 4
BLOCK_M: tl.constexpr = 4,Kernel source
submission.py202 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quantization is implemented as a custom Triton kernel.
Scales are written directly into the shuffled layout expected by gemm_a4w4, fusing
the e8m0_shuffle step. BLOCK_M is autotuned per (M, K) shape to minimize wasted work
on small-M shapes and maximize occupancy on larger ones.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.autotune(
configs=[
triton.Config({"BLOCK_M": 4}, num_warps=1),
triton.Config({"BLOCK_M": 8}, num_warps=2),
triton.Config({"BLOCK_M": 16}, num_warps=4),
triton.Config({"BLOCK_M": 32}, num_warps=4),
],
key=["M", "K"],
)
@triton.jit
def _mxfp4_quant_kernel(
x_ptr,
out_ptr,
scale_sh_ptr, # pre-shuffled scale buffer: flat [M_padded * (K//32)]
M, K,
stride_xm,
GROUP_SIZE: tl.constexpr = 32,
BLOCK_M: tl.constexpr = 4,
):
row_start = tl.program_id(0) * BLOCK_M
group_id = tl.program_id(1)
k_start = group_id * GROUP_SIZE
# 1D row offsets and 2D tensors
row_offs_1d = row_start + tl.arange(0, BLOCK_M) # [BLOCK_M]
row_offs = row_offs_1d[:, None] # [BLOCK_M, 1]
half_offs = tl.arange(0, GROUP_SIZE // 2)[None, :] # [1, 16]
k_even = k_start + half_offs * 2 # [1, 16]
k_odd = k_start + half_offs * 2 + 1 # [1, 16]
row_mask_1d = row_offs_1d < M # [BLOCK_M]
row_mask = row_mask_1d[:, None] # [BLOCK_M, 1] — broadcasts over columns
mask_e = row_mask & (k_even < K) # [BLOCK_M, 16]
mask_o = row_mask & (k_odd < K) # [BLOCK_M, 16]
x_even = tl.load(x_ptr + row_offs * stride_xm + k_even, mask=mask_e, other=0.0).to(tl.float32)
x_odd = tl.load(x_ptr + row_offs * stride_xm + k_odd, mask=mask_o, other=0.0).to(tl.float32)
# E8M0 scale per row: reduce over 16 cols -> [BLOCK_M]
abs_max_1d = tl.maximum(tl.max(tl.abs(x_even), axis=1),
tl.max(tl.abs(x_odd), axis=1)) # [BLOCK_M]
# Match reference _mxfp4_quant_op scale exactly via bitwise rounding
abs_max_1d = tl.maximum(abs_max_1d, 1e-38).to(tl.float32)
abs_max_int = abs_max_1d.to(tl.int32, bitcast=True)
abs_max_rounded = ((abs_max_int + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000).to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.floor(tl.math.log2(abs_max_rounded)).to(tl.int32) - 2
scale_e8m0_unbiased = tl.minimum(tl.maximum(scale_e8m0_unbiased, -127), 127)
e8m0_exp = (scale_e8m0_unbiased + 127).to(tl.uint8) # [BLOCK_M]
quant_scale = tl.math.exp2(-scale_e8m0_unbiased.to(tl.float32))[:, None] # [BLOCK_M, 1]
# Store scales directly into shuffled layout (fusing e8m0_shuffle).
sn = K // GROUP_SIZE
j = group_id
sh_off = (
(row_offs_1d // 32) * (32 * sn)
+ (j // 8) * 256
+ (j % 4) * 64
+ (row_offs_1d % 16) * 4
+ ((j // 4) % 2) * 2
+ (row_offs_1d // 16) % 2
)
tl.store(scale_sh_ptr + sh_off, e8m0_exp, mask=row_mask_1d)
# E2M1 bit-manipulation encoding matching ROCm/aiter _mxfp4_quant_op exactly
# Even elements -> lo nibbles
xs_e = x_even * quant_scale
xs_e_uint = xs_e.to(tl.int32, bitcast=True).to(tl.uint32)
s_e = xs_e_uint & 0x80000000
xs_e_pos_uint = xs_e_uint ^ s_e
xs_e_pos = xs_e_pos_uint.to(tl.float32, bitcast=True)
sat_e = xs_e_pos >= 6.0
den_e = xs_e_pos < 1.0
mant_odd_e = (xs_e_pos_uint >> 22) & 1
norm_e = ((xs_e_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_e.to(tl.int32)) >> 22
norm_e = norm_e.to(tl.uint8)
den_val_e = (xs_e_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
den_val_e = den_val_e.to(tl.uint8)
q_e = tl.full(xs_e.shape, 7, dtype=tl.uint8)
q_e = tl.where(~sat_e, norm_e, q_e)
q_e = tl.where(den_e, den_val_e, q_e)
sign_e_lp = (s_e >> 28).to(tl.uint8)
lo = (q_e | sign_e_lp) & 0xF
# Odd elements -> hi nibbles
xs_o = x_odd * quant_scale
xs_o_uint = xs_o.to(tl.int32, bitcast=True).to(tl.uint32)
s_o = xs_o_uint & 0x80000000
xs_o_pos_uint = xs_o_uint ^ s_o
xs_o_pos = xs_o_pos_uint.to(tl.float32, bitcast=True)
sat_o = xs_o_pos >= 6.0
den_o = xs_o_pos < 1.0
mant_odd_o = (xs_o_pos_uint >> 22) & 1
norm_o = ((xs_o_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_o.to(tl.int32)) >> 22
norm_o = norm_o.to(tl.uint8)
den_val_o = (xs_o_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
den_val_o = den_val_o.to(tl.uint8)
q_o = tl.full(xs_o.shape, 7, dtype=tl.uint8)
q_o = tl.where(~sat_o, norm_o, q_o)
q_o = tl.where(den_o, den_val_o, q_o)
sign_o_lp = (s_o >> 28).to(tl.uint8)
hi = ((q_o | sign_o_lp) & 0xF) << 4
# Pack two FP4 nibbles per byte: lo nibble = even index, hi nibble = odd index
packed = lo | hi # [BLOCK_M, 16]
# Store packed output: [BLOCK_M, 16] at row_offs*(K//2) + k_start//2 + half_offs
out_base = row_offs * (K // 2) + k_start // 2 + half_offs # [BLOCK_M, 16]
tl.store(out_ptr + out_base, packed.to(tl.uint8), mask=mask_e)
# Module-level workspace cache
_quant_ws: dict = {}
_gemm_ws: dict = {}
def _triton_mxfp4_quant(x: torch.Tensor):
"""x: [M, K] bf16 -> (fp4_packed [M, K//2] uint8, scale_sh [M_padded, K//32] uint8 in shuffled layout)"""
M, K = x.shape
assert K % 32 == 0, "K must be a multiple of 32 for per-1x32 MXFP4 quant"
x = x.contiguous()
key = (M, K)
if key not in _quant_ws:
M_padded = (M + 255) // 256 * 256
_quant_ws[key] = (
torch.empty(M, K // 2, dtype=torch.uint8, device=x.device),
torch.empty(M_padded * (K // 32), dtype=torch.uint8, device=x.device),
)
out, scale_sh_flat = _quant_ws[key]
M_padded = (M + 255) // 256 * 256
def grid(meta):
return ((M + meta["BLOCK_M"] - 1) // meta["BLOCK_M"], K // 32)
_mxfp4_quant_kernel[grid](
x, out, scale_sh_flat,
M, K,
x.stride(0),
)
return out, scale_sh_flat.view(M_padded, K // 32)
def custom_kernel(data: input_t) -> output_t:
import aiter
from aiter import dtypes
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
M, K = A.shape
# Quantize A; scale is returned already in the shuffled layout
A_fp4, A_scale_sh = _triton_mxfp4_quant(A)
if M <= 64:
N = B.shape[0] # B is [N, K] weight matrix
gemm_key = (M, N, K)
if gemm_key not in _gemm_ws:
_gemm_ws[gemm_key] = torch.empty(M, N, dtype=torch.bfloat16, device=A.device)
out = _gemm_ws[gemm_key]
aiter.gemm_a4w4_asm(
A_fp4.view(dtypes.fp4x2),
B_shuffle,
A_scale_sh.view(dtypes.fp8_e8m0),
B_scale_sh,
out,
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
bpreshuffle=True,
log2_k_split=None,
)
return out
out_gemm = aiter.gemm_a4w4(
A_fp4.view(dtypes.fp4x2),
B_shuffle,
A_scale_sh.view(dtypes.fp8_e8m0),
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out_gemm
scrolls · 202 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