submission 588876
StephenCao422 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 220 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-588876?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:acafa11c07e3ed46bb5685b47bacfe57718af38b5ea71b6aed428c86050382dc
license declaredunknown
license concludedunknown
authorsStephenCao422
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps = 4persistent-kernel
num_sk = tl.num_programs(axis=1) if SPLIT_K > 1 else 1split-k
SPLIT_K: tl.constexpr,stages = 3
num_stages=3tile-k = 256
BLOCK_SIZE_K = 256tile-m = 16
BLOCK_SIZE_M = 16tile-n = 128
BLOCK_SIZE_N = 128Kernel source
submission.py220 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import aiter
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.jit
def _fused_quant_gemm_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
SPLIT_K: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
):
pid = tl.program_id(axis=0)
pid_sk = tl.program_id(axis=1) if SPLIT_K > 1 else 0
num_sk = tl.num_programs(axis=1) if SPLIT_K > 1 else 1
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_k = tl.arange(0, BLOCK_SIZE_K)
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak)
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
offs_k_b = tl.arange(0, BLOCK_SIZE_K // 2)
b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_b[None, :] * stride_bk)
a_ptrs += pid_sk * BLOCK_SIZE_K * stride_ak
SCALE_GROUP_SIZE: tl.constexpr = 32
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
# Mathematical un-shuffling mappings for B_scale_sh (6D coordinate inversion)
n_idx_s = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
i0 = n_idx_s // 32
i1 = (n_idx_s % 32) // 16
i2 = n_idx_s % 16
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
# Precompute mask for A if M is not a multiple of BLOCK_SIZE_M
mask_am = offs_am < M
EXP_BIAS_FP4: tl.constexpr = 1
EXP_BIAS_FP32: tl.constexpr = 127
MBITS_F32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
max_int: tl.constexpr = (1 << (EBITS_FP4 + MBITS_FP4)) - 1
max_normal = 2 ** (3 - EXP_BIAS_FP4) * (3 / 2)
min_normal = 2 ** (1 - EXP_BIAS_FP4)
denorm_exp = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
denorm_mask_int = denorm_exp << MBITS_F32
denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
for k in range(pid_sk, tl.cdiv(K, BLOCK_SIZE_K), num_sk):
a = tl.load(a_ptrs, mask=mask_am[:, None] & (offs_k[None, :] < K - k * BLOCK_SIZE_K), other=0.0).to(tl.float32)
# ---------------- A Dynamic Quantization to MXFP4 ---------------- #
x = a.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, SCALE_GROUP_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
a_scales_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
quant_scale = tl.exp2(-scale_e8m0_unbiased)
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 = (not saturate_mask) & (qx_fp32 < min_normal)
normal_mask = not (saturate_mask | denormal_mask)
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_x = qx
mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 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)
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
e2m1_value = tl.reshape(e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, SCALE_GROUP_SIZE // 2, 2])
evens, odds = tl.split(e2m1_value)
a_fp4 = evens | (odds << 4)
a_fp4 = a_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)
a_scales_e8m0 = a_scales_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
# ----------------------------------------------------------------- #
# ---------------- B Load Native ---------------- #
k_b_offset = k * (BLOCK_SIZE_K // 2) * stride_bk
b_T = tl.load(b_ptrs + k_b_offset, mask=offs_k_b[None, :] < (K - k * BLOCK_SIZE_K) // 2, other=0.0)
b = b_T.trans(1, 0)
# ---------------- B_scale Decode ---------------- #
k_idx_base = k * (BLOCK_SIZE_K // 32)
k_idx = k_idx_base + tl.arange(0, BLOCK_SIZE_K // 32)
i3 = k_idx // 8
i4 = (k_idx % 8) // 4
i5 = k_idx % 4
flat_idx = i0[:, None] * K + i3[None, :] * 256 + i5[None, :] * 64 + i2[:, None] * 4 + i4[None, :] * 2 + i1[:, None]
b_scales_e8m0 = tl.load(b_scales_ptr + flat_idx, mask=k_idx[None, :] < K // 32, other=0)
# ----------------------------------------------------------------- #
accumulator = tl.dot_scaled(a_fp4, a_scales_e8m0, "e2m1", b, b_scales_e8m0, "e2m1", accumulator)
a_ptrs += num_sk * BLOCK_SIZE_K * stride_ak
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
if SPLIT_K > 1:
tl.atomic_add(c_ptrs, c, mask=c_mask)
else:
tl.store(c_ptrs, c, mask=c_mask)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
M, K = A.shape
N, _ = B.shape
# Dynamic Tuning for extremely skewed M shapes in MoE
BLOCK_SIZE_N = 128
BLOCK_SIZE_K = 256
GROUP_SIZE_M = 8
# Dynamic Tuning for extremely skewed M shapes in MoE
if M <= 16:
BLOCK_SIZE_M = 16
num_warps = 4
elif M <= 32:
BLOCK_SIZE_M = 32
num_warps = 4
elif M <= 64:
BLOCK_SIZE_M = 64
num_warps = 8
else:
BLOCK_SIZE_M = 128
num_warps = 8
TOTAL_SPATIAL_BLOCKS = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
if K <= 512:
SPLIT_K = 1
elif TOTAL_SPATIAL_BLOCKS < 120:
desired_sk = 120 // TOTAL_SPATIAL_BLOCKS
k_chunks = triton.cdiv(K, BLOCK_SIZE_K)
SPLIT_K = max(1, min(desired_sk, k_chunks))
SPLIT_K = min(SPLIT_K, 16)
else:
SPLIT_K = 1
if SPLIT_K > 1:
C_out = torch.zeros((M, N), device=A.device, dtype=torch.float32)
else:
C_out = torch.empty((M, N), device=A.device, dtype=torch.bfloat16)
grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), META['SPLIT_K'])
_fused_quant_gemm_kernel[grid](
A, B_q.view(torch.uint8), C_out, B_scale_sh.view(torch.uint8),
M, N, K,
A.stride(0), A.stride(1),
B_q.stride(0), B_q.stride(1),
C_out.stride(0), C_out.stride(1),
SPLIT_K=SPLIT_K,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
BLOCK_SIZE_K=BLOCK_SIZE_K,
GROUP_SIZE_M=GROUP_SIZE_M,
num_warps=num_warps,
num_stages=3
)
if SPLIT_K > 1:
return C_out.to(torch.bfloat16)
return C_out
scrolls · 220 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