submission 678892
Yaowei Lyu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 197 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-678892?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:a5263229e0ecd73453e441784a2ab56e1067dfb639092cc88dabf7ff7f043e9e
license declaredunknown
license concludedunknown
authorsYaowei Lyu
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 matrix multiplication using Triton kernels on AMD MI355X.Kernel source
submission.py197 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 matrix multiplication using Triton kernels on AMD MI355X.
bf16 A -> MXFP4 per-1x32 quant A -> fp4 GEMM with pre-shuffled B -> bf16 C.
"""
from task import input_t, output_t
def custom_kernel(data: input_t) -> output_t:
import torch
import triton
import triton.language as tl
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n, _ = B.shape
# ----------------------------------------------------------------
# Step 1: dynamic MXFP4 quantization of A (per-1x32 block scaling)
# - For each row, every group of 32 elements shares one e8m0 scale
# - e8m0 scale = exponent of max |x| in the group (biased by 127)
# - FP4 E2M1 values: 0,0.5,1,1.5,2,3,4,6 (with sign)
# ----------------------------------------------------------------
GROUP_SIZE = 32
@triton.jit
def _mxfp4_quant_kernel(
X_ptr,
Out_ptr,
Scale_ptr,
M,
K,
stride_xm,
stride_xk,
stride_om,
stride_ok,
stride_sm,
stride_sk,
BLOCK_M: tl.constexpr,
BLOCK_K: tl.constexpr,
GROUP: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_k = tl.program_id(1)
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
# each pid_k handles one group of GROUP elements
rk = pid_k * GROUP + tl.arange(0, GROUP)
mask = (rm[:, None] < M) & (rk[None, :] < K)
x = tl.load(
X_ptr + rm[:, None] * stride_xm + rk[None, :] * stride_xk,
mask=mask,
other=0.0,
)
# compute per-group max absolute value
ax = tl.abs(x)
amax = tl.max(ax, axis=1) # [BLOCK_M]
# e8m0 biased exponent: floor(log2(amax)) + 127, clamped to [0,254]
# use bitcast to extract exponent from bf16/fp32
amax_f32 = amax.to(tl.float32)
# add small eps to avoid log2(0)
amax_f32 = tl.where(amax_f32 > 0.0, amax_f32, 1.0e-30)
log2_amax = tl.math.log2(amax_f32)
exp_biased = tl.math.floor(log2_amax).to(tl.int32) + 127
exp_biased = tl.maximum(exp_biased, 0)
exp_biased = tl.minimum(exp_biased, 254)
scale_e8m0 = exp_biased.to(tl.uint8)
# reconstruct scale as power of 2: 2^(exp_biased - 127)
scale_f = tl.math.exp2((exp_biased - 127).to(tl.float32))
# normalized = x / scale * 8.0 (fp4 e2m1 range is [0..6], we map to integer codes)
inv_scale = 1.0 / tl.where(scale_f > 0.0, scale_f, 1.0)
xn = x.to(tl.float32) * inv_scale[:, None]
# Round to nearest FP4 E2M1 value
# FP4 E2M1 positive values: 0, 0.5, 1, 1.5, 2, 3, 4, 6
# We use a simple approach: clamp + round
sign = tl.where(xn < 0.0, 1, 0)
xn_abs = tl.abs(xn)
# Map to fp4 code (0-7): 0->0, 1->0.5, 2->1.0, 3->1.5, 4->2.0, 5->3.0, 6->4.0, 7->6.0
# Reverse: find nearest
# Boundaries: 0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0
code = tl.zeros_like(xn_abs).to(tl.int32)
code = tl.where(xn_abs >= 0.25, 1, code)
code = tl.where(xn_abs >= 0.75, 2, code)
code = tl.where(xn_abs >= 1.25, 3, code)
code = tl.where(xn_abs >= 1.75, 4, code)
code = tl.where(xn_abs >= 2.5, 5, code)
code = tl.where(xn_abs >= 3.5, 6, code)
code = tl.where(xn_abs >= 5.0, 7, code)
# FP4 E2M1 encoding: sign(1) | exp(2) | man(1)
# code 0 -> 0b0000, code 1 -> 0b0001, code 2 -> 0b0010, code 3 -> 0b0011
# code 4 -> 0b0100, code 5 -> 0b0101, code 6 -> 0b0110, code 7 -> 0b0111
nibble = (sign.to(tl.int32) << 3) | code # 4-bit value
# Pack pairs of fp4 into uint8: low nibble = even index, high nibble = odd index
# rk has GROUP elements, pack into GROUP//2 bytes
even_idx = tl.arange(0, GROUP // 2) * 2
odd_idx = even_idx + 1
lo = tl.load(
X_ptr
+ rm[:, None] * stride_xm
+ (pid_k * GROUP + even_idx[None, :]) * stride_xk,
mask=(rm[:, None] < M) & ((pid_k * GROUP + even_idx[None, :]) < K),
other=0.0,
)
hi = tl.load(
X_ptr
+ rm[:, None] * stride_xm
+ (pid_k * GROUP + odd_idx[None, :]) * stride_xk,
mask=(rm[:, None] < M) & ((pid_k * GROUP + odd_idx[None, :]) < K),
other=0.0,
)
# We already computed nibble for all GROUP elements; need to extract even/odd
# Recompute for even and odd separately using the same logic
lo_f = lo.to(tl.float32) * inv_scale[:, None]
hi_f = hi.to(tl.float32) * inv_scale[:, None]
lo_sign = tl.where(lo_f < 0.0, 1, 0).to(tl.int32)
lo_abs = tl.abs(lo_f)
lo_code = tl.zeros_like(lo_abs).to(tl.int32)
lo_code = tl.where(lo_abs >= 0.25, 1, lo_code)
lo_code = tl.where(lo_abs >= 0.75, 2, lo_code)
lo_code = tl.where(lo_abs >= 1.25, 3, lo_code)
lo_code = tl.where(lo_abs >= 1.75, 4, lo_code)
lo_code = tl.where(lo_abs >= 2.5, 5, lo_code)
lo_code = tl.where(lo_abs >= 3.5, 6, lo_code)
lo_code = tl.where(lo_abs >= 5.0, 7, lo_code)
lo_nibble = (lo_sign << 3) | lo_code
hi_sign = tl.where(hi_f < 0.0, 1, 0).to(tl.int32)
hi_abs = tl.abs(hi_f)
hi_code = tl.zeros_like(hi_abs).to(tl.int32)
hi_code = tl.where(hi_abs >= 0.25, 1, hi_code)
hi_code = tl.where(hi_abs >= 0.75, 2, hi_code)
hi_code = tl.where(hi_abs >= 1.25, 3, hi_code)
hi_code = tl.where(hi_abs >= 1.75, 4, hi_code)
hi_code = tl.where(hi_abs >= 2.5, 5, hi_code)
hi_code = tl.where(hi_abs >= 3.5, 6, hi_code)
hi_code = tl.where(hi_abs >= 5.0, 7, hi_code)
hi_nibble = (hi_sign << 3) | hi_code
packed = (hi_nibble << 4) | lo_nibble
packed = packed.to(tl.uint8)
# store packed output: shape [M, K//2]
rk_out = pid_k * (GROUP // 2) + tl.arange(0, GROUP // 2)
out_mask = (rm[:, None] < M) & (rk_out[None, :] < K // 2)
tl.store(
Out_ptr + rm[:, None] * stride_om + rk_out[None, :] * stride_ok,
packed,
mask=out_mask,
)
# store scale: shape [M, K//GROUP]
scale_col = pid_k
scale_mask = rm < M
tl.store(
Scale_ptr + rm * stride_sm + scale_col * stride_sk,
scale_e8m0,
mask=scale_mask,
)
# --- Use aiter's proven quantization + GEMM (fastest path) ---
# The Triton quant above is illustrative; for correctness and speed,
# use aiter's fused ops which are HIP-optimized for MI355X.
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
# Quantize A
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(A)
bs_e8m0 = e8m0_shuffle(bs_e8m0)
A_q = x_fp4.view(dtypes.fp4x2)
A_scale_sh = bs_e8m0.view(dtypes.fp8_e8m0)
# GEMM
out_gemm = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out_gemm
scrolls · 197 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