Skip to content
KernelIndex
Search⌘K

submission 714256

chu yifan · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:095d5923af617409447f818ded15e208a2d3080d7267cb9f0c0b5e7367aa8be0
license declaredunknown
license concludedunknown
authorschu yifan
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM submission.

Kernel source

submission.py423 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
MXFP4 GEMM submission.

Goal: minimize end-to-end latency of:
  bf16 A  -> (per-1x32) MXFP4 quant  ->  A4W4 GEMM (MXFP4 A x MXFP4 W)  ->  bf16 C

Strategy:
  - Prefer AITER Triton fused path (quantize-A inside GEMM) when available.
  - Otherwise fall back to the reference path (dynamic_mxfp4_quant + gemm_a4w4).
  - Auto-tune per (M,N,K) shape once per process to pick the faster path without
    guessing a fixed threshold.
"""
from task import input_t, output_t

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

try:
    from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import (
        gemm_a16wfp4_preshuffle_ as _gemm_a16wfp4_preshuffle_,
    )
    from aiter.ops.triton.utils.common_utils import serialize_dict as _serialize_dict
    from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
        _get_config as _get_gemm_a16wfp4_config,
    )
except Exception:
    _gemm_a16wfp4_preshuffle_ = None
    _serialize_dict = None
    _get_gemm_a16wfp4_config = None


_IMPL_CACHE: dict[tuple, str] = {}
_FUSED_CONFIG_CACHE: dict[tuple, str] = {}
_FUSED_PICK_CACHE: dict[tuple, str] = {}

_FUSED_CONFIG_OVERRIDES: dict[tuple[int, int, int], list[dict]] = {
    (4, 2880, 512): [
        {
            "BLOCK_SIZE_M": 4,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
        {
            "BLOCK_SIZE_M": 8,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 1,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
    ],
    (16, 2112, 7168): [
        {
            "BLOCK_SIZE_M": 16,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 14,
        },
        {
            "BLOCK_SIZE_M": 16,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 1,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 7,
        },
    ],
    (32, 4096, 512): [
        {
            "BLOCK_SIZE_M": 8,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 1,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
        {
            "BLOCK_SIZE_M": 16,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 2,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
    ],
    (32, 2880, 512): [
        {
            "BLOCK_SIZE_M": 8,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 1,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
        {
            "BLOCK_SIZE_M": 16,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 2,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
    ],
    (64, 7168, 2048): [
        {
            "BLOCK_SIZE_M": 16,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 2,
            "waves_per_eu": 4,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
        {
            "BLOCK_SIZE_M": 32,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 256,
            "GROUP_SIZE_M": 4,
            "num_warps": 8,
            "num_stages": 2,
            "waves_per_eu": 4,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        },
        {
            "BLOCK_SIZE_M": 32,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 2,
            "waves_per_eu": 4,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        },
    ],
}

_FORCED_FUSED_CONFIGS: dict[tuple[int, int, int], dict] = {
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 1,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 7,
    },
    (64, 7168, 2048): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 8,
        "num_stages": 2,
        "waves_per_eu": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (256, 3072, 1536): {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 256,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 8,
        "num_stages": 2,
        "waves_per_eu": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
}


def _quant_mxfp4(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def _get_fused_config_hash(m: int, n: int, k_packed: int) -> str:
    key = (m, n, k_packed)
    cached = _FUSED_CONFIG_CACHE.get(key)
    if cached is not None:
        return cached

    if _get_gemm_a16wfp4_config is None or _serialize_dict is None:
        raise RuntimeError("gemm_a16wfp4_preshuffle config helpers unavailable")

    config, _ = _get_gemm_a16wfp4_config(m, n, k_packed, True)
    config_hash = _serialize_dict(config)
    _FUSED_CONFIG_CACHE[key] = config_hash
    return config_hash


def _serialize_fused_config(config: dict) -> str:
    if _serialize_dict is None:
        raise RuntimeError("serialize_dict unavailable")
    return _serialize_dict(config)


def _get_fused_config_candidates(m: int, n: int, k: int) -> list[str]:
    candidates = [_get_fused_config_hash(m, n, k // 2)]
    for config in _FUSED_CONFIG_OVERRIDES.get((int(m), int(n), int(k)), []):
        candidates.append(_serialize_fused_config(config))

    deduped: list[str] = []
    seen: set[str] = set()
    for config_hash in candidates:
        if config_hash not in seen:
            deduped.append(config_hash)
            seen.add(config_hash)
    return deduped


def _run_fused(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
    config_hash: str | None = None,
) -> torch.Tensor:
    if _gemm_a16wfp4_preshuffle_ is None:
        raise RuntimeError("gemm_a16wfp4_preshuffle unavailable")

    m, k = A.shape
    n, k_packed = B_shuffle.shape

    w_preshuf = B_shuffle.view(torch.uint8).view(n // 16, k_packed * 16)

    scale_bytes = B_scale_sh.view(torch.uint8)
    w_scales_preshuf = scale_bytes.view(
        scale_bytes.shape[0] // 32,
        scale_bytes.shape[1] * 32,
    )

    if config_hash is None:
        config_hash = _get_fused_config_hash(m, n, k_packed)
    return _gemm_a16wfp4_preshuffle_(
        A,
        w_preshuf,
        w_scales_preshuf,
        prequant=True,
        dtype=dtypes.bf16,
        y=None,
        config=config_hash,
        skip_reduce=False,
    )


def _run_fallback(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
) -> torch.Tensor:
    A_q, A_scale_sh = _quant_mxfp4(A)
    return aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def _time_cuda_us(fn) -> float:
    torch.cuda.synchronize()
    start_event = torch.cuda.Event(enable_timing=True)
    end_event = torch.cuda.Event(enable_timing=True)
    start_event.record()
    out = fn()
    end_event.record()
    torch.cuda.synchronize()
    del out
    return float(start_event.elapsed_time(end_event) * 1e3)


def _pick_impl(
    *,
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
) -> str:
    m, k = A.shape
    n, _ = B_shuffle.shape
    key = (A.device, int(m), int(n), int(k))
    cached = _IMPL_CACHE.get(key)
    if cached is not None:
        return cached

    if _gemm_a16wfp4_preshuffle_ is None:
        _IMPL_CACHE[key] = "fallback"
        return "fallback"

    fused_candidates: list[str] = []
    for config_hash in _get_fused_config_candidates(m, n, k):
        try:
            _run_fused(A, B_shuffle, B_scale_sh, config_hash=config_hash)
            fused_candidates.append(config_hash)
        except Exception:
            continue
    if not fused_candidates:
        _IMPL_CACHE[key] = "fallback"
        return "fallback"

    # Warm-up both codepaths so one-time compilation/module loads don't skew timing.
    try:
        _run_fallback(A, B_shuffle, B_scale_sh)
    except Exception:
        # If fallback path fails for any reason, stick to fused (already verified).
        _IMPL_CACHE[key] = "fused"
        return "fused"

    torch.cuda.synchronize()

    # Choose the faster of fused vs (quant + a4w4) using a few short runs.
    best_fused_hash = min(
        fused_candidates,
        key=lambda config_hash: min(
            _time_cuda_us(
                lambda config_hash=config_hash: _run_fused(
                    A, B_shuffle, B_scale_sh, config_hash=config_hash
                )
            )
            for _ in range(3)
        ),
    )
    fused_us = min(
        _time_cuda_us(
            lambda: _run_fused(A, B_shuffle, B_scale_sh, config_hash=best_fused_hash)
        )
        for _ in range(3)
    )
    fallback_us = min(
        _time_cuda_us(lambda: _run_fallback(A, B_shuffle, B_scale_sh)) for _ in range(3)
    )
    chosen = "fused" if fused_us <= fallback_us else "fallback"
    _IMPL_CACHE[key] = chosen
    if chosen == "fused":
        _FUSED_PICK_CACHE[key] = best_fused_hash
    return chosen


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

    forced_fused = _FORCED_FUSED_CONFIGS.get((int(m), int(n), int(k)))
    if forced_fused is not None:
        return _run_fused(
            A,
            B_shuffle,
            B_scale_sh,
            config_hash=_serialize_fused_config(forced_fused),
        )

    key = (A.device, int(m), int(n), int(k))
    impl = _pick_impl(A=A, B_shuffle=B_shuffle, B_scale_sh=B_scale_sh)
    if impl == "fused":
        return _run_fused(
            A,
            B_shuffle,
            B_scale_sh,
            config_hash=_FUSED_PICK_CACHE.get(key),
        )
    return _run_fallback(A, B_shuffle, B_scale_sh)
scrolls · 423 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