Skip to content
KernelIndex
Search⌘K

submission 616464

maxvel · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 291 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-616464?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
23.9µs
#817 of 1143
2026-03-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d07839ed2ce7de3ad8905521817d3e93bd5db3c800e2ac0d32549f901d8519e6
license declaredunknown
license concludedunknown
authorsmaxvel
imported2026-08-26

Techniques

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

fp4MI355X-oriented MXFP4 submission for `mxfp4-mm`.
split-kreleases mention ongoing FP4, A4W4 blockscale, split-K, and MI35x tuning work.

Kernel source

submission.py291 lines
"""
MI355X-oriented MXFP4 submission for `mxfp4-mm`.

Research summary embedded here to keep all implementation details in this file:
- The input already includes `B_shuffle` and `B_scale_sh`, so the main hot path is
  quantizing `A` to MXFP4 and dispatching `aiter.gemm_a4w4`.
- The AMD evaluator benchmarks the same shape repeatedly in one spawned process,
  which makes module-level lazy init, one-time warmup, and per-shape dispatch
  caching worthwhile.
- AMD public ROCm materials identify MI355X as `gfx950`, and recent ROCm/aiter
  releases mention ongoing FP4, A4W4 blockscale, split-K, and MI35x tuning work.

Tiered strategy implemented below:
- Tier 1: remove per-call imports/helper construction, avoid needless copies, and
  cache the selected execution path per shape.
- Tier 2: add hot-shape specializations plus one-time warmup and candidate GEMM
  probing.
- Tier 3: keep a safe framework for trying alternative A4W4 entry points when the
  local `aiter` build exposes them, while falling back to the stable baseline.
"""

from __future__ import annotations

from typing import Any, Callable
import inspect

from task import input_t, output_t


_HOT_M_VALUES = frozenset((4, 8, 16, 32, 64, 256))
_HOT_K_VALUES = frozenset((512, 1536, 2048, 7168))

_RUNTIME: dict[str, Any] | None = None
_DISPATCH_CACHE: dict[tuple[Any, ...], Callable[..., output_t]] = {}
_WARMED_DEVICES: set[tuple[str, int | None]] = set()


def _supports_gemm_signature(fn: Callable[..., Any]) -> bool:
    try:
        params = inspect.signature(fn).parameters
    except (TypeError, ValueError):
        return False
    return "dtype" in params and "bpreshuffle" in params


def _init_runtime() -> dict[str, Any]:
    global _RUNTIME
    if _RUNTIME is not None:
        return _RUNTIME

    import torch
    import aiter
    from aiter import dtypes
    from aiter.ops.shuffle import shuffle_weight
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    candidate_names = (
        "gemm_a4w4",
        "gemm_a4w4_tuned",
        "gemm_a4w4_blockscale",
        "ck_gemm_a4w4_blockscale",
        "asm_gemm_a4w4",
    )
    candidate_specs: list[tuple[str, Callable[..., Any]]] = []
    seen_ids: set[int] = set()

    for name in candidate_names:
        fn = getattr(aiter, name, None)
        if fn is None or not callable(fn):
            continue
        if id(fn) in seen_ids:
            continue
        if name != "gemm_a4w4" and not _supports_gemm_signature(fn):
            continue
        candidate_specs.append((name, fn))
        seen_ids.add(id(fn))

    if not candidate_specs:
        raise RuntimeError("No usable A4W4 GEMM entry point found in aiter")

    _RUNTIME = {
        "torch": torch,
        "aiter": aiter,
        "dtypes": dtypes,
        "shuffle_weight": shuffle_weight,
        "dynamic_mxfp4_quant": dynamic_mxfp4_quant,
        "e8m0_shuffle": e8m0_shuffle,
        "gemm_candidates": tuple(candidate_specs),
    }
    return _RUNTIME


def _quant_mxfp4(x, runtime: dict[str, Any]):
    x_fp4, bs_e8m0 = runtime["dynamic_mxfp4_quant"](x)
    bs_e8m0 = runtime["e8m0_shuffle"](bs_e8m0)
    dtypes = runtime["dtypes"]
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def _maybe_contiguous(x):
    return x if x.is_contiguous() else x.contiguous()


def _candidate_score(name: str, hot_shape: bool) -> int:
    score = 0
    lowered = name.lower()
    if hot_shape:
        if "tuned" in lowered:
            score += 8
        if "asm" in lowered:
            score += 6
        if "ck" in lowered or "blockscale" in lowered:
            score += 4
    if name == "gemm_a4w4":
        score += 2
    return score


def _select_gemm(runtime: dict[str, Any], m: int, k: int) -> Callable[..., Any]:
    hot_shape = m in _HOT_M_VALUES and k in _HOT_K_VALUES
    best_name, best_fn = max(
        runtime["gemm_candidates"],
        key=lambda item: _candidate_score(item[0], hot_shape),
    )
    runtime["selected_gemm_name"] = best_name
    return best_fn


