submission 651591
Hamza · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 120 lines, June 9 Researcher Reciprocity License v1.0.
submission_tuned.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-651591?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:13e97f76437f368035c95882a769d61541e6226a15a760a7976de8fe59105448
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"]tile-m = 16
BLOCK_M = 16 if M > 8 else 8Kernel source
submission_tuned.py120 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
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from task import input_t, output_t
_PRESHUFFLE_CACHE: dict = {}
_OUT_BUF: 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_preshuffle_config(M: int, N: int, K_real: int) -> dict:
if M < 32:
BLOCK_M = 16 if M > 8 else 8
KSPLIT = 1
if K_real >= 4096:
KSPLIT = 14
elif K_real >= 1536:
KSPLIT = 4
return {
"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
elif M <= 32:
# M=32: BLOCK_M=16 doubles tile count, BLOCK_K=512 (low register pressure)
KSPLIT = 1
if K_real >= 4096:
KSPLIT = 14
elif K_real >= 1536:
KSPLIT = 4
return {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
else:
# M>=64: BLOCK_M=32, BLOCK_K=256 (proven, avoids register spilling)
if K_real >= 4096:
target_ksplit = 4
elif K_real >= 1536:
target_ksplit = 2
else:
target_ksplit = 1
return {
"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": target_ksplit,
}
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
okey = (A.device.index, M, N)
if okey not in _OUT_BUF:
_OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
out = _OUT_BUF[okey]
# Preshuffle for all M (non-preshuffle has ranked correctness issues)
B_w, B_s = _get_preshuffle_b(data)
config = _get_preshuffle_config(M, N, K_real)
result = gemm_a16wfp4_preshuffle(A_2d, B_w, B_s, y=out, config=config)
return result.view(*shape_prefix, N)
scrolls · 120 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 526320.
- """- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).- """- import aiter- from aiter import QuantType, dtypes+ #!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+ from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshufflefrom task import input_t, output_t- QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)- GEMM_A4W4 = aiter.gemm_a4w4- OUTPUT_DTYPE = dtypes.bf16+ _PRESHUFFLE_CACHE: dict = {}+ _OUT_BUF: 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_preshuffle_config(M: int, N: int, K_real: int) -> dict:+ if M < 32:+ BLOCK_M = 16 if M > 8 else 8+ KSPLIT = 1+ if K_real >= 4096:+ KSPLIT = 14+ elif K_real >= 1536:+ KSPLIT = 4+ return {+ "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,+ }+ elif M <= 32:+ # M=32: BLOCK_M=16 doubles tile count, BLOCK_K=512 (low register pressure)+ KSPLIT = 1+ if K_real >= 4096:+ KSPLIT = 14+ elif K_real >= 1536:+ KSPLIT = 4+ return {+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,+ "waves_per_eu": 2, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,+ }+ else:+ # M>=64: BLOCK_M=32, BLOCK_K=256 (proven, avoids register spilling)+ if K_real >= 4096:+ target_ksplit = 4+ elif K_real >= 1536:+ target_ksplit = 2+ else:+ target_ksplit = 1+ return {+ "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,+ "waves_per_eu": 2, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": target_ksplit,+ }++def custom_kernel(data: input_t) -> output_t:- """- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.- gemm_a4w4 with bpreshuffle=True.- """- A, _, _, B_shuffle, B_scale_sh = data- A_q, A_scale_sh = QUANT_FUNC(A, shuffle=True)+ A = data[0]+ if not A.is_contiguous():+ A = A.contiguous()- return GEMM_A4W4(- A_q,- B_shuffle,- A_scale_sh,- B_scale_sh,- dtype=OUTPUT_DTYPE,- bpreshuffle=True,- )+ 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++ okey = (A.device.index, M, N)+ if okey not in _OUT_BUF:+ _OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)+ out = _OUT_BUF[okey]++ # Preshuffle for all M (non-preshuffle has ranked correctness issues)+ B_w, B_s = _get_preshuffle_b(data)+ config = _get_preshuffle_config(M, N, K_real)+ result = gemm_a16wfp4_preshuffle(A_2d, B_w, B_s, y=out, config=config)++ return result.view(*shape_prefix, N)
scrolls · 142 diff lines total
Best evidence level for this revision: reported
JSON