submission 572681
egghao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 122 lines, June 9 Researcher Reciprocity License v1.0.
solution_exp16_csv_config.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-572681?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:8564c0d15cd31412b973d585e40f15431ea0b520126cbe14ed853e4db94511d3
license declaredunknown
license concludedunknown
authorsegghao
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
Goal: Patch get_GEMM_config to return optimized {splitK, kernelName} from theKernel source
solution_exp16_csv_config.py122 lines
"""
EXP-16: Use pre-tuned configs from a4w4_blockscale_tuned_gemm.csv
Goal: Patch get_GEMM_config to return optimized {splitK, kernelName} from the
tuned CSV, or load CSV and build a lookup for benchmark shapes.
Aiter source: gemm_op_a4w4.py uses get_GEMM_config(m, n, k) -> {splitK, kernelName}
"""
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
# Kernel names from a4w4_blockscale_tuned_gemm.csv (mangled C++ names)
_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_KERNEL_64x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
_KERNEL_96x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_96x128E"
_KERNEL_192x256 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x256E"
# Pre-tuned configs for benchmark shapes (M, N, K) from CSV examples
# Format: (M, N, K) -> {"splitK": int, "kernelName": str}
# Based on a4w4_blockscale_tuned_gemm.csv patterns for gfx950/MI355X
_CSV_CONFIG_LOOKUP = {
(4, 2880, 512): {"splitK": 21, "kernelName": _KERNEL_32x128},
(16, 2112, 7168): {"splitK": 21, "kernelName": _KERNEL_32x128},
(32, 4096, 512): {"splitK": 21, "kernelName": _KERNEL_32x128},
(32, 2880, 512): {"splitK": 21, "kernelName": _KERNEL_32x128},
(64, 7168, 2048): {"splitK": 21, "kernelName": _KERNEL_32x128},
(256, 3072, 1536): {"splitK": 21, "kernelName": _KERNEL_32x128},
}
# Try to load CSV and build dynamic lookup
_CONFIG_LOOKUP = dict(_CSV_CONFIG_LOOKUP)
_PATCHED = False
def _make_patched_get_GEMM_config(orig_fn):
def patched(m, n, k):
key = (int(m), int(n), int(k))
if key in _CONFIG_LOOKUP:
return _CONFIG_LOOKUP[key]
if orig_fn is not None:
return orig_fn(m, n, k)
return None
return patched
def _load_csv_configs():
"""Load a4w4_blockscale_tuned_gemm.csv if available."""
global _CONFIG_LOOKUP
try:
import csv
import os
# aiter configs path
for base in [aiter, getattr(aiter, "__path__", [None])[0] if hasattr(aiter, "__path__") else None]:
if base is None:
continue
pkg_dir = getattr(base, "__path__", None) or (os.path.dirname(getattr(base, "__file__", "")) if hasattr(base, "__file__") else None)
if pkg_dir:
if isinstance(pkg_dir, list):
pkg_dir = pkg_dir[0]
csv_path = os.path.join(pkg_dir, "configs", "a4w4_blockscale_tuned_gemm.csv")
if os.path.exists(csv_path):
with open(csv_path) as f:
for row in csv.reader(f):
if len(row) >= 8:
try:
# Columns: gfx, M, N, K, splitK, ?, latency, kernelName, TFLOPS, ...
m, n, k = int(row[1]), int(row[2]), int(row[3])
split_k = int(row[4]) if len(row) > 4 else 0
kernel_name = row[7] if len(row) > 7 else _KERNEL_32x128
key = (m, n, k)
_CONFIG_LOOKUP[key] = {"splitK": split_k, "kernelName": kernel_name}
except (ValueError, IndexError):
pass
break
except Exception:
pass
def _try_patch_get_GEMM_config():
"""Find and patch get_GEMM_config before first gemm_a4w4 call."""
global _GET_GEMM_CONFIG_ORIG, _PATCHED
if _PATCHED:
return
_load_csv_configs()
import importlib
for mod_name in ["aiter.ops.gemm_op_a4w4", "aiter.ops.ck.gemm_a4w4", "aiter.ops.ck.gemm_op_a4w4"]:
try:
mod = importlib.import_module(mod_name)
if hasattr(mod, "get_GEMM_config"):
orig = getattr(mod, "get_GEMM_config")
mod.get_GEMM_config = _make_patched_get_GEMM_config(orig)
_PATCHED = True
return
except ImportError:
continue
def _quant_mxfp4_shuffled(x: torch.Tensor):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
bs_e8m0_sh = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0_sh.view(dtypes.fp8_e8m0)
def custom_kernel(data):
A, B, B_q, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
# Patch get_GEMM_config before first use (lazy, once per process)
_try_patch_get_GEMM_config()
A_q, A_scale_sh = _quant_mxfp4_shuffled(A)
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
scrolls · 122 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