Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
13.3µs
#419 of 1143
2026-03-11

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-kVersion 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
+ break
except 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