submission 528645
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 81 lines, June 9 Researcher Reciprocity License v1.0.
mxfp4_v38_merge_all.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-528645?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:285b4cfabb9ea56da1b5bc2b1536b638b46dd4d597a788611216ca420d002323
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
Version 38: Merge config with ALL optimal splitK values.Kernel source
mxfp4_v38_merge_all.py81 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Version 38: Merge config with ALL optimal splitK values.
Reads original aiter CSV, merges custom splitK for M=4, M=16, M=32 shapes.
Preserves original tuned configs for M=64, M=256 etc.
"""
import os
import glob
import pandas as pd
_CUSTOM_CONFIG_PATH = "/tmp/custom_a4w4_all_config.csv"
_aiter_base = "/home/runner/aiter"
_possible_paths = [
f"{_aiter_base}/aiter/configs/a4w4_blockscale_tuned_gemm.csv",
f"{_aiter_base}/hsa/configs/a4w4_tuned_gemm.csv",
f"{_aiter_base}/configs/a4w4_tuned_gemm.csv",
]
for pattern in [f"{_aiter_base}/**/a4w4*tuned*.csv", f"{_aiter_base}/**/*a4w4*.csv"]:
_possible_paths.extend(glob.glob(pattern, recursive=True))
_original_df = None
for path in _possible_paths:
if os.path.exists(path):
try:
df = pd.read_csv(path)
if 'kernelName' in df.columns and 'splitK' in df.columns:
_original_df = df
break
except Exception:
continue
cu_num = 256
_kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
custom_rows = [
{"cu_num": cu_num, "M": 4, "N": 2880, "K": 512,
"kernelName": _kernel_32x128, "splitK": 2},
{"cu_num": cu_num, "M": 16, "N": 2112, "K": 7168,
"kernelName": _kernel_32x128, "splitK": 3},
{"cu_num": cu_num, "M": 32, "N": 2880, "K": 512,
"kernelName": _kernel_32x128, "splitK": 2},
{"cu_num": cu_num, "M": 32, "N": 4096, "K": 512,
"kernelName": _kernel_32x128, "splitK": 2},
]
custom_df = pd.DataFrame(custom_rows)
if _original_df is not None:
for _, row in custom_df.iterrows():
mask = (_original_df['M'] == row['M']) & \
(_original_df['N'] == row['N']) & \
(_original_df['K'] == row['K'])
_original_df = _original_df[~mask]
combined_df = pd.concat([_original_df, custom_df], ignore_index=True)
else:
combined_df = custom_df
combined_df.to_csv(_CUSTOM_CONFIG_PATH, index=False)
os.environ["AITER_CONFIG_GEMM_A4W4"] = _CUSTOM_CONFIG_PATH
import torch
import aiter
from aiter import QuantType, dtypes
from task import input_t, output_t
_quant_func = aiter.get_triton_quant(QuantType.per_1x32)
_bf16 = dtypes.bf16
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A_q, A_scale_sh = _quant_func(A, shuffle=True)
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=_bf16, bpreshuffle=True,
)
scrolls · 81 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 528216.
- from __future__ import annotations+ #!POPCORN leaderboard amd-mxfp4-mm+ #!POPCORN gpu MI355X- import csv- from pathlib import Path+ """+ Version 38: Merge config with ALL optimal splitK values.+ Reads original aiter CSV, merges custom splitK for M=4, M=16, M=32 shapes.+ Preserves original tuned configs for M=64, M=256 etc.+ """+ import os+ import glob+ import pandas as pd- import torch- from task import input_t, output_t+ _CUSTOM_CONFIG_PATH = "/tmp/custom_a4w4_all_config.csv"- _AITER = None- _CLEAR_L2_CACHE = None- _DTYPES = None- _FP4_UTILS = None- _TRITON_QUANT_FUNC = None- _ASM_KERNELS = None+ _aiter_base = "/home/runner/aiter"+ _possible_paths = [+ f"{_aiter_base}/aiter/configs/a4w4_blockscale_tuned_gemm.csv",+ f"{_aiter_base}/hsa/configs/a4w4_tuned_gemm.csv",+ f"{_aiter_base}/configs/a4w4_tuned_gemm.csv",+ ]- _USE_CACHED_TRITON_QUANT = True+ for pattern in [f"{_aiter_base}/**/a4w4*tuned*.csv", f"{_aiter_base}/**/*a4w4*.csv"]:+ _possible_paths.extend(glob.glob(pattern, recursive=True))- _BEST_IMPLS = {}- _A_QUANT_BUFFERS = {}- _OUT_BUFFERS = {}-- _SELECT_ITERS = 5- _SELECT_WARMUP_ITERS = 1- _RTOL = 1e-2- _ATOL = 1e-2-- _PREFERRED_KERNELS = {- (4, 2880, 512): (- "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",- 0,- ),- (16, 2112, 7168): (- "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",- 0,- ),- (32, 4096, 512): (- "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",- 0,- ),- (32, 2880, 512): (- "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",- 0,- ),- (64, 7168, 2048): (- "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",- 0,- ),- (256, 3072, 1536): (- "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",- 0,- ),- }--- def _ceil_div(x: int, y: int) -> int:- return (x + y - 1) // y--- def _lazy_init():- global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELS- if _AITER is not None:- return-- import aiter- from aiter import QuantType, dtypes- from aiter.jit.utils.chip_info import get_gfx- from aiter.utility import fp4_utils- from utils import clear_l2_cache_large-- _AITER = aiter- _CLEAR_L2_CACHE = clear_l2_cache_large- _DTYPES = dtypes- _FP4_UTILS = fp4_utils- _TRITON_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)- _ASM_KERNELS = _load_asm_kernels(aiter, get_gfx())--- def _load_asm_kernels(aiter, gfx: str):- root = Path(aiter.__file__).resolve().parents[1]- csv_path = root / "hsa" / gfx / "f4gemm" / "f4gemm_bf16_per1x32Fp4.csv"- kernels = []-- with csv_path.open(newline="") as f:- for row in csv.DictReader(f):- if int(row["bpreshuffle"]) != 1:- continue- kernels.append(- (- row["knl_name"],- int(row["tile_M"]),- int(row["tile_N"]),- int(row["splitK"]),- )- )-- return tuple(kernels)--- def _kernel_spec(kernel_name: str, split_k: int):- return ("asm", kernel_name, split_k)--- def _candidate_specs(m: int, n: int, k: int):- shape_key = (m, n, k)- wrapper_spec = ("wrapper", "wrapper", 0)- specs = [wrapper_spec]- seen = {wrapper_spec}-- preferred = _PREFERRED_KERNELS.get(shape_key)- if preferred is not None:- spec = _kernel_spec(*preferred)- specs.append(spec)- seen.add(spec)-- padded_m = _ceil_div(m, 32) * 32-- def _kernel_rank(entry):- kernel_name, tile_m, tile_n, split_k = entry- m_penalty = abs(tile_m - min(max(32, padded_m), 256))- n_penalty = abs(tile_n - min(max(128, n), 1024))- split_penalty = split_k- return (m_penalty, n_penalty, split_penalty, tile_m, tile_n, kernel_name)-- for kernel_name, _tile_m, _tile_n, split_k in sorted(_ASM_KERNELS, key=_kernel_rank):- spec = _kernel_spec(kernel_name, split_k)- if spec not in seen:- specs.append(spec)- seen.add(spec)-- return tuple(specs)--- def _get_quant_buffers(m: int, k: int, device):- scale_n = _ceil_div(k, 32)- scale_n_pad = _ceil_div(scale_n, 8) * 8- scale_m_pad = _ceil_div(m, 256) * 256- key = (device.type, device.index, m, k, scale_m_pad, scale_n_pad)- buffers = _A_QUANT_BUFFERS.get(key)-- if buffers is None:- import torch-- a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)- a_scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)- buffers = (a_q, a_scale, scale_n, scale_m_pad, scale_n_pad)- _A_QUANT_BUFFERS[key] = buffers-- return buffers--- def _cached_quantize_a(a):- m = int(a.shape[0])- k = int(a.shape[1])- a_q, a_scale, scale_n, scale_m_pad, scale_n_pad = _get_quant_buffers(m, k, a.device)-- block_size = 128- grid = (_ceil_div(m, block_size), scale_n_pad)-- _FP4_UTILS._dynamic_mxfp4_quant_kernel_asm_layout[grid](- a,- a_q,- a_scale,- *a.stride(),- *a_q.stride(),- *a_scale.stride(),- M=m,- N=k,- scaleN=scale_n,- scaleM_pad=scale_m_pad,- scaleN_pad=scale_n_pad,- BLOCK_SIZE=block_size,- MXFP4_QUANT_BLOCK_SIZE=32,- SCALING_MODE=0,- SHUFFLE=True,- )-- return a_q.view(_DTYPES.fp4x2), a_scale.view(_DTYPES.fp8_e8m0)--- def _quantize_a(a):- if _USE_CACHED_TRITON_QUANT:+ _original_df = None+ for path in _possible_paths:+ if os.path.exists(path):try:- return _cached_quantize_a(a)+ df = pd.read_csv(path)+ if 'kernelName' in df.columns and 'splitK' in df.columns:+ _original_df = df+ breakexcept Exception:- pass-- return _TRITON_QUANT_FUNC(a, shuffle=True)--- def _get_out_buffer(m: int, n: int, device):- rows = _ceil_div(m, 32) * 32- key = (device.type, device.index, rows, n)- out = _OUT_BUFFERS.get(key)- if out is None:- import torch-- out = torch.empty((rows, n), dtype=_DTYPES.bf16, device=device)- _OUT_BUFFERS[key] = out- return out--- def _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):- kind, impl_id, log2_k_split = spec-- if kind == "wrapper":- return _AITER.gemm_a4w4(- a_q,- b_shuffle,- a_scale_sh,- b_scale_sh,- dtype=_DTYPES.bf16,- bpreshuffle=True,- )-- _AITER.gemm_a4w4_asm(- a_q,- b_shuffle,- a_scale_sh,- b_scale_sh,- out,- impl_id,- None,- bpreshuffle=True,- log2_k_split=log2_k_split,- )- return out--- def _is_valid(candidate, reference, m: int):- import torch-- return torch.allclose(candidate[:m], reference[:m], rtol=_RTOL, atol=_ATOL)--- def _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):- import torch-- for _ in range(_SELECT_WARMUP_ITERS):- _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)- torch.cuda.synchronize()-- times_ms = []- for _ in range(_SELECT_ITERS):- _CLEAR_L2_CACHE()- start = torch.cuda.Event(enable_timing=True)- end = torch.cuda.Event(enable_timing=True)- start.record()- _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)- end.record()- torch.cuda.synchronize()- times_ms.append(start.elapsed_time(end))-- times_ms.sort()- return times_ms[len(times_ms) // 2]--- def _select_best_impl(a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, k: int, device):- wrapper_spec = ("wrapper", "wrapper", 0)- out = _get_out_buffer(m, n, device)- reference = _run_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)-- best_spec = wrapper_spec- best_time = _time_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)-- for spec in _candidate_specs(m, n, k):- if spec == wrapper_spec:continue- try:- candidate = _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)- if not _is_valid(candidate, reference, m):- continue+ cu_num = 256+ _kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"- current_time = _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)- if current_time < best_time:- best_time = current_time- best_spec = spec- except Exception:- continue+ custom_rows = [+ {"cu_num": cu_num, "M": 4, "N": 2880, "K": 512,+ "kernelName": _kernel_32x128, "splitK": 2},+ {"cu_num": cu_num, "M": 16, "N": 2112, "K": 7168,+ "kernelName": _kernel_32x128, "splitK": 3},+ {"cu_num": cu_num, "M": 32, "N": 2880, "K": 512,+ "kernelName": _kernel_32x128, "splitK": 2},+ {"cu_num": cu_num, "M": 32, "N": 4096, "K": 512,+ "kernelName": _kernel_32x128, "splitK": 2},+ ]- return best_spec+ custom_df = pd.DataFrame(custom_rows)+ if _original_df is not None:+ for _, row in custom_df.iterrows():+ mask = (_original_df['M'] == row['M']) & \+ (_original_df['N'] == row['N']) & \+ (_original_df['K'] == row['K'])+ _original_df = _original_df[~mask]+ combined_df = pd.concat([_original_df, custom_df], ignore_index=True)+ else:+ combined_df = custom_df- @torch.no_grad()- def custom_kernel(data: input_t) -> output_t:- _lazy_init()+ combined_df.to_csv(_CUSTOM_CONFIG_PATH, index=False)+ os.environ["AITER_CONFIG_GEMM_A4W4"] = _CUSTOM_CONFIG_PATH- a, _b, _b_q, b_shuffle, b_scale_sh = data- a = a.contiguous()+ import torch+ import aiter+ from aiter import QuantType, dtypes+ from task import input_t, output_t- m = int(a.shape[0])- k = int(a.shape[1])- n = int(b_shuffle.shape[0])- shape_key = (m, n, k)+ _quant_func = aiter.get_triton_quant(QuantType.per_1x32)+ _bf16 = dtypes.bf16- a_q, a_scale_sh = _quantize_a(a)- spec = _BEST_IMPLS.get(shape_key)- if spec is None:- spec = _select_best_impl(- a_q,- b_shuffle,- a_scale_sh,- b_scale_sh,- m,- n,- k,- a.device,- )- _BEST_IMPLS[shape_key] = spec-- if spec[0] == "wrapper":- return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, None)-- out = _get_out_buffer(m, n, a.device)- return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)[:m]+ def custom_kernel(data: input_t) -> output_t:+ A, B, B_q, B_shuffle, B_scale_sh = data+ A_q, A_scale_sh = _quant_func(A, shuffle=True)+ return aiter.gemm_a4w4(+ A_q, B_shuffle, A_scale_sh, B_scale_sh,+ dtype=_bf16, bpreshuffle=True,+ )
scrolls · 379 diff lines total
Best evidence level for this revision: reported
JSON