submission 728249
TraceByWind · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 308 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-728249?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:b7189954d4d0d5b77285b029456c4eac5f7d15560c49f26e6c28819fd8117e66
license declaredunknown
license concludedunknown
authorsTraceByWind
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
A16WFP4 GEMM with BPRESHUFFLE: BF16 A + MXFP4 B (shuffled) -> BF16 C.split-k
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)tile-m = 16
REDUCE_BLOCK_SIZE_M = 16tile-n = 64
REDUCE_BLOCK_SIZE_N = 64Kernel source
submission.py308 lines
"""
A16WFP4 GEMM with BPRESHUFFLE: BF16 A + MXFP4 B (shuffled) -> BF16 C.
Kernel definitions at module top-level; custom_kernel() prepares data and launches.
"""
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config
from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr
import torch
import triton
import triton.language as tl
import triton._utils as _tu
# --- ROCm custom dtype monkey-patch ---
_tu.type_canonicalisation_dict["float4_e2m1fn_x2"] = "u8"
_tu.type_canonicalisation_dict["float8_e8m0fnu"] = "u8"
# --- Global helper: pid_grid ---
@triton.jit
def pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr = 1):
if GROUP_SIZE_M == 1:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
else:
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)
tl.assume(group_size_m >= 0)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
return pid_m, pid_n
# --- Global: _mxfp4_quant_op (from aiter) ---
@triton.jit
def _mxfp4_quant_op(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
MBITS_FP32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
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)
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)
bs_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)
denorm_exp: tl.constexpr = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_FP32 - MBITS_FP4) + 1
denorm_mask_int: tl.constexpr = denorm_exp << MBITS_FP32
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_x = qx
mant_odd = (normal_x >> (MBITS_FP32 - MBITS_FP4)) & 1
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_FP32) + (1 << 21) - 1
normal_x += val_to_add
normal_x += mant_odd
normal_x = normal_x >> (MBITS_FP32 - 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_FP32 + 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, 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)
# --- Global: _gemm_a16wfp4_preshuffle_repr & kernel ---
_gemm_a16wfp4_preshuffle_repr = make_kernel_repr(
"_gemm_a16wfp4_preshuffle_kernel",
["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K","GROUP_SIZE_M","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim","cache_modifier","NUM_KSPLIT"],
)
@triton.heuristics({
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"]) * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
})
@triton.jit(repr=_gemm_a16wfp4_preshuffle_repr)
def _gemm_a16wfp4_preshuffle_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_ck,
stride_cm, stride_cn,
stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr, num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr, GRID_MN: tl.constexpr,
PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0); tl.assume(stride_ak > 0); tl.assume(stride_bk > 0); tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0); tl.assume(stride_cn > 0); tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
tl.assume(pid_m >= 0); tl.assume(pid_n >= 0); tl.assume(pid_k >= 0)
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
b_scale_ptrs = b_scales_ptr + (offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0,5,3,1,4,2,6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE))
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0,1,4,2,3,5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1,0))
if PREQUANT:
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
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, :] + pid_k * stride_ck)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
# --- Global: get_splitk (Python heuristic) ---
def get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
NUM_KSPLIT_STEP = 2
BLOCK_SIZE_K_STEP = 2
SPLITK_BLOCK_SIZE = (triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K)
while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
if (K % (SPLITK_BLOCK_SIZE // 2) == 0 and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0 and K % (BLOCK_SIZE_K // 2) == 0):
break
elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // NUM_KSPLIT_STEP
elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
if NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // NUM_KSPLIT_STEP
elif BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // BLOCK_SIZE_K_STEP
elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // BLOCK_SIZE_K_STEP
else:
break
SPLITK_BLOCK_SIZE = (triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K)
NUM_KSPLIT = triton.cdiv(K, (SPLITK_BLOCK_SIZE // 2))
return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
# --- Global: _gemm_afp4wfp4_reduce_repr & kernel ---
_gemm_afp4wfp4_reduce_repr = make_kernel_repr("_gemm_afp4wfp4_reduce_kernel", ["BLOCK_SIZE_M","BLOCK_SIZE_N","ACTUAL_KSPLIT","MAX_KSPLIT"])
@triton.heuristics({})
@triton.jit(repr=_gemm_afp4wfp4_reduce_repr)
def _gemm_afp4wfp4_reduce_kernel(
c_in_ptr, c_out_ptr,
M, N,
stride_c_in_k, stride_c_in_m, stride_c_in_n,
stride_c_out_m, stride_c_out_n,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
offs_k = tl.arange(0, MAX_KSPLIT)
c_in_ptrs = (c_in_ptr + (offs_k[:,None,None] * stride_c_in_k) + (offs_m[None,:,None] * stride_c_in_m) + (offs_n[None,None,:] * stride_c_in_n))
if ACTUAL_KSPLIT == MAX_KSPLIT:
c = tl.load(c_in_ptrs)
else:
c = tl.load(c_in_ptrs, mask=offs_k[:,None,None] < ACTUAL_KSPLIT)
c = tl.sum(c, axis=0)
c = c.to(c_out_ptr.type.element_ty)
c_out_ptrs = (c_out_ptr + (offs_m[:,None] * stride_c_out_m) + (offs_n[None,:] * stride_c_out_n))
tl.store(c_out_ptrs, c)
# --- Entry point: custom_kernel ---
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
B_shuffle = B_shuffle.reshape(B_shuffle.shape[0] // 16, B_shuffle.shape[1] * 16).contiguous()
B_scale_sh = B_scale_sh.reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32).contiguous()
M, K_orig = A.shape
N_packed, K_packed = B_shuffle.shape
N = N_packed * 16
K = K_packed // 16
# Load configuration
config, _ = _get_config(M, N, K, shuffle=True)
# Coarse split & BLOCK_SIZE_K adjustment
if config["BLOCK_SIZE_K"] >= 2 * K:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
config["SPLITK_BLOCK_SIZE"] = 2 * K
config["NUM_KSPLIT"] = 1
if config["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
K, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
config["BLOCK_SIZE_K"] = BLOCK_SIZE_K
config["NUM_KSPLIT"] = NUM_KSPLIT
# Post‑get_splitk check (mirrors official)
if config["BLOCK_SIZE_K"] >= 2 * K:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
config["SPLITK_BLOCK_SIZE"] = 2 * K
config["NUM_KSPLIT"] = 1
# BLOCK_SIZE_N lower bound
config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
# Decide output buffer layout
use_y_pp = config["NUM_KSPLIT"] > 1
if use_y_pp:
y_pp = torch.empty((config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=A.device)
y = None
else:
config["SPLITK_BLOCK_SIZE"] = 2 * K
y_pp = None
y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
ck_stride = 0 if not use_y_pp else y_pp.stride(0)
grid = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]),)
_gemm_a16wfp4_preshuffle_kernel[grid](
A,
B_shuffle,
y if not use_y_pp else y_pp,
B_scale_sh,
M, N, K,
A.stride(0), A.stride(1),
B_shuffle.stride(0), B_shuffle.stride(1),
ck_stride,
y.stride(0) if not use_y_pp else y_pp.stride(1),
y.stride(1) if not use_y_pp else y_pp.stride(2),
B_scale_sh.stride(0), B_scale_sh.stride(1),
PREQUANT=True,
**config,
)
if use_y_pp:
REDUCE_BLOCK_SIZE_M = 16
REDUCE_BLOCK_SIZE_N = 64
ACTUAL_KSPLIT = triton.cdiv(K, (config["SPLITK_BLOCK_SIZE"] // 2))
MAX_KSPLIT = triton.next_power_of_2(config["NUM_KSPLIT"])
if y is None:
y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
grid_reduce = (triton.cdiv(M, REDUCE_BLOCK_SIZE_M), triton.cdiv(N, REDUCE_BLOCK_SIZE_N))
_gemm_afp4wfp4_reduce_kernel[grid_reduce](
y_pp,
y,
M, N,
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
y.stride(0), y.stride(1),
REDUCE_BLOCK_SIZE_M, REDUCE_BLOCK_SIZE_N,
ACTUAL_KSPLIT, MAX_KSPLIT,
)
return y[:M, :N]
scrolls · 308 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