submission 658767
Hamza · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 256 lines, June 9 Researcher Reciprocity License v1.0.
submission_direct.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-658767?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:4ed9b595efaebef4b4acff9b32c5adbd7f8b262617072edd9390a272061f3aae
license declaredunknown
license concludedunknown
authorsHamza
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]stages = 1
STAGES = 1tile-m = 8
BLOCK_M = 8tile-n = 128
BLOCK_N = 128Kernel source
submission_direct.py256 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# --- Config injection (prevents extra module_gemm_common build) ---
import os as _os
_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_CSV_PATH = "/tmp/_mxfp4_mm_config.csv"
_CU = 256
_NK_FAMILIES = [
(2880, 512), (2112, 7168), (4096, 512), (7168, 2048), (3072, 1536),
(2880, 1536), (4096, 1536), (2112, 512), (2112, 2048),
(7168, 512), (7168, 1536), (7168, 7168), (3072, 512),
(3072, 7168), (3072, 2048), (4096, 2048), (4096, 7168),
(2880, 2048), (2880, 7168),
]
_M_VALUES = [1, 2, 4, 8, 16, 32, 64, 128, 256]
_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
for _n, _k in _NK_FAMILIES:
for _m in _M_VALUES:
_tile_num = ((_m + 31) // 32) * ((_n + 127) // 128)
_cus_per_tile = _CU / max(_tile_num, 1)
_split = 0
while _cus_per_tile >= pow(2, _split + 1) and (pow(2, _split + 1) * 128) < 2 * _k:
_split += 1
_split = min(_split, 3)
_lines.append(f"{_CU},{_m},{_n},{_k},21,{_split},1.0,{_KERNEL_32x128},0,0,0.0")
with open(_CSV_PATH, "w") as _f:
_f.write("\n".join(_lines))
_os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
# --- End config injection ---
import torch
import triton
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
from task import input_t, output_t
def _get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
"""Adjust KSPLIT/BLOCK_K for EVEN_K alignment (inlined from aiter)."""
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 // 2
elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
if NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
else:
break
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
_PRESHUFFLE_CACHE: dict = {}
_OUT_BUF: dict = {}
_YPP_BUF: dict = {}
_CFG_CACHE: dict = {}
def _get_preshuffle_b(data):
key = data[3].data_ptr()
if key not in _PRESHUFFLE_CACHE:
N = data[3].shape[0]
K_bytes = data[3].shape[1]
sm, sn = data[4].shape
N_groups = N // 32
B_w = data[3].view(torch.uint8).reshape(N // 16, K_bytes * 16)
B_s = data[4].view(torch.uint8).reshape(sm // 32, sn * 32)[:N_groups].contiguous()
_PRESHUFFLE_CACHE[key] = (B_w, B_s)
return _PRESHUFFLE_CACHE[key]
def _get_cfg(M: int, N: int, K_real: int):
key = (M, N, K_real)
if key in _CFG_CACHE:
return _CFG_CACHE[key]
K = K_real // 2 # Internal K dimension
if M <= 32:
BLOCK_M = 8
BLOCK_N = 128
KSPLIT = 1
STAGES = 1
if K_real >= 4096:
KSPLIT = 7
STAGES = 2
elif K_real >= 2048:
KSPLIT = 2
STAGES = 2
elif K_real >= 1536:
KSPLIT = 3
STAGES = 2
# Use BLOCK_N=64 when CU utilization is low
tiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
if tiles_128 * KSPLIT < (_CU * 3) // 4:
BLOCK_N = 64
wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
cfg = {
"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,
"waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
else:
tiles = ((M + 15) // 16) * ((N + 127) // 128)
BLOCK_N = 128
KSPLIT = 1
STAGES = 1 if K_real <= 512 else 2 # 1 K-iter → no pipeline benefit
# Use KSPLIT to boost CU utilization for shapes with few tiles
if K_real >= 7168 and _CU // 2 <= tiles < _CU:
# K=7168 has 14 K-iters: KSPLIT=2 halves to 7 with small reduce overhead
KSPLIT = 2
elif tiles < _CU // 2 and K_real > 512:
if K_real >= 4096:
if tiles * 2 >= _CU:
KSPLIT = 2
else:
KSPLIT = 7
elif K_real >= 2048:
KSPLIT = 2
elif K_real >= 1536:
KSPLIT = 3
STAGES = 2
# Use BLOCK_N=64 when CU utilization is low
if tiles * KSPLIT < (_CU * 3) // 4:
BLOCK_N = 64
wgs = ((M + 15) // 16) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
cfg = {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,
"waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
# Apply get_splitk to adjust KSPLIT
if cfg["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = _get_splitk(
K, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
)
cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
cfg["BLOCK_SIZE_K"] = BLOCK_SIZE_K
cfg["NUM_KSPLIT"] = NUM_KSPLIT
# Handle BLOCK_K >= 2*K edge case
if cfg["BLOCK_SIZE_K"] >= 2 * K:
cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
cfg["SPLITK_BLOCK_SIZE"] = 2 * K
cfg["NUM_KSPLIT"] = 1
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
if cfg["NUM_KSPLIT"] == 1:
cfg["SPLITK_BLOCK_SIZE"] = 2 * K
# Pre-compute reduce kernel params
actual_ksplit = None
nk_pow2 = None
if cfg["NUM_KSPLIT"] > 1:
actual_ksplit = triton.cdiv(K, cfg["SPLITK_BLOCK_SIZE"] // 2)
nk_pow2 = triton.next_power_of_2(cfg["NUM_KSPLIT"])
# Pre-compute grid
num_m_tiles = triton.cdiv(M, cfg["BLOCK_SIZE_M"])
num_n_tiles = triton.cdiv(N, cfg["BLOCK_SIZE_N"])
total_tiles = num_m_tiles * num_n_tiles
grid_main = (cfg["NUM_KSPLIT"] * total_tiles,)
grid_reduce = None
if cfg["NUM_KSPLIT"] > 1:
grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 16))
result = (cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce)
_CFG_CACHE[key] = result
return result
def custom_kernel(data: input_t) -> output_t:
A = data[0]
if not A.is_contiguous():
A = A.contiguous()
shape_prefix = tuple(A.shape[:-1])
A_2d = A.view(-1, A.shape[-1])
M = A_2d.shape[0]
N = data[3].shape[0]
K_bytes = data[3].shape[1]
K_real = K_bytes * 2
K = K_real // 2
cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce = _get_cfg(M, N, K_real)
dev = A.device
okey = (dev.index, M, N)
if okey not in _OUT_BUF:
_OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device=dev)
y = _OUT_BUF[okey]
B_w, B_s = _get_preshuffle_b(data)
# Pre-allocated y_pp for KSPLIT > 1
if cfg["NUM_KSPLIT"] > 1:
ppkey = (dev.index, nk_pow2, M, N)
if ppkey not in _YPP_BUF:
_YPP_BUF[ppkey] = torch.empty(
(nk_pow2, M, N), dtype=torch.float32, device=dev
)
y_pp = _YPP_BUF[ppkey]
else:
y_pp = None
# Launch main GEMM kernel
_gemm_a16wfp4_preshuffle_kernel[grid_main](
A_2d, B_w,
y if y_pp is None else y_pp,
B_s,
M, N, K,
A_2d.stride(0), A_2d.stride(1),
B_w.stride(0), B_w.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
B_s.stride(0), B_s.stride(1),
PREQUANT=True,
**cfg,
)
# Reduce if KSPLIT > 1
if y_pp is not None:
_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),
16, 16,
actual_ksplit, nk_pow2,
)
return y.view(*shape_prefix, N)
scrolls · 256 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 657965.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X- # --- Config injection ---+ # --- Config injection (prevents extra module_gemm_common build) ---import os as _os_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"⋯ 30 unchanged linesfrom aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (_gemm_afp4wfp4_reduce_kernel,)- from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitkfrom task import input_t, output_t+ def _get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):+ """Adjust KSPLIT/BLOCK_K for EVEN_K alignment (inlined from aiter)."""+ 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 // 2+ elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:+ if NUM_KSPLIT > 1:+ NUM_KSPLIT = NUM_KSPLIT // 2+ elif BLOCK_SIZE_K > 16:+ BLOCK_SIZE_K = BLOCK_SIZE_K // 2+ elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:+ BLOCK_SIZE_K = BLOCK_SIZE_K // 2+ else:+ break+ SPLITK_BLOCK_SIZE = (+ triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K+ )+ return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT++_PRESHUFFLE_CACHE: dict = {}_OUT_BUF: dict = {}_YPP_BUF: dict = {}⋯ 33 unchanged linesSTAGES = 2elif K_real >= 1536:KSPLIT = 3+ STAGES = 2# Use BLOCK_N=64 when CU utilization is lowtiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)if tiles_128 * KSPLIT < (_CU * 3) // 4:BLOCK_N = 64+ wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLITcfg = {"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,- "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,}else:tiles = ((M + 15) // 16) * ((N + 127) // 128)BLOCK_N = 128KSPLIT = 1- STAGES = 2+ STAGES = 1 if K_real <= 512 else 2 # 1 K-iter → no pipeline benefit# Use KSPLIT to boost CU utilization for shapes with few tiles- if tiles < _CU // 2 and K_real > 512:+ if K_real >= 7168 and _CU // 2 <= tiles < _CU:+ # K=7168 has 14 K-iters: KSPLIT=2 halves to 7 with small reduce overhead+ KSPLIT = 2+ elif tiles < _CU // 2 and K_real > 512:if K_real >= 4096:if tiles * 2 >= _CU:KSPLIT = 2⋯ 3 unchanged linesKSPLIT = 2elif K_real >= 1536:KSPLIT = 3- STAGES = 1+ STAGES = 2# Use BLOCK_N=64 when CU utilization is lowif tiles * KSPLIT < (_CU * 3) // 4:BLOCK_N = 64⋯ 7 unchanged lines# Apply get_splitk to adjust KSPLITif cfg["NUM_KSPLIT"] > 1:- SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(+ SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = _get_splitk(K, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"])cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE⋯ 10 unchanged linesif cfg["NUM_KSPLIT"] == 1:cfg["SPLITK_BLOCK_SIZE"] = 2 * K- # Use BLOCK_K=1024 for KSPLIT=1 K>=2048 with waves_per_eu=1- # Halves K-iterations (4→2 for k=2048, 14→7 for k=7168)- # Only safe with waves_per_eu=1 (register pressure too high for wave overlap)- if (cfg["NUM_KSPLIT"] == 1 and K_real >= 2048- and cfg["waves_per_eu"] == 1 and cfg["BLOCK_SIZE_K"] == 512):- cfg["BLOCK_SIZE_K"] = 1024-# Pre-compute reduce kernel paramsactual_ksplit = Nonenk_pow2 = None
scrolls · 112 diff lines total
Best evidence level for this revision: reported
JSON