Skip to content
KernelIndex
Search⌘K

submission 739858

gxtzhuxi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mxfp4_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-739858?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
#823 of 1143
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7262cd01e8f6c8664b84f50b8c1c66cba06078ee79405256047de31fe89cbcfc
license declaredunknown
license concludedunknown
authorsgxtzhuxi
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM — bf16 A × MXFP4 B → bf16 C on MI355X (CDNA4).

Kernel source

mxfp4_gemm.py154 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
MXFP4 GEMM — bf16 A × MXFP4 B → bf16 C on MI355X (CDNA4).

== Architecture ==
  MI355X: 256 CUs, 8 XCDs, 64KB LDS/CU, 32MB L2, 256MB Infinity Cache,
          8 TB/s HBM3E, FP4 MFMA (10 PFLOPS).

== Explicit Kernel Control ==
  Calls gemm_a4w4_asm DIRECTLY with pre-allocated output buffer,
  bypassing the gemm_a4w4 wrapper to eliminate per-call allocation
  and kernel-selection overhead.

  Pipeline (3 kernel launches):
    1. dynamic_mxfp4_quant(A) → A_fp4, A_scale   [Triton kernel]
    2. e8m0_shuffle(A_scale)  → A_scale_shuffled  [CK MFMA layout]
    3. gemm_a4w4_asm(...)     → C                 [FP4 MFMA ASM]

  First call per shape uses gemm_a4w4 wrapper for JIT warmup and
  ASM kernel-name discovery (verified against reference). Subsequent
  calls bypass the wrapper with pre-allocated output + padded A buffers.
"""

import torch
import torch.nn.functional as F
from task import input_t, output_t

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


# ===== Global state =====
_CACHE = {}
_DIRECT_FNS = {}

_KERNEL_PREFIX = "f4gemm_bf16_per1x32Fp4_BpreShuffle_"
_KNOWN_TILES = [32, 192]


def _discover_direct_config(M, N, K, A_q, B_shuffle, A_sc, B_scale_sh,
                            reference, device):
    """Try CK ASM kernel names to find one matching the wrapper's output.

    Verifies correctness with torch.equal before committing. On success,
    pre-allocates output + A-padding buffers for the direct path.
    """
    asm = _DIRECT_FNS.get("asm")
    get_pm = _DIRECT_FNS.get("get_padded_m")
    if asm is None or get_pm is None:
        return None

    padded_m = get_pm(M, N, K, 1)

    if padded_m > M:
        a_pad = F.pad(A_q.view(torch.uint8),
                      (0, 0, 0, padded_m - M)).view(dtypes.fp4x2)
        sc_pad = F.pad(A_sc.view(torch.uint8),
                       (0, 0, 0, padded_m - M)).view(dtypes.fp8_e8m0)
    else:
        a_pad = A_q
        sc_pad = A_sc

    out = torch.empty(padded_m, N, dtype=torch.bfloat16, device=device)

    candidates = sorted(set([padded_m] + _KNOWN_TILES), reverse=True)
    for tile_m in candidates:
        kname = f"{_KERNEL_PREFIX}{tile_m}x128"
        for log2_ks in [None, 0, 1, 2]:
            try:
                out.zero_()
                asm(a_pad, B_shuffle, sc_pad, B_scale_sh, out, kname,
                    bpreshuffle=True, log2_k_split=log2_ks)
                if torch.equal(out[:M, :], reference):
                    K_half = A_q.view(torch.uint8).shape[1]
                    K_sc = A_sc.view(torch.uint8).shape[1]
                    return {
                        "kname": kname,
                        "log2_ks": log2_ks,
                        "padded_m": padded_m,
                        "needs_pad": padded_m > M,
                        "out": out,
                        "A_q_buf": torch.zeros(padded_m, K_half,
                                               dtype=torch.uint8, device=device),
                        "A_sc_buf": torch.zeros(padded_m, K_sc,
                                                dtype=torch.uint8, device=device),
                    }
            except Exception:
                continue
    return None


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.shape[0]
    device = A.device

    # Stage 1: Quantize A (dynamic activations — every call)
    A_fp4, A_scale = dynamic_mxfp4_quant(A)

    # Stage 2: Shuffle A's E8M0 scale to CK's MFMA-friendly layout
    A_scale_sh = e8m0_shuffle(A_scale)

    A_q = A_fp4.view(dtypes.fp4x2)
    A_sc = A_scale_sh.view(dtypes.fp8_e8m0)

    key = (M, N, K)

    if key not in _CACHE:
        # JIT warmup: use wrapper for reference result + kernel auto-selection
        ref = aiter.gemm_a4w4(
            A_q, B_shuffle, A_sc, B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True,
        )
        if "asm" not in _DIRECT_FNS:
            _DIRECT_FNS["asm"] = getattr(aiter, "gemm_a4w4_asm", None)
            _DIRECT_FNS["get_padded_m"] = getattr(aiter, "get_padded_m", None)

        _CACHE[key] = _discover_direct_config(
            M, N, K, A_q, B_shuffle, A_sc, B_scale_sh, ref, device,
        )
        return ref

    c = _CACHE[key]

    # === Direct ASM path: pre-allocated output, no wrapper overhead ===
    if c is not None:
        if c["needs_pad"]:
            c["A_q_buf"][:M].copy_(A_q.view(torch.uint8))
            c["A_sc_buf"][:M].copy_(A_sc.view(torch.uint8))
            a_in = c["A_q_buf"].view(dtypes.fp4x2)
            sc_in = c["A_sc_buf"].view(dtypes.fp8_e8m0)
        else:
            a_in = A_q
            sc_in = A_sc

        _DIRECT_FNS["asm"](
            a_in, B_shuffle, sc_in, B_scale_sh,
            c["out"], c["kname"],
            bpreshuffle=True, log2_k_split=c["log2_ks"],
        )
        return c["out"][:M, :]

    # === Wrapper fallback (if kernel discovery failed) ===
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_sc, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )
scrolls · 154 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