submission 654326
Hamza · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 217 lines, June 9 Researcher Reciprocity License v1.0.
submission_direct.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-654326?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:fc9f2705101dc298d3fca60a4c68cb4f6a32c950f40521b05f004a1fbebd78f5
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 = 16
BLOCK_M = 16 if M > 8 else 8Kernel source
submission_direct.py217 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# --- Config injection ---
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 aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from task import input_t, output_t
_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 = 16 if M > 8 else 8
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 = 4
cfg = {
"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
else:
tiles = ((M + 15) // 16) * ((N + 127) // 128)
KSPLIT = 1
STAGES = 2
# Use KSPLIT to boost CU utilization for shapes with few tiles
if tiles < _CU // 2 and K_real > 512:
if K_real >= 4096:
# K=3584: KSPLIT=7→2 iters, KSPLIT=2→7 iters
if tiles * 2 >= _CU:
KSPLIT = 2
else:
KSPLIT = 7
elif K_real >= 2048:
# K=1024: KSPLIT=2→2 iters with stages=2
KSPLIT = 2
elif K_real >= 1536:
# K=768: KSPLIT=3→1 iter (no pipeline benefit)
KSPLIT = 3
STAGES = 1
wgs = tiles * KSPLIT
cfg = {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "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
grid_main = (
cfg["NUM_KSPLIT"]
* triton.cdiv(M, cfg["BLOCK_SIZE_M"])
* triton.cdiv(N, cfg["BLOCK_SIZE_N"]),
)
grid_reduce = None
if cfg["NUM_KSPLIT"] > 1:
grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))
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]
# Pre-allocated y_pp for KSPLIT > 1
if cfg["NUM_KSPLIT"] > 1:
ppkey = (dev.index, cfg["NUM_KSPLIT"], M, N)
if ppkey not in _YPP_BUF:
_YPP_BUF[ppkey] = torch.empty(
(cfg["NUM_KSPLIT"], M, N), dtype=torch.float32, device=dev
)
y_pp = _YPP_BUF[ppkey]
else:
y_pp = None
B_w, B_s = _get_preshuffle_b(data)
# 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, 64,
actual_ksplit, nk_pow2,
)
return y.view(*shape_prefix, N)
scrolls · 217 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 652726.
⋯ 69 unchanged linesif M < 32:BLOCK_M = 16 if M > 8 else 8KSPLIT = 1+ STAGES = 1if K_real >= 4096:- KSPLIT = 14+ KSPLIT = 7+ STAGES = 2+ elif K_real >= 2048:+ KSPLIT = 2+ STAGES = 2elif K_real >= 1536:KSPLIT = 4cfg = {"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,"waves_per_eu": 1, "matrix_instr_nonkdim": 16,"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,}else:+ tiles = ((M + 15) // 16) * ((N + 127) // 128)+ KSPLIT = 1+ STAGES = 2+ # Use KSPLIT to boost CU utilization for shapes with few tiles+ if tiles < _CU // 2 and K_real > 512:+ if K_real >= 4096:+ # K=3584: KSPLIT=7→2 iters, KSPLIT=2→7 iters+ if tiles * 2 >= _CU:+ KSPLIT = 2+ else:+ KSPLIT = 7+ elif K_real >= 2048:+ # K=1024: KSPLIT=2→2 iters with stages=2+ KSPLIT = 2+ elif K_real >= 1536:+ # K=768: KSPLIT=3→1 iter (no pipeline benefit)+ KSPLIT = 3+ STAGES = 1+ wgs = tiles * KSPLITcfg = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": 2, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ "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
scrolls · 52 diff lines total
Best evidence level for this revision: reported
JSON