submission 662095
Hamza · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 265 lines, June 9 Researcher Reciprocity License v1.0.
submission_direct.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-662095?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:356286fc10c0da3f20a7b174a7ccebaa3dc430af47f328f49326785876aaa343
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 = 2
STAGES = 2tile-k = 256
BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512tile-m = 8
BLOCK_M = 8tile-n = 128
BLOCK_N = 128Kernel source
submission_direct.py265 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 = 2
if K_real >= 4096:
KSPLIT = 7
elif K_real >= 2048:
KSPLIT = 4
elif K_real >= 1536:
KSPLIT = 3
# Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline
BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512
# 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": BLOCK_K,
"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:
# Use BLOCK_M=8 for M<=64 when CU utilization with BM=16 is low
BLOCK_M = 16
if M <= 128:
tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)
if tiles_bm16 < (_CU * 3) // 4:
BLOCK_M = 8
tiles = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
BLOCK_N = 128
KSPLIT = 1
STAGES = 2
# 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
# KSPLIT=2 to reduce wave tail for 1.x-wave shapes with K=2048
if KSPLIT == 1 and _CU < tiles <= 2 * _CU and K_real == 2048:
KSPLIT = 2
# Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline
BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512
# Use BLOCK_N=64 when CU utilization is low
if tiles * 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": BLOCK_K,
"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 · 265 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 658767.
⋯ 99 unchanged linesBLOCK_M = 8BLOCK_N = 128KSPLIT = 1- STAGES = 1+ STAGES = 2if K_real >= 4096:KSPLIT = 7- STAGES = 2elif K_real >= 2048:- KSPLIT = 2- STAGES = 2+ KSPLIT = 4elif K_real >= 1536:KSPLIT = 3- STAGES = 2+ # Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline+ BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512# 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 = 64wgs = ((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,+ "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,"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)+ # Use BLOCK_M=8 for M<=64 when CU utilization with BM=16 is low+ BLOCK_M = 16+ if M <= 128:+ tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)+ if tiles_bm16 < (_CU * 3) // 4:+ BLOCK_M = 8+ tiles = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)BLOCK_N = 128KSPLIT = 1- STAGES = 1 if K_real <= 512 else 2 # 1 K-iter → no pipeline benefit+ STAGES = 2# Use KSPLIT to boost CU utilization for shapes with few tiles- if K_real >= 7168 and _CU // 2 <= tiles < _CU:+ if K_real >= 7168 and _CU // 2 <= tiles <= _CU:# K=7168 has 14 K-iters: KSPLIT=2 halves to 7 with small reduce overheadKSPLIT = 2elif tiles < _CU // 2 and K_real > 512:⋯ 6 unchanged linesKSPLIT = 2elif K_real >= 1536:KSPLIT = 3- STAGES = 2+ # KSPLIT=2 to reduce wave tail for 1.x-wave shapes with K=2048+ if KSPLIT == 1 and _CU < tiles <= 2 * _CU and K_real == 2048:+ KSPLIT = 2+ # Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline+ BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512# Use BLOCK_N=64 when CU utilization is lowif tiles * KSPLIT < (_CU * 3) // 4:BLOCK_N = 64- wgs = ((M + 15) // 16) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT+ wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLITcfg = {- "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,+ "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,"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,
scrolls · 70 diff lines total
Best evidence level for this revision: reported
JSON