Skip to content
KernelIndex
Search⌘K

submission 673983

rujutafujuta · 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.

mxfp4-mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-673983?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
24.0µs
#851 of 1143
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3b270a148946aff8e401878f680f735860330ebee4cd53e30b0680a983df0147
license declaredunknown
license concludedunknown
authorsrujutafujuta
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4MXFP4-MM optimized submission.
split-k- Runtime-tune across all valid kernelIds × splitK {0} per shape.

Kernel source

mxfp4-mm.py129 lines
#!POPCORN leaderboard amd-mxfp4-mm
"""
MXFP4-MM optimized submission.

Strategy:
- Discover valid kernelIds for gemm_a4w4_blockscale_tune at startup (server has
  different compiled kernel IDs than the local CSV, starting from 0).
- Probe up to 24 IDs (reduced from 64 to avoid timeouts).
- Runtime-tune across all valid kernelIds × splitK {0} per shape.
  splitK=1,2 causes numerical errors for large-K shapes (k=7168) — removed.
- Fall back to aiter.gemm_a4w4 (reference path) if no blockscale kernel found.
- Reduced timing iterations (1 warmup, 3 timed) to stay within timeout budget.
"""

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
from task import input_t, output_t

# Cached best (kernelId, splitK) or None (→ fallback) per (m, n, k)
_TUNE_CACHE: dict = {}

# Valid kernelIds discovered once at first use
_VALID_KERNEL_IDS: list | None = None
_SPLIT_KS = [0]  # splitk=1,2 causes numerical errors for large-K shapes (k=7168)
_MAX_KERNEL_PROBE = 24  # don't probe all 64 — stop early to avoid timeouts


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 _discover_kernel_ids(out, A_flat, B_shuffle, A_scale_sh, B_scale_sh) -> list:
    """
    Find all valid kernelId values by probing 0, 1, 2, ... until out-of-range.
    Stops at the first 'out of range' error or after _MAX_KERNEL_PROBE attempts.
    """
    valid = []
    for kid in range(_MAX_KERNEL_PROBE):
        try:
            aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, 0)
            torch.cuda.synchronize()
            valid.append(kid)
        except RuntimeError as e:
            if "out of range" in str(e).lower():
                break
            # Other RuntimeErrors (shape mismatch etc.) — skip this id but keep going
            continue
        except Exception:
            continue
    return valid


def _tune_shape(A_flat, B_shuffle, A_scale_sh, B_scale_sh, m, n):
    """Return (kernelId, splitK) for fastest config, or None to use fallback."""
    global _VALID_KERNEL_IDS

    padded_m = ((m + 31) // 32) * 32
    out = torch.empty((padded_m, n), dtype=dtypes.bf16, device="cuda")

    # Discover valid IDs once using the first shape we see
    if _VALID_KERNEL_IDS is None:
        _VALID_KERNEL_IDS = _discover_kernel_ids(out, A_flat, B_shuffle, A_scale_sh, B_scale_sh)

    if not _VALID_KERNEL_IDS:
        return None  # no blockscale kernels available, use fallback

    best_us = float("inf")
    best_config = None  # only set on a successful timed run

    for kid in _VALID_KERNEL_IDS:
        for sk in _SPLIT_KS:
            try:
                # 1 warmup run
                aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, sk)
                torch.cuda.synchronize()

                start = torch.cuda.Event(enable_timing=True)
                end = torch.cuda.Event(enable_timing=True)
                start.record()
                for _ in range(3):
                    aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, sk)
                end.record()
                torch.cuda.synchronize()

                us = start.elapsed_time(end) * 1e3 / 3
                if us < best_us:
                    best_us = us
                    best_config = (kid, sk)
            except Exception:
                continue

    return best_config


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n = B_shuffle.shape[0]

    A_q, A_scale_sh = _quant_mxfp4_shuffled(A)
    A_flat = A_q.view(m, k // 2)

    shape_key = (m, n, k)
    if shape_key not in _TUNE_CACHE:
        _TUNE_CACHE[shape_key] = _tune_shape(A_flat, B_shuffle, A_scale_sh, B_scale_sh, m, n)

    config = _TUNE_CACHE[shape_key]

    padded_m = ((m + 31) // 32) * 32
    out = torch.empty((padded_m, n), dtype=dtypes.bf16, device="cuda")

    if config is None:
        # No blockscale kernels available — fall back to reference path
        return aiter.gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)

    kid, sk = config
    try:
        aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, sk)
        return out[:m]
    except RuntimeError:
        _TUNE_CACHE[shape_key] = None
        return aiter.gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
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