Skip to content
KernelIndex
Search⌘K

submission 595640

wsxhjnb1 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_mxfp4_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-595640?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.8µs
#810 of 1143
2026-03-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7c8264e6ac0fd7fe0117df1061318ef6977ab7df81617fa2eaa4d128ce4015f9
license declaredunknown
license concludedunknown
authorswsxhjnb1
imported2026-08-26

Kernel source

submission_mxfp4_mm.py355 lines
from __future__ import annotations

import os
from typing import Tuple

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


# -----------------------------------------------------------------------------
# MXFP4 GEMM skeleton for AMD MI355X qualification round.
#
# Safe default:
#   patched dynamic_mxfp4_quant + aiter.gemm_a4w4(bpreshuffle=True)
#
# Where to optimize next:
#   1) _dispatch_bucket
#   2) _run_custom_bucket_kernel
# -----------------------------------------------------------------------------

SCALE_GROUP_SIZE = 32
_WARMED_BUCKETS: set[Tuple[str, int, int, int]] = set()


def _env_flag(name: str, default: bool = False) -> bool:
    value = os.getenv(name)
    if value is None:
        return default
    return value.lower() in {"1", "true", "yes", "on"}


def _env_str(name: str, default: str) -> str:
    value = os.getenv(name)
    return default if value is None else value


def _env_int(name: str, default: int) -> int:
    value = os.getenv(name)
    if value is None or value == "":
        return default
    try:
        return int(value)
    except ValueError:
        return default


MM_ENABLE_WARMUP = _env_flag("MM_ENABLE_WARMUP", True)
MM_WARMUP_SYNC = _env_flag("MM_WARMUP_SYNC", False)
MM_WARMUP_SCOPE = _env_str("MM_WARMUP_SCOPE", "bucket_k")

MM_CONTIG_MODE = _env_str("MM_CONTIG_MODE", "never")
MM_BSHUFFLE_CONTIG_MODE = _env_str("MM_BSHUFFLE_CONTIG_MODE", "never")
MM_BSCALE_CONTIG_MODE = _env_str("MM_BSCALE_CONTIG_MODE", "never")

MM_IMPL = _env_str("MM_IMPL", "aiter_ck")
MM_FORCE_TORCH_DEBUG = _env_flag("MM_FORCE_TORCH_DEBUG", False)

MM_ASM_KERNEL_SMALL_K_ALIGNED_N = _env_str("MM_ASM_KERNEL_SMALL_K_ALIGNED_N", "")
MM_ASM_KERNEL_SMALL_K_UNALIGNED_N = _env_str("MM_ASM_KERNEL_SMALL_K_UNALIGNED_N", "")
MM_ASM_KERNEL_VERY_LARGE_K = _env_str("MM_ASM_KERNEL_VERY_LARGE_K", "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E")
MM_ASM_KERNEL_THIN_M_WIDE_N = _env_str("MM_ASM_KERNEL_THIN_M_WIDE_N", "")
MM_ASM_KERNEL_LARGE_M = _env_str("MM_ASM_KERNEL_LARGE_M", "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E")
MM_ASM_KERNEL_DEFAULT = _env_str("MM_ASM_KERNEL_DEFAULT", "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E")

MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N = _env_int("MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N", 0)
MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N = _env_int("MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N", 0)
MM_ASM_LOG2_SPLIT_VERY_LARGE_K = _env_int("MM_ASM_LOG2_SPLIT_VERY_LARGE_K", 1)
MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N = _env_int("MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N", 0)
MM_ASM_LOG2_SPLIT_LARGE_M = _env_int("MM_ASM_LOG2_SPLIT_LARGE_M", 0)
MM_ASM_LOG2_SPLIT_DEFAULT = _env_int("MM_ASM_LOG2_SPLIT_DEFAULT", 0)