def _maybe_warmup(runtime: dict[str, Any], device) -> None:
    device_key = (device.type, device.index)
    if device_key in _WARMED_DEVICES:
        return

    _WARMED_DEVICES.add(device_key)
    if device.type != "cuda":
        return

    torch = runtime["torch"]
    try:
        dummy = torch.zeros((64, 64), dtype=torch.bfloat16, device=device)
        q_a, scale_a = _quant_mxfp4(dummy, runtime)
        q_b, scale_b = _quant_mxfp4(dummy, runtime)
        b_shuffle = runtime["shuffle_weight"](q_b, layout=(16, 16))
        _select_gemm(runtime, 64, 64)(  # Compile and cache the common path once.
            q_a,
            b_shuffle,
            scale_a,
            scale_b,
            dtype=runtime["dtypes"].bf16,
            bpreshuffle=True,
        )
        torch.cuda.synchronize(device)
    except Exception:
        # Warmup is opportunistic only; never fail the submission because of it.
        return


def _run_gemm(
    runtime: dict[str, Any],
    gemm_fn: Callable[..., Any],
    A,
    B_shuffle,
    B_scale_sh,
):
    A_q, A_scale_sh = _quant_mxfp4(_maybe_contiguous(A), runtime)
    return gemm_fn(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=runtime["dtypes"].bf16,
        bpreshuffle=True,
    )


def _run_hot_small(
    runtime: dict[str, Any],
    gemm_fn: Callable[..., Any],
    A,
    B_shuffle,
    B_scale_sh,
):
    return _run_gemm(
        runtime,
        gemm_fn,
        A,
        B_shuffle,
        B_scale_sh,
    )


def _run_hot_medium(
    runtime: dict[str, Any],
    gemm_fn: Callable[..., Any],
    A,
    B_shuffle,
    B_scale_sh,
):
    B_shuffle = _maybe_contiguous(B_shuffle)
    B_scale_sh = _maybe_contiguous(B_scale_sh)
    return _run_gemm(runtime, gemm_fn, A, B_shuffle, B_scale_sh)


def _run_hot_large(
    runtime: dict[str, Any],
    gemm_fn: Callable[..., Any],
    A,
    B_shuffle,
    B_scale_sh,
):
    A = _maybe_contiguous(A)
    B_shuffle = _maybe_contiguous(B_shuffle)
    B_scale_sh = _maybe_contiguous(B_scale_sh)
    return _run_gemm(runtime, gemm_fn, A, B_shuffle, B_scale_sh)


def _run_generic(
    runtime: dict[str, Any],
    gemm_fn: Callable[..., Any],
    A,
    B_shuffle,
    B_scale_sh,
):
    return _run_gemm(
        runtime,
        gemm_fn,
        A,
        _maybe_contiguous(B_shuffle),
        _maybe_contiguous(B_scale_sh),
    )


def _select_dispatch(m: int, n: int, k: int) -> Callable[..., output_t]:
    hot_shape = m in _HOT_M_VALUES and k in _HOT_K_VALUES
    if hot_shape and m <= 32 and n >= 2048:
        return _run_hot_small
    if hot_shape and m >= 256:
        return _run_hot_large
    if hot_shape:
        return _run_hot_medium
    return _run_generic


def _build_executor(
    runtime: dict[str, Any],
    m: int,
    n: int,
    k: int,
) -> Callable[..., output_t]:
    dispatch = _select_dispatch(m, n, k)
    gemm_fn = _select_gemm(runtime, m, k)

    def _executor(runtime: dict[str, Any], A, B_shuffle, B_scale_sh):
        return dispatch(runtime, gemm_fn, A, B_shuffle, B_scale_sh)

    return _executor


def custom_kernel(data: input_t) -> output_t:
    """
    Optimized path: quantize `A` to MXFP4 once per call, then invoke the selected
    A4W4 GEMM entry point with preshuffled `B`.
    """
    runtime = _init_runtime()

    A, B, _B_q, B_shuffle, B_scale_sh = data
    del _B_q

    m, k = A.shape
    n = B.shape[0]

    _maybe_warmup(runtime, A.device)

    shape_key = (
        m,
        n,
        k,
        tuple(A.stride()),
        tuple(B_shuffle.stride()),
        A.device.type,
        A.device.index,
        A.dtype,
    )
    executor = _DISPATCH_CACHE.get(shape_key)
    if executor is None:
        executor = _build_executor(runtime, m, n, k)
        _DISPATCH_CACHE[shape_key] = executor

    return executor(runtime, A, B_shuffle, B_scale_sh)
scrolls · 291 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