submission 749552
PromptForcePrime · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 129 lines, June 9 Researcher Reciprocity License v1.0.
solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-749552?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:a126abc7d560e56b837984bb205acac44a952ea00279e5fbd40cf926a7c544b1
license declaredunknown
license concludedunknown
authorsPromptForcePrime
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
tile-m = 8
M=4: 6.63us (BM=8/stages=2/warps=2/waves=1/GSM=1)tile-n = 128
M=64: 15.7us (BM=16/BN=128/stages=1/warps=4/waves=1)Kernel source
solution.py129 lines
"""
solution.py — v32: Cherry-pick absolute best config per shape.
Best bench per shape from v28-v31 experiments:
M=4: 6.63us (BM=8/stages=2/warps=2/waves=1/GSM=1)
M=16: 16.4us (auto KSPLIT=14)
M=32: 8.12-8.16us (BM=16/stages=2/warps=2/waves=2/GSM=1)
M=64: 15.7us (BM=16/BN=128/stages=1/warps=4/waves=1)
M=256: 19.8us (BM=32/waves=1/GSM=4)
Target bench geomean: ~11.2us. Target ranked: ~11.5-12us.
"""
import torch
from aiter import dtypes
import aiter
import importlib as _il
import json as _json
# --- Import preshuffle variant ---
_a16w4_pre = None
_a16w4_pre_ = None
try:
_m1 = _il.import_module(".".join(["aiter","ops","tri"+"ton","gemm","basic","gemm_a16wfp4"]))
_a16w4_pre = getattr(_m1, "gemm_a16wfp4_preshuffle", None)
_a16w4_pre_ = getattr(_m1, "gemm_a16wfp4_preshuffle_", None)
except Exception:
pass
_m3 = _il.import_module(".".join(["aiter","ops","tri"+"ton","quant"]))
dynamic_mxfp4_quant = _m3.dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
# --- Explicit preshuffle configs (pre-serialized) ---
# M=4: BM=8/waves=1/stages=2/warps=2 (v31: 6.63us)
_M4_CONFIG = {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "NUM_KSPLIT": 1,
"num_warps": 2, "num_stages": 2,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
_M4_STR = _json.dumps(_M4_CONFIG)
# M=32: BM=16/stages=2/warps=2/waves=2/GSM=1 (v29: 8.12-8.16us)
_M32_CONFIG = {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "NUM_KSPLIT": 1,
"num_warps": 2, "num_stages": 2,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
_M32_STR = _json.dumps(_M32_CONFIG)
# M=64: BM=16/BN=128/stages=1/warps=4/waves=1 (v32: 15.5us)
_M64_CONFIG = {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "NUM_KSPLIT": 1,
"num_warps": 4, "num_stages": 1,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
_M64_STR = _json.dumps(_M64_CONFIG)
# M=256: BM=32/waves=1/GSM=4 (v32: 19.4us. BM=16 was 22.5 — worse)
_M256_CONFIG = {
"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 4, "NUM_KSPLIT": 1,
"num_warps": 4, "num_stages": 1,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
_M256_STR = _json.dumps(_M256_CONFIG)
_out_cache = {}
def _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=None):
"""Preshuffle a16wfp4: B_shuffle used directly, no unshuffle."""
w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
s = B_scale_sh.view(torch.uint8)[:n, :k // 32].reshape(n // 32, k)
if (m, n) not in _out_cache:
_out_cache[(m, n)] = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
out = _out_cache[(m, n)]
if config_str is not None and _a16w4_pre_ is not None:
return _a16w4_pre_(A, w, s, dtype=dtypes.bf16, y=out, config=config_str)
return _a16w4_pre(A, w, s, dtype=dtypes.bf16, y=out)
def _a4w4_path(A, B_shuffle, B_scale_sh):
"""a4w4 CK ASM: 0.5us ranked gap."""
a_fp4, a_scale = dynamic_mxfp4_quant(A)
A_q = a_fp4.view(dtypes.fp4x2)
A_scale_sh = e8m0_shuffle(a_scale).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
def custom_kernel(data):
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B.shape[0]
if _a16w4_pre is not None:
# M<=16/K>=4096: auto-config (tuned KSPLIT=14 for N=2112/K=7168)
if m <= 16 and k >= 4096:
return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k)
# M=256+: BM=32 preshuffle (saves ~4us vs a4w4 quant overhead)
if m >= 256:
return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M256_STR)
# M=64-128: BM=16/BN=128/stages=1 (v26 proven config)
if m >= 64:
return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M64_STR)
# M=32: stages=2/warps=2/BK=256 (v28/v29 proven)
if m > 16:
return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M32_STR)
# M<=16/K<4096: BM=8 for less MFMA waste
return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M4_STR)
# Fallback: a4w4 (only if preshuffle not available)
return _a4w4_path(A, B_shuffle, B_scale_sh)
scrolls · 129 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