MM_FAST_CK_PATH = (
    MM_IMPL == "aiter_ck"
    and not (
        MM_ASM_KERNEL_SMALL_K_ALIGNED_N
        or MM_ASM_KERNEL_SMALL_K_UNALIGNED_N
        or MM_ASM_KERNEL_VERY_LARGE_K
        or MM_ASM_KERNEL_THIN_M_WIDE_N
        or MM_ASM_KERNEL_LARGE_M
        or MM_ASM_KERNEL_DEFAULT
    )
    and MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N == 0
    and MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N == 0
    and MM_ASM_LOG2_SPLIT_VERY_LARGE_K == 0
    and MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N == 0
    and MM_ASM_LOG2_SPLIT_LARGE_M == 0
    and MM_ASM_LOG2_SPLIT_DEFAULT == 0
)


def _maybe_contiguous(x: torch.Tensor, mode: str) -> torch.Tensor:
    if mode == "always":
        return x.contiguous()
    if mode == "never":
        return x
    return x if x.is_contiguous() else x.contiguous()


def _shape_bucket(m: int, n: int, k: int) -> str:
    # Tuned for the public benchmark regimes.
    if k <= 512:
        if n % 512 == 0:
            return "small_k_aligned_n"
        return "small_k_unaligned_n"
    if k >= 4096:
        return "very_large_k"
    if m <= 32 and n >= 4096:
        return "thin_m_wide_n"
    if m >= 128:
        return "large_m"
    return "default"


