Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
22.5µs
#758 of 1143
2026-03-17

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-kGoal: Patch get_GEMM_config to return optimized {splitK, kernelName} from the

Kernel 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