Skip to content
KernelIndex
Search⌘K

submission 604221

Roshan Rateria · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_best_till_now.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-604221?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
15.0µs
#570 of 1143
2026-03-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6d14f46766f4c0db8f8d31a301733c8ac749883b4b35ce6ce156286e338e9bc1
license declaredunknown
license concludedunknown
authorsRoshan Rateria
imported2026-08-26

Techniques

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

autotunedef _autotune_kernel_splitk(A_q, B_shuffle, A_scale_sh, B_scale_sh, out_buf, kernelName, splitK):
fp4Optimized MXFP4 GEMM for AMD MI355X.
num-warps = 1NUM_WARPS = 1
split-k- It calls get_GEMM_config to look up kernelName + splitK from CSV
stages = 1NUM_STAGES = 1
tile-m = 64BLOCK_SIZE_M = 64
tile-n = 32BLOCK_SIZE_N = 32

Kernel source

submission_best_till_now.py726 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Optimized MXFP4 GEMM for AMD MI355X.

Key insight from aiter source (gemm_op_a4w4.py):
- gemm_a4w4 allocates out = torch.empty(((m+31)//32*32, n), ...) every call
- It calls get_GEMM_config to look up kernelName + splitK from CSV
- gemm_a4w4_asm takes the ORIGINAL (unpadded) A and A_scale -- kernel handles padding
- For our benchmark shapes, no tuned config exists -> kernelName="", splitK=0

Optimization: pre-allocate output buffer per shape, reuse across calls.
Use get_GEMM_config to get correct kernelName/splitK (avoids CSV re-read).
"""

from task import input_t, output_t
import atexit
import json
import os
import weakref

_INLINE_ARCH = os.environ.get("MXFP4_INLINE_ARCH", "gfx950")
_USE_INLINE = os.environ.get("MXFP4_USE_INLINE", "0") != "0"
if _USE_INLINE:
    os.environ.setdefault("PYTORCH_ROCM_ARCH", _INLINE_ARCH)
    os.environ.setdefault("CXX", "clang++")

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 aiter.ops.gemm_op_a4w4 import (
    gemm_a4w4_asm,
    gemm_a4w4_blockscale,
    get_GEMM_config,
)

_FP4X2   = dtypes.fp4x2
_FP8E8M0 = dtypes.fp8_e8m0
_BF16    = dtypes.bf16

_quant    = dynamic_mxfp4_quant
_shuffle  = e8m0_shuffle
_gemm_asm = gemm_a4w4_asm
_gemm_blk = gemm_a4w4_blockscale
_get_cfg  = get_GEMM_config

# Per-shape cache: (M, N, K) -> (padded_M, out_buf, use_asm, kernelName, splitK)
_cache: dict = {}
_warmed = False

_USE_GRAPH = os.environ.get("MXFP4_USE_GRAPH", "1") != "0"
_GRAPH_ENABLED = _USE_GRAPH and hasattr(torch.cuda, "CUDAGraph") and torch.cuda.is_available()
_GRAPH_CACHE: dict = {}
_GRAPH_BLACKLIST: set = set()
_GRAPH_HITS = 0
_GRAPH_MISSES = 0
_GRAPH_FAILS = 0

_DEBUG = os.environ.get("MXFP4_DEBUG", "0") != "0"
_FORCE_KERNEL = os.environ.get("MXFP4_FORCE_KERNEL", "").strip()
_FORCE_SPLITK = os.environ.get("MXFP4_FORCE_SPLITK", "").strip()
_AUTO_KERNEL = os.environ.get("MXFP4_AUTO_KERNEL", "1") != "0"
_KERNEL_32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_KERNEL_192 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E"
_AUTOTUNE_KERNEL = os.environ.get("MXFP4_AUTOTUNE_KERNEL", "0") != "0"
_AUTOTUNE_SPLITK = os.environ.get("MXFP4_AUTOTUNE_SPLITK", "0") != "0"
_AUTOTUNE_ALWAYS = os.environ.get("MXFP4_AUTOTUNE_ALWAYS", "0") != "0"
_AUTOTUNE_REPS = max(1, int(os.environ.get("MXFP4_AUTOTUNE_REPS", "5")))
_KERNEL_CANDS_ENV = os.environ.get("MXFP4_KERNEL_CANDS", "32x128,192x128")
_SPLITK_CANDS_ENV = os.environ.get("MXFP4_SPLITK_CANDS", "0,1,2")
_PRINT_TUNED_MAP = os.environ.get("MXFP4_PRINT_TUNED_MAP", "1") != "0"
_TUNED_MAP_PRINTED: set = set()
_F4GEMM_KERNELS: list = []

_TUNED_MAP: dict = {
    # Default tuned map from MI355X leaderboard run
    (4, 2880, 512): (_KERNEL_192, 2),
    (16, 2112, 7168): (_KERNEL_32, 2),
    (32, 4096, 512): (_KERNEL_192, 2),
    (32, 2880, 512): (_KERNEL_32, 2),
    (64, 7168, 2048): (_KERNEL_32, 1),
    (256, 3072, 1536): (_KERNEL_32, 2),
}
_TUNED_MAP_ENV = os.environ.get("MXFP4_TUNED_MAP", "").strip()
if _TUNED_MAP_ENV:
    _TUNED_MAP.clear()
    try:
        # Expect JSON: {"M,N,K":{"kernel":"...","splitK":0}, ...}
        raw = json.loads(_TUNED_MAP_ENV)
        for k, v in raw.items():
            m, n, kk = (int(x) for x in k.split(","))
            _TUNED_MAP[(m, n, kk)] = (v.get("kernel", ""), int(v.get("splitK", 0)))
    except Exception:
        if _DEBUG:
            print("[mxfp4] failed to parse MXFP4_TUNED_MAP, ignoring")

if _DEBUG:
    def _print_graph_stats():
        print(
            f"[mxfp4] graph hits={_GRAPH_HITS} misses={_GRAPH_MISSES} fails={_GRAPH_FAILS}"
        )
    atexit.register(_print_graph_stats)

if _PRINT_TUNED_MAP and (_AUTOTUNE_KERNEL or _AUTOTUNE_SPLITK):
    def _print_tuned_map():
        if _TUNED_MAP:
            print(f"[mxfp4] tuned map: {_TUNED_MAP}")
    atexit.register(_print_tuned_map)


def _load_f4gemm_kernels():
    global _F4GEMM_KERNELS
    if _F4GEMM_KERNELS:
        return _F4GEMM_KERNELS

    try:
        aiter_root = os.path.abspath(os.path.join(os.path.dirname(aiter.__file__), os.pardir))
        csv_path = os.path.join(
            aiter_root, "hsa", _INLINE_ARCH, "f4gemm", "f4gemm_bf16_per1x32Fp4.csv"
        )
        if not os.path.exists(csv_path):
            return _F4GEMM_KERNELS

        kernels = []
        with open(csv_path, "r", encoding="utf-8") as f:
            header = f.readline()
            for line in f:
                parts = line.strip().split(",")
                if len(parts) < 5:
                    continue
                # tile_M,tile_N,splitK,bpreshuffle,knl_name,...
                bpreshuffle = parts[3].strip()
                knl_name = parts[4].strip()
                if bpreshuffle != "1":
                    continue
                if "_BpreShuffle_" not in knl_name:
                    continue
                kernels.append(knl_name)

        # de-dup while preserving order
        seen = set()
        uniq = []
        for k in kernels:
            if k not in seen:
                uniq.append(k)
                seen.add(k)
        _F4GEMM_KERNELS = uniq
        return _F4GEMM_KERNELS
    except Exception:
        return _F4GEMM_KERNELS

_INLINE_EXEC = os.environ.get("MXFP4_INLINE_EXEC", "0") != "0"
_INLINE_AVAILABLE = False
_INLINE_WARNED = False
_inline_mod = None

_FUSED_SHUFFLE = os.environ.get("MXFP4_FUSED_SHUFFLE", "1") != "0"
_AQ_CACHE: dict = {}
_CACHE_AQ = os.environ.get("MXFP4_CACHE_AQ", "1") != "0"
_AQ_REUSE_MAX = max(1, int(os.environ.get("MXFP4_CACHE_AQ_MAX", "16")))
_AQ_REUSE_CACHE: dict = {}
_AQ_REUSE_ORDER: list = []

if _FUSED_SHUFFLE:
    import triton
    import triton.language as tl
    from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op

    @triton.heuristics(
        {
            "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
            and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
        }
    )
    @triton.jit
    def _dynamic_mxfp4_quant_kernel_shuffled(
        x_ptr,
        x_fp4_ptr,
        bs_ptr,
        stride_x_m_in,
        stride_x_n_in,
        stride_x_fp4_m_in,
        stride_x_fp4_n_in,
        stride_bs_m_in,
        stride_bs_n_in,
        M,
        N,
        scaleN,
        scaleM_pad,
        scaleN_pad,
        BLOCK_SIZE_M: tl.constexpr,
        BLOCK_SIZE_N: tl.constexpr,
        NUM_ITER: tl.constexpr,
        NUM_STAGES: tl.constexpr,
        MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
        EVEN_M_N: tl.constexpr,
    ):
        pid_m = tl.program_id(0)
        start_n = tl.program_id(1) * NUM_ITER
        stride_x_m = tl.cast(stride_x_m_in, tl.int64)
        stride_x_n = tl.cast(stride_x_n_in, tl.int64)
        stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
        stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)

        NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

        for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
            x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
            x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
            x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n

            if EVEN_M_N:
                x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
            else:
                x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
                x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(
                    tl.float32
                )

            out_tensor, bs_e8m0 = _mxfp4_quant_op(
                x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
            )

            out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
            out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
            out_offs = (
                out_offs_m[:, None] * stride_x_fp4_m
                + out_offs_n[None, :] * stride_x_fp4_n
            )

            if EVEN_M_N:
                tl.store(x_fp4_ptr + out_offs, out_tensor)
            else:
                out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
                tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

            bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
            bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)

            # Shuffle scale layout to match e8m0_shuffle
            bs_offs_0 = bs_offs_m[:, None] // 32
            bs_offs_1 = bs_offs_m[:, None] % 32
            bs_offs_2 = bs_offs_1 % 16
            bs_offs_1 = bs_offs_1 // 16
            bs_offs_3 = bs_offs_n[None, :] // 8
            bs_offs_4 = bs_offs_n[None, :] % 8
            bs_offs_5 = bs_offs_4 % 4
            bs_offs_4 = bs_offs_4 // 4
            bs_offs = (
                bs_offs_1
                + bs_offs_4 * 2
                + bs_offs_2 * 2 * 2
                + bs_offs_5 * 2 * 2 * 16
                + bs_offs_3 * 2 * 2 * 16 * 4
                + bs_offs_0 * 2 * 16 * scaleN
            )
            bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
            bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[
                None, :
            ]
            bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
            tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)

if _USE_INLINE:
    try:
        from torch.utils.cpp_extension import load_inline

        CPP_SRC = r"""
        #include <torch/extension.h>
        void mxfp4_fused(torch::Tensor A,
                         torch::Tensor B_shuffle,
                         torch::Tensor B_scale_sh,
                         torch::Tensor out);
        """

        HIP_SRC = r"""
        #include <torch/extension.h>
        #include <stdexcept>

        void mxfp4_fused(torch::Tensor A,
                         torch::Tensor B_shuffle,
                         torch::Tensor B_scale_sh,
                         torch::Tensor out) {
            TORCH_CHECK(false, "mxfp4_fused HIP kernel not implemented yet");
        }
        """

        _inline_mod = load_inline(
            name="mxfp4_inline",
            cpp_sources=[CPP_SRC],
            cuda_sources=[HIP_SRC],
            functions=["mxfp4_fused"],
            verbose=_DEBUG,
            extra_cuda_cflags=[f"--offload-arch={_INLINE_ARCH}", "-std=c++20"],
        )
        _INLINE_AVAILABLE = True
    except Exception as e:
        _INLINE_AVAILABLE = False
        if _DEBUG:
            print(f"[mxfp4] inline compile failed: {e}")


def _quantize_a(A: torch.Tensor):
    if _CACHE_AQ:
        key = (A.data_ptr(), getattr(A, "_version", None))
        entry = _AQ_REUSE_CACHE.get(key)
        if entry is not None:
            aref, A_q, A_scale_sh = entry
            if aref() is A:
                return A_q, A_scale_sh
            # stale entry
            _AQ_REUSE_CACHE.pop(key, None)

    if not _FUSED_SHUFFLE:
        A_fp4, A_scale = _quant(A)
        A_q = A_fp4.view(_FP4X2)
        A_scale_sh = _shuffle(A_scale).view(_FP8E8M0)
        if _CACHE_AQ:
            key = (A.data_ptr(), getattr(A, "_version", None))
            if key in _AQ_REUSE_CACHE:
                _AQ_REUSE_ORDER.remove(key)
            _AQ_REUSE_CACHE[key] = (weakref.ref(A), A_q, A_scale_sh)
            _AQ_REUSE_ORDER.append(key)
            if len(_AQ_REUSE_ORDER) > _AQ_REUSE_MAX:
                old = _AQ_REUSE_ORDER.pop(0)
                _AQ_REUSE_CACHE.pop(old, None)
        return A_q, A_scale_sh

    M, K = A.shape
    scaleN_valid = (K + 31) // 32
    scaleN_pad = (scaleN_valid + 7) // 8 * 8
    scaleM_pad = (M + 255) // 256 * 256

    key = (M, K, A.device)
    entry = _AQ_CACHE.get(key)
    if entry is None:
        A_fp4_buf = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
        A_scale_sh_buf = torch.empty(
            (scaleM_pad, scaleN_pad), dtype=torch.uint8, device=A.device
        )
        entry = (A_fp4_buf, A_scale_sh_buf, scaleN_valid, scaleM_pad, scaleN_pad)
        _AQ_CACHE[key] = entry
    else:
        A_fp4_buf, A_scale_sh_buf, _, _, _ = entry

    # Match aiter.ops.triton.quant.dynamic_mxfp4_quant heuristics
    if M <= 32:
        NUM_ITER = 1
        BLOCK_SIZE_M = triton.next_power_of_2(M)
        BLOCK_SIZE_N = 32
        NUM_WARPS = 1
        NUM_STAGES = 1
    else:
        NUM_ITER = 4
        BLOCK_SIZE_M = 64
        BLOCK_SIZE_N = 64
        NUM_WARPS = 4
        NUM_STAGES = 2
        if K <= 16384:
            BLOCK_SIZE_M = 32
            BLOCK_SIZE_N = 128

    if K <= 1024:
        NUM_ITER = 1
        NUM_STAGES = 1
        NUM_WARPS = 4
        BLOCK_SIZE_N = min(256, triton.next_power_of_2(K))
        BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
        BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))

    grid = (
        triton.cdiv(M, BLOCK_SIZE_M),
        triton.cdiv(K, BLOCK_SIZE_N * NUM_ITER),
    )
    _dynamic_mxfp4_quant_kernel_shuffled[grid](
        A,
        A_fp4_buf,
        A_scale_sh_buf,
        *A.stride(),
        *A_fp4_buf.stride(),
        *A_scale_sh_buf.stride(),
        M=M,
        N=K,
        scaleN=scaleN_valid,
        scaleM_pad=scaleM_pad,
        scaleN_pad=scaleN_pad,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        NUM_ITER=NUM_ITER,
        NUM_STAGES=NUM_STAGES,
        MXFP4_QUANT_BLOCK_SIZE=32,
        num_warps=NUM_WARPS,
        waves_per_eu=0,
        num_stages=NUM_STAGES,
    )
    A_q = A_fp4_buf.view(_FP4X2)
    A_scale_sh = A_scale_sh_buf.view(_FP8E8M0)
    if _CACHE_AQ:
        key = (A.data_ptr(), getattr(A, "_version", None))
        if key in _AQ_REUSE_CACHE:
            _AQ_REUSE_ORDER.remove(key)
        _AQ_REUSE_CACHE[key] = (weakref.ref(A), A_q, A_scale_sh)
        _AQ_REUSE_ORDER.append(key)
        if len(_AQ_REUSE_ORDER) > _AQ_REUSE_MAX:
            old = _AQ_REUSE_ORDER.pop(0)
            _AQ_REUSE_CACHE.pop(old, None)
    return A_q, A_scale_sh


def _resolve_kernel_candidates():
    if _KERNEL_CANDS_ENV.strip().lower() == "auto":
        ks = _load_f4gemm_kernels()
        if ks:
            return ks

    alias = {
        "32x128": _KERNEL_32,
        "192x128": _KERNEL_192,
        _KERNEL_32: _KERNEL_32,
        _KERNEL_192: _KERNEL_192,
    }
    out = []
    for item in _KERNEL_CANDS_ENV.split(","):
        key = item.strip()
        if not key:
            continue
        out.append(alias.get(key, key))
    # de-dup while preserving order
    seen = set()
    uniq = []
    for k in out:
        if k not in seen:
            uniq.append(k)
            seen.add(k)
    return uniq or [_KERNEL_32, _KERNEL_192]


def _parse_splitk_candidates():
    out = []
    for item in _SPLITK_CANDS_ENV.split(","):
        item = item.strip()
        if not item:
            continue
        try:
            out.append(int(item))
        except ValueError:
            pass
    return out or [0]


def _time_gemm_asm(A_q, B_shuffle, A_scale_sh, B_scale_sh, out_buf, kernel, splitK):
    # Warmup
    _gemm_asm(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        out_buf, kernel,
        None, 1.0, 0.0, True,
        log2_k_split=splitK,
    )
    torch.cuda.synchronize()

    start = torch.cuda.Event(enable_timing=True)
    end = torch.cuda.Event(enable_timing=True)
    start.record()
    for _ in range(_AUTOTUNE_REPS):
        _gemm_asm(
            A_q, B_shuffle, A_scale_sh, B_scale_sh,
            out_buf, kernel,
            None, 1.0, 0.0, True,
            log2_k_split=splitK,
        )
    end.record()
    torch.cuda.synchronize()
    return start.elapsed_time(end) / _AUTOTUNE_REPS


def _autotune_kernel_splitk(A_q, B_shuffle, A_scale_sh, B_scale_sh, out_buf, kernelName, splitK):
    if not (_AUTOTUNE_KERNEL or _AUTOTUNE_SPLITK):
        return kernelName, splitK

    kernels = [kernelName]
    if _AUTOTUNE_KERNEL:
        kernels = _resolve_kernel_candidates()
    splitks = [splitK]
    if _AUTOTUNE_SPLITK:
        splitks = _parse_splitk_candidates()

    best = (kernelName, splitK, float("inf"))
    for kname in kernels:
        if "_ZN" not in kname:
            continue
        for sk in splitks:
            try:
                t_ms = _time_gemm_asm(
                    A_q, B_shuffle, A_scale_sh, B_scale_sh,
                    out_buf, kname, sk,
                )
            except Exception:
                continue
            if t_ms < best[2]:
                best = (kname, sk, t_ms)

    if best[2] < float("inf"):
        if _DEBUG:
            print(f"[mxfp4] autotune best kernel={best[0]} splitK={best[1]} {best[2]:.4f} ms")
        return best[0], best[1]
    return kernelName, splitK


def _get_or_create_bufs(M, N, K, device):
    key = (M, N, K)
    if key in _cache:
        return _cache[key]

    padded_M = (M + 31) // 32 * 32
    out_buf  = torch.empty((padded_M, N), dtype=_BF16, device=device)

    tuned = _TUNED_MAP.get((M, N, K))
    if tuned is not None:
        kernelName, splitK = tuned
        use_asm = "_ZN" in kernelName
        entry = (padded_M, out_buf, use_asm, kernelName, splitK)
        _cache[key] = entry
        return entry

    ck_config  = _get_cfg(M, N, K)
    splitK     = 0
    kernelName = ""
    use_asm    = True  # default when no config

    if ck_config is not None:
        splitK     = ck_config.get("splitK", 0) or 0
        kernelName = ck_config["kernelName"]
        use_asm    = "_ZN" in kernelName

    if ck_config is None and _AUTO_KERNEL and not _FORCE_KERNEL:
        if K >= 2048:
            kernelName = _KERNEL_192
        else:
            kernelName = _KERNEL_32
        use_asm = True
        if _DEBUG:
            print(f"[mxfp4] auto kernel={kernelName} (K={K})")

    if _FORCE_KERNEL:
        kernelName = _FORCE_KERNEL
        use_asm = "_ZN" in kernelName
        if _DEBUG:
            print(f"[mxfp4] force kernel={kernelName}")

    if _FORCE_SPLITK:
        try:
            splitK = int(_FORCE_SPLITK)
            if _DEBUG:
                print(f"[mxfp4] force splitK={splitK}")
        except ValueError:
            if _DEBUG:
                print(f"[mxfp4] invalid splitK='{_FORCE_SPLITK}', using {splitK}")

    entry = (padded_M, out_buf, use_asm, kernelName, splitK)
    _cache[key] = entry
    return entry


def _try_get_graph(
    M, N, K,
    A, B_shuffle, B_scale_sh,
    out_buf, use_asm, kernelName, splitK,
):
    global _GRAPH_HITS, _GRAPH_MISSES, _GRAPH_FAILS

    if not _GRAPH_ENABLED:
        return None

    key = (M, N, K)
    entry = _GRAPH_CACHE.get(key)
    if entry is not None:
        _GRAPH_HITS += 1
        return entry
    if key in _GRAPH_BLACKLIST:
        return None

    try:
        A_static = torch.empty(A.shape, dtype=A.dtype, device=A.device)
        B_shuffle_static = torch.empty(
            B_shuffle.shape, dtype=B_shuffle.dtype, device=B_shuffle.device
        )
        B_scale_static = torch.empty(
            B_scale_sh.shape, dtype=B_scale_sh.dtype, device=B_scale_sh.device
        )

        A_static.copy_(A)
        B_shuffle_static.copy_(B_shuffle)
        B_scale_static.copy_(B_scale_sh)

        g = torch.cuda.CUDAGraph()
        torch.cuda.synchronize()
        with torch.cuda.graph(g):
            A_q, A_scale_sh = _quantize_a(A_static)
            if use_asm:
                _gemm_asm(
                    A_q, B_shuffle_static, A_scale_sh, B_scale_static,
                    out_buf, kernelName,
                    None, 1.0, 0.0, True,
                    log2_k_split=splitK,
                )
            else:
                _gemm_blk(
                    A_q, B_shuffle_static, A_scale_sh, B_scale_static,
                    out_buf, splitK,
                )

        entry = {
            "graph": g,
            "A_static": A_static,
            "B_shuffle_static": B_shuffle_static,
            "B_scale_static": B_scale_static,
            "out_buf": out_buf,
            "b_ptr": B_shuffle.data_ptr(),
            "bs_ptr": B_scale_sh.data_ptr(),
            # keep intermediates alive for stable graph addresses
            "A_q": A_q,
            "A_scale_sh": A_scale_sh,
        }
        _GRAPH_CACHE[key] = entry
        _GRAPH_MISSES += 1
        return entry
    except Exception:
        _GRAPH_FAILS += 1
        _GRAPH_BLACKLIST.add(key)
        return None


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    global _warmed, _INLINE_AVAILABLE, _INLINE_WARNED

    A, B, B_q, B_shuffle, B_scale_sh = data

    if not A.is_contiguous():
        A = A.contiguous()
    if not B_shuffle.is_contiguous():
        B_shuffle = B_shuffle.contiguous()
    if not B_scale_sh.is_contiguous():
        B_scale_sh = B_scale_sh.contiguous()

    M, K = A.shape
    N    = B.shape[0]

    # Warm up JIT on first call (builds .so files)
    if not _warmed:
        A_q, A_scale_sh = _quantize_a(A)
        _warmed = True
        return aiter.gemm_a4w4(
            A_q, B_shuffle, A_scale_sh, B_scale_sh,
            dtype=_BF16, bpreshuffle=True,
        )

    padded_M, out_buf, use_asm, kernelName, splitK = \
        _get_or_create_bufs(M, N, K, A.device)

    should_autotune = (_AUTOTUNE_KERNEL or _AUTOTUNE_SPLITK) and (
        _AUTOTUNE_ALWAYS or (M, N, K) not in _TUNED_MAP
    )
    if should_autotune:
        A_q, A_scale_sh = _quantize_a(A)
        kernelName, splitK = _autotune_kernel_splitk(
            A_q, B_shuffle, A_scale_sh, B_scale_sh,
            out_buf, kernelName, splitK,
        )
        use_asm = "_ZN" in kernelName
        _TUNED_MAP[(M, N, K)] = (kernelName, splitK)
        _cache[(M, N, K)] = (padded_M, out_buf, use_asm, kernelName, splitK)
        if _PRINT_TUNED_MAP and (M, N, K) not in _TUNED_MAP_PRINTED:
            _TUNED_MAP_PRINTED.add((M, N, K))
            print(f"[mxfp4] tuned {M},{N},{K} -> kernel={kernelName} splitK={splitK}")

    if _INLINE_AVAILABLE and _INLINE_EXEC:
        try:
            _inline_mod.mxfp4_fused(A, B_shuffle, B_scale_sh, out_buf)
            return out_buf[:M]
        except Exception as e:
            _INLINE_AVAILABLE = False
            if _DEBUG:
                print(f"[mxfp4] inline exec failed, fallback: {e}")
    elif _INLINE_AVAILABLE and _DEBUG and not _INLINE_WARNED:
        _INLINE_WARNED = True
        print("[mxfp4] inline kernel compiled but not executed (set MXFP4_INLINE_EXEC=1)")

    graph_entry = _try_get_graph(
        M, N, K,
        A, B_shuffle, B_scale_sh,
        out_buf, use_asm, kernelName, splitK,
    )
    if graph_entry is not None:
        graph_entry["A_static"].copy_(A)
        if graph_entry["b_ptr"] != B_shuffle.data_ptr():
            graph_entry["B_shuffle_static"].copy_(B_shuffle)
            graph_entry["b_ptr"] = B_shuffle.data_ptr()
        if graph_entry["bs_ptr"] != B_scale_sh.data_ptr():
            graph_entry["B_scale_static"].copy_(B_scale_sh)
            graph_entry["bs_ptr"] = B_scale_sh.data_ptr()

        graph_entry["graph"].replay()
        return out_buf[:M]

    A_q, A_scale_sh = _quantize_a(A)

    # Pass unpadded A_q and A_scale_sh -- the ASM kernel handles M-padding internally
    if use_asm:
        _gemm_asm(
            A_q, B_shuffle, A_scale_sh, B_scale_sh,
            out_buf, kernelName,
            None, 1.0, 0.0, True,
            log2_k_split=splitK,
        )
    else:
        _gemm_blk(
            A_q, B_shuffle, A_scale_sh, B_scale_sh,
            out_buf, splitK,
        )

    return out_buf[:M]
scrolls · 726 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