def _quant_mxfp4_shuffled(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    # Important: this uses the patched kernel from aiter.ops.triton.quant.
    x_fp4, scale_e8m0 = dynamic_mxfp4_quant(x)
    scale_e8m0 = e8m0_shuffle(scale_e8m0)
    return x_fp4.view(dtypes.fp4x2), scale_e8m0.view(dtypes.fp8_e8m0)


def _run_aiter_ck(
    a_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    return aiter.gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def _run_aiter_asm(
    a_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_scale_sh: torch.Tensor,
    kernel_name: str = "",
    log2_k_split: int | None = None,
) -> torch.Tensor:
    m = a_q.shape[0]
    n = b_shuffle.shape[0]
    out = torch.empty(((m + 31) // 32) * 32, n, dtype=dtypes.bf16, device=a_q.device)
    split = log2_k_split if log2_k_split is not None and log2_k_split > 0 else None
    aiter.gemm_a4w4_asm(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        out,
        kernel_name,
        None,
        bpreshuffle=True,
        log2_k_split=split,
    )
    return out[:m]


def _warmup_key(bucket: str, m: int, n: int, k: int) -> tuple:
    if MM_WARMUP_SCOPE == "bucket":
        return (bucket,)
    if MM_WARMUP_SCOPE == "bucket_k":
        return (bucket, k)
    return (bucket, m, n, k)


def _default_asm_kernel(bucket: str) -> str:
    if bucket == "small_k_aligned_n":
        return MM_ASM_KERNEL_SMALL_K_ALIGNED_N
    if bucket == "small_k_unaligned_n":
        return MM_ASM_KERNEL_SMALL_K_UNALIGNED_N
    if bucket == "very_large_k":
        return MM_ASM_KERNEL_VERY_LARGE_K
    if bucket == "thin_m_wide_n":
        return MM_ASM_KERNEL_THIN_M_WIDE_N
    if bucket == "large_m":
        return MM_ASM_KERNEL_LARGE_M
    return MM_ASM_KERNEL_DEFAULT


def _default_asm_split(bucket: str) -> int:
    if bucket == "small_k_aligned_n":
        return MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N
    if bucket == "small_k_unaligned_n":
        return MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N
    if bucket == "very_large_k":
        return MM_ASM_LOG2_SPLIT_VERY_LARGE_K
    if bucket == "thin_m_wide_n":
        return MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N
    if bucket == "large_m":
        return MM_ASM_LOG2_SPLIT_LARGE_M
    return MM_ASM_LOG2_SPLIT_DEFAULT


def _resolve_asm_spec(bucket: str) -> tuple[str, int]:
    return _default_asm_kernel(bucket), _default_asm_split(bucket)


def _run_custom_bucket_kernel(
    bucket: str,
    a: torch.Tensor,
    a_q: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    # TODO(user): replace this function with your bucket-specific Triton / Helion /
    # asm implementation. Keep the signature stable so the dispatcher and warmup
    # logic do not need to change.
    _ = (bucket, a)
    kernel_name, log2_k_split = _resolve_asm_spec(bucket)
    return _run_aiter_asm(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        kernel_name=kernel_name,
        log2_k_split=log2_k_split,
    )


def _dispatch_bucket(
    a: torch.Tensor,
    a_q: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    if MM_FAST_CK_PATH:
        return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)

    m, k = a.shape
    n = b_shuffle.shape[0]
    bucket = _shape_bucket(m, n, k)
    impl = MM_IMPL
    kernel_name, log2_k_split = _resolve_asm_spec(bucket)

    if impl == "custom":
        return _run_custom_bucket_kernel(bucket, a, a_q, a_scale_sh, b_shuffle, b_scale_sh)
    if kernel_name:
        return _run_aiter_asm(
            a_q,
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            kernel_name=kernel_name,
            log2_k_split=log2_k_split,
        )
    if impl == "aiter_asm":
        return _run_aiter_asm(
            a_q,
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            kernel_name="",
            log2_k_split=log2_k_split,
        )

    return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)


def _run_selected_impl(
    a: torch.Tensor,
    a_q: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    if MM_IMPL == "aiter_ck":
        m, k = a.shape
        n = b_shuffle.shape[0]
        bucket = _shape_bucket(m, n, k)
        kernel_name, log2_k_split = _resolve_asm_spec(bucket)
        if kernel_name or log2_k_split > 0:
            return _run_aiter_asm(
                a_q,
                b_shuffle,
                a_scale_sh,
                b_scale_sh,
                kernel_name=kernel_name,
                log2_k_split=log2_k_split,
            )
        return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)

    return _dispatch_bucket(a, a_q, a_scale_sh, b_shuffle, b_scale_sh)


def _maybe_warmup(
    a: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> None:
    if not MM_ENABLE_WARMUP:
        return

    m, k = a.shape
    n = b_shuffle.shape[0]
    bucket = _warmup_key(_shape_bucket(m, n, k), m, n, k)
    if bucket in _WARMED_BUCKETS:
        return

    a_q, a_scale_sh = _quant_mxfp4_shuffled(a)
    if MM_FAST_CK_PATH:
        _ = _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)
    else:
        _ = _run_selected_impl(a, a_q, a_scale_sh, b_shuffle, b_scale_sh)
    if MM_WARMUP_SYNC:
        torch.cuda.synchronize()
    _WARMED_BUCKETS.add(bucket)


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    """
    Expected input:
        (A, B, B_q, B_shuffle, B_scale_sh)

    Timed-path rules for this task:
      - do NOT reshuffle B inside custom_kernel
      - do NOT rebuild B_scale_sh inside custom_kernel
      - keep Python-side branching minimal
    """
    a, b, b_q, b_shuffle, b_scale_sh = data

    # Only A is needed on the timed path. Keep B/B_q available for debug checks.
    bshuffle_contig_mode = MM_BSHUFFLE_CONTIG_MODE if MM_BSHUFFLE_CONTIG_MODE != "if_needed" else MM_CONTIG_MODE
    bscale_contig_mode = MM_BSCALE_CONTIG_MODE if MM_BSCALE_CONTIG_MODE != "if_needed" else MM_CONTIG_MODE

    a = _maybe_contiguous(a, MM_CONTIG_MODE)
    b_shuffle = _maybe_contiguous(b_shuffle, bshuffle_contig_mode)
    b_scale_sh = _maybe_contiguous(b_scale_sh, bscale_contig_mode)

    _maybe_warmup(a, b_shuffle, b_scale_sh)

    a_q, a_scale_sh = _quant_mxfp4_shuffled(a)

    if MM_FORCE_TORCH_DEBUG:
        # Kept as a non-failing knob so you can switch on extra instrumentation
        # in local experiments without changing the submission surface.
        _ = (b, b_q)

    if MM_FAST_CK_PATH:
        return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)
    return _run_selected_impl(a, a_q, a_scale_sh, b_shuffle, b_scale_sh)
scrolls · 355 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