submission 600312
fluudgate · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 267 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-600312?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:78ba0f912d0c7578379765953b980425f6feb31bdc725c8e8f322ebc6a94bc61
license declaredunknown
license concludedunknown
authorsfluudgate
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM — Exp 6: Address real bottlenecks.split-k
d[(cu, 4, 2880, 512)] = {**base, 'kernelId': 21, 'splitK': 0, 'kernelName': k32}tile-m = 4
BLOCK_M = 4tile-n = 128
BLOCK_N = 128Kernel source
submission.py267 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM — Exp 6: Address real bottlenecks.
1. Cache workspaces (no per-call allocation/memset)
2. Remove log2/exp2 — use exponent bit extraction
3. Aligned fast path (no masks for benchmark shapes)
4. Complete 6-shape GEMM config sweep (inject all 6)
5. Keep fused quant+shuffle (still saves 1 launch)
"""
import os
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_USE_NT"] = "1"
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# ---------------------------------------------------------------------------
# Fused MXFP4 quant + e8m0 shuffle — optimized
# - Exponent-bit scale extraction (no log2/exp2)
# - Aligned fast path (no masks)
# ---------------------------------------------------------------------------
@triton.jit
def _mxfp4_quant_op_fast(
x,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
"""Optimized FP4 quant: exponent-bit extraction instead of log2/exp2."""
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
MBITS_F32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
max_normal: tl.constexpr = 6
min_normal: tl.constexpr = 1
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Block max — round to power of 2
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax_int = amax.to(tl.int32, bitcast=True)
amax_int = (amax_int + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
# Extract exponent directly from integer bits (no log2!)
# For a power-of-2 float: exponent = (amax_int >> 23) - 127
# scale_e8m0_unbiased = exponent - 2
raw_exp = (amax_int >> 23).to(tl.int32)
scale_e8m0_unbiased = raw_exp - 127 - 2
# tl.clamp doesn't support int32 — use manual min/max
scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127, scale_e8m0_unbiased)
scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased > 127, 127, scale_e8m0_unbiased)
bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.uint8)
# Reconstruct quant_scale as 2^(-scale_e8m0_unbiased) via integer bit construction (no exp2!)
quant_exp = (-scale_e8m0_unbiased + 127).to(tl.uint32)
quant_scale_int = quant_exp << 23
quant_scale = quant_scale_int.to(tl.float32, bitcast=True)
qx = x * quant_scale
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
qx = qx ^ s
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
normal_mask = ~(saturate_mask | denormal_mask)
# Denormal path
denorm_exp: tl.constexpr = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
denormal_x = qx_fp32 + denorm_mask_float
denormal_x = denormal_x.to(tl.uint32, bitcast=True)
denormal_x -= denorm_mask_int
denormal_x = denormal_x.to(tl.uint8)
# Normal path
normal_x = qx
mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
normal_x += val_to_add
normal_x += mant_odd
normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
normal_x = normal_x.to(tl.uint8)
# Merge
e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
sign_lp = sign_lp.to(tl.uint8)
e2m1_value = e2m1_value | sign_lp
# Pack pairs
e2m1_value = tl.reshape(
e2m1_value,
[BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2],
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _fused_quant_shuffle_aligned(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m, stride_x_n, stride_fp4_m, stride_fp4_n,
M, N, scaleN_pad,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
"""Aligned fast path: no masks, no validity checks. For benchmark shapes only."""
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_xm = tl.cast(stride_x_m, tl.int64)
stride_xn = tl.cast(stride_x_n, tl.int64)
stride_fm = tl.cast(stride_fp4_m, tl.int64)
stride_fn = tl.cast(stride_fp4_n, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_xm + x_offs_n[None, :] * stride_xn
# No mask — shapes are aligned
x = tl.load(x_ptr + x_offs).to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op_fast(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
fp4_offs_n = pid_n * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
fp4_offs = x_offs_m[:, None] * stride_fm + fp4_offs_n[None, :] * stride_fn
tl.store(x_fp4_ptr + fp4_offs, out_tensor)
# Shuffled scale store
bs_m = x_offs_m
bs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
i0 = bs_m[:, None] // 32
i1 = (bs_m[:, None] % 32) // 16
i2 = bs_m[:, None] % 16
i3 = bs_n[None, :] // 8
i4 = (bs_n[None, :] % 8) // 4
i5 = bs_n[None, :] % 4
shuffled_offs = i1 + i4 * 2 + i2 * 4 + i5 * 64 + i3 * 256 + i0 * (32 * scaleN_pad)
tl.store(bs_ptr + shuffled_offs, bs_e8m0)
# ---------------------------------------------------------------------------
# Workspace cache — avoid per-call allocation/memset
# ---------------------------------------------------------------------------
_workspace_cache = {}
def _get_workspace(M, K, device, dtypes):
key = (M, K)
if key in _workspace_cache:
return _workspace_cache[key]
MXFP4_BLOCK = 32
scaleN_valid = triton.cdiv(K, MXFP4_BLOCK)
M_p = triton.cdiv(M, 256) * 256
N_sp = triton.cdiv(scaleN_valid, 8) * 8
x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
# Initialize padded scale buffer ONCE with 127
bs_shuffled = torch.full((M_p * N_sp,), 127, dtype=torch.uint8, device=device)
_workspace_cache.clear() # Only cache one shape at a time
_workspace_cache[key] = (x_fp4, bs_shuffled, scaleN_valid, M_p, N_sp)
return _workspace_cache[key]
# ---------------------------------------------------------------------------
# Initialization
# ---------------------------------------------------------------------------
_inited = False
_aiter = None
_bf16 = None
_dtypes = None
def _ensure_init():
global _inited, _aiter, _bf16, _dtypes
if _inited:
return
_inited = True
import aiter
from aiter import dtypes
_aiter = aiter
_bf16 = dtypes.bf16
_dtypes = dtypes
# EVOLVE-BLOCK-START gemm_config_patch
# Inject configs for ALL 6 benchmark shapes
try:
from aiter.ops.gemm_op_a4w4 import get_GEMM_config
_ = get_GEMM_config(1, 1, 1)
if hasattr(get_GEMM_config, "gemm_dict"):
d = get_GEMM_config.gemm_dict
cu = 256
k32 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E'
k64 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E'
k96 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_96x128E'
k128 = '_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_128x128E'
base = {'us': 0, 'tflops': 0, 'bw': 0, 'errRatio': 0}
# Small M: bandwidth-optimized (32x128)
d[(cu, 4, 2880, 512)] = {**base, 'kernelId': 21, 'splitK': 0, 'kernelName': k32}
d[(cu, 16, 2112, 7168)] = {**base, 'kernelId': 21, 'splitK': 0, 'kernelName': k32}
# Medium M: transitional (64x128)
d[(cu, 32, 4096, 512)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}
d[(cu, 32, 2880, 512)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}
# Large M: compute-optimized (96x128 or 128x128)
d[(cu, 64, 7168, 2048)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}
d[(cu, 256, 3072, 1536)] = {**base, 'kernelId': 29, 'splitK': 0, 'kernelName': k64}
get_GEMM_config.cache_clear()
except Exception:
pass
# EVOLVE-BLOCK-END gemm_config_patch
MXFP4_BLOCK = 32
BLOCK_M = 4
BLOCK_N = 128
# EVOLVE-BLOCK-START gemm_dispatch
def custom_kernel(data: input_t) -> output_t:
_ensure_init()
A, B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
# Cached workspace — no per-call allocation or memset
x_fp4, bs_shuffled, scaleN_valid, M_p, N_sp = _get_workspace(M, K, A.device, _dtypes)
grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, BLOCK_N))
_fused_quant_shuffle_aligned[grid](
A, x_fp4, bs_shuffled,
A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1),
M, K, N_sp,
BLOCK_SIZE_M=BLOCK_M, BLOCK_SIZE_N=BLOCK_N, MXFP4_QUANT_BLOCK_SIZE=MXFP4_BLOCK,
)
A_q = x_fp4.view(_dtypes.fp4x2)
A_scale_sh = bs_shuffled.view(M_p, N_sp).view(_dtypes.fp8_e8m0)
return _aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=_bf16, bpreshuffle=True,
)
# EVOLVE-BLOCK-END gemm_dispatch
scrolls · 267 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