Skip to content
KernelIndex
Search⌘K

submission 619457

Coalwood · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6e1ab7a4f6ac321f45e861a837e927fe65adafee848492a0977f45a639d98553
license declaredunknown
license concludedunknown
authorsCoalwood
imported2026-08-26

Techniques

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

fp4MXFP4 per-1x32 quant on A, then A4W4 GEMM on MI355X.
num-warps = 1num_warps = 1
split-k_A4W4_TUNED_HEADER = "cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n"
stages = 1NUM_STAGES=1,
tile-m = 32BLOCK_M=32,
tile-n = 8BLOCK_N=8,

Kernel source

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

"""
MXFP4 per-1x32 quant on A, then A4W4 GEMM on MI355X.
Formal submission path:
- exact-shape asm dispatch for the fixed benchmark shapes
- specialized quant+shuffle path for those same shapes
- unified aiter fallback for everything else
"""
import importlib.util
import os

from task import input_t, output_t


_ASM_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_32X128_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co"
_A4W4_TUNED_HEADER = "cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n"
_A4W4_TUNED_ROWS = [
    (256, 4, 2880, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
    (256, 16, 2112, 7168, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
    (256, 32, 4096, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
    (256, 32, 2880, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
]
_SPECIALIZED_QUANT_ENV_VAR = "MXFP4_ENABLE_SPECIALIZED_QUANT"
_A4W4_TUNED_OVERRIDE = None
_UNIFIED_PLAN = {"kind": "unified"}
_SHAPE_PLANS = {
    (4, 2880, 512): {
        "kind": "asm",
        "kernel_name": _ASM_32X128,
        "co_name": _ASM_32X128_CO,
        "log2_k_split": None,
        "specialized_quant_candidate": True,
    },
    (16, 2112, 7168): {
        "kind": "asm",
        "kernel_name": _ASM_32X128,
        "co_name": _ASM_32X128_CO,
        "log2_k_split": None,
        "specialized_quant_candidate": True,
    },
    (32, 4096, 512): {
        "kind": "asm",
        "kernel_name": _ASM_32X128,
        "co_name": _ASM_32X128_CO,
        "log2_k_split": None,
        "specialized_quant_candidate": True,
    },
    (32, 2880, 512): {
        "kind": "asm",
        "kernel_name": _ASM_32X128,
        "co_name": _ASM_32X128_CO,
        "log2_k_split": None,
        "specialized_quant_candidate": True,
    },
}
_RUNTIME = None
_SPECIALIZED_QUANT_RUNTIME = None
_SPECIALIZED_QUANT_WORKSPACES = {}
_SPECIALIZED_QUANT_ERROR = None
_SPECIALIZED_QUANT_INFO_PRINTED = False


def _select_gemm_plan(m: int, n: int, k: int):
    return _SHAPE_PLANS.get((m, n, k), _UNIFIED_PLAN)


def _find_aiter_config_path():
    try:
        spec = importlib.util.find_spec("aiter")
    except (ImportError, ValueError):
        spec = None
    locations = getattr(spec, "submodule_search_locations", None) if spec is not None else None
    if locations:
        return os.path.join(locations[0], "configs", "a4w4_blockscale_tuned_gemm.csv")
    runtime = globals().get("_RUNTIME")
    if runtime is None:
        return None
    _, aiter, _, _, _ = runtime
    package_file = getattr(aiter, "__file__", None)
    if package_file:
        return os.path.join(
            os.path.dirname(os.path.abspath(package_file)),
            "configs",
            "a4w4_blockscale_tuned_gemm.csv",
        )
    return None


def _render_a4w4_tuned_override():
    rows = ["{},{},{},{},{},{},{},{},{},{},{}".format(*row) for row in _A4W4_TUNED_ROWS]
    return _A4W4_TUNED_HEADER + "\n".join(rows) + "\n"


def _ensure_a4w4_tuned_override():
    global _A4W4_TUNED_OVERRIDE
    override_path = _A4W4_TUNED_OVERRIDE or "/tmp/mxfp4_a4w4_tuned_override.csv"
    content = _render_a4w4_tuned_override()
    try:
        existing = None
        if os.path.exists(override_path):
            with open(override_path, "r", encoding="utf-8") as handle:
                existing = handle.read()
        if existing != content:
            with open(override_path, "w", encoding="utf-8") as handle:
                handle.write(content)
    except Exception:
        return None
    _A4W4_TUNED_OVERRIDE = override_path

    default_path = _find_aiter_config_path()
    if not default_path:
        return override_path
    os.environ["AITER_CONFIG_GEMM_A4W4"] = os.pathsep.join([default_path, override_path])
    return override_path


def _get_runtime():
    global _RUNTIME
    if _RUNTIME is None:
        _ensure_a4w4_tuned_override()
        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

        _RUNTIME = (torch, aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
    return _RUNTIME


def _alloc_gemm_output(a_q, dtypes, m: int, n: int, zero_init: bool = False):
    out_rows = ((m + 31) // 32) * 32
    if zero_init:
        return a_q.new_zeros((out_rows, n), dtype=dtypes.bf16)
    return a_q.new_empty((out_rows, n), dtype=dtypes.bf16)


def _quant_mxfp4(x, dtypes, dynamic_mxfp4_quant, e8m0_shuffle):
    x_fp4, scale = dynamic_mxfp4_quant(x)
    scale = e8m0_shuffle(scale)
    return x_fp4.view(dtypes.fp4x2), scale.view(dtypes.fp8_e8m0)


def _specialized_quant_runtime_enabled(torch, plan):
    if not plan.get("specialized_quant_candidate"):
        return False
    if os.environ.get(_SPECIALIZED_QUANT_ENV_VAR, "1") == "0":
        return False
    return getattr(getattr(torch, "version", None), "hip", None) is not None


def _get_specialized_quant_runtime():
    global _SPECIALIZED_QUANT_RUNTIME
    if _SPECIALIZED_QUANT_RUNTIME is not None:
        return _SPECIALIZED_QUANT_RUNTIME

    import triton
    import triton.language as tl
    from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel

    @triton.jit
    def _shuffle_e8m0_scale_kernel(
        src_ptr,
        dst_ptr,
        stride_src_m,
        stride_src_n,
        M,
        N_VALID,
        N_PAD,
        BLOCK_M: tl.constexpr,
        BLOCK_N: tl.constexpr,
    ):
        pid_m = tl.program_id(0)
        pid_n = tl.program_id(1)
        offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
        offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
        mask = (offs_m[:, None] < M) & (offs_n[None, :] < N_VALID)
        src_offs = offs_m[:, None] * stride_src_m + offs_n[None, :] * stride_src_n
        vals = tl.load(src_ptr + src_offs, mask=mask, other=127)

        g0 = offs_m[:, None] // 32
        rem_m = offs_m[:, None] % 32
        g1 = rem_m // 16
        g2 = rem_m % 16
        g3 = offs_n[None, :] // 8
        rem_n = offs_n[None, :] % 8
        g4 = rem_n // 4
        g5 = rem_n % 4
        dst_offs = (
            g1
            + g4 * 2
            + g2 * 4
            + g5 * 64
            + g3 * 256
            + g0 * (32 * N_PAD)
        )
        tl.store(dst_ptr + dst_offs, vals, mask=mask)

    _SPECIALIZED_QUANT_RUNTIME = (triton, _dynamic_mxfp4_quant_kernel, _shuffle_e8m0_scale_kernel)
    return _SPECIALIZED_QUANT_RUNTIME


def _get_specialized_quant_workspace(torch, x):
    m, n = x.shape
    scale_n_valid = (n + 31) // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256
    cache_key = (tuple(x.shape), str(getattr(x, "device", "")), str(getattr(x, "dtype", "")))
    workspace = _SPECIALIZED_QUANT_WORKSPACES.get(cache_key)
    if workspace is None:
        workspace = {
            "a_q_raw": torch.empty((m, n // 2), dtype=torch.uint8, device=x.device),
            "a_scale_raw": torch.empty((m, scale_n_valid), dtype=torch.uint8, device=x.device),
            "a_scale_shuffled_raw": torch.full(
                (scale_m_pad, scale_n_pad),
                127,
                dtype=torch.uint8,
                device=x.device,
            ),
            "scale_n_valid": scale_n_valid,
            "scale_n_pad": scale_n_pad,
        }
        _SPECIALIZED_QUANT_WORKSPACES[cache_key] = workspace
    return workspace


def _quant_mxfp4_specialized(torch, dtypes, x):
    triton, quant_kernel, shuffle_kernel = _get_specialized_quant_runtime()
    workspace = _get_specialized_quant_workspace(torch, x)
    m, n = x.shape
    a_q_raw = workspace["a_q_raw"]
    a_scale_raw = workspace["a_scale_raw"]
    a_scale_shuffled_raw = workspace["a_scale_shuffled_raw"]
    scale_n_valid = workspace["scale_n_valid"]
    scale_n_pad = workspace["scale_n_pad"]

    if m <= 32:
        num_iter = 1
        block_size_m = triton.next_power_of_2(m)
        block_size_n = 32
        num_warps = 1
    else:
        num_iter = 4
        block_size_m = 64
        block_size_n = 64
        num_warps = 4
        if n <= 16384:
            block_size_m = 32
            block_size_n = 128

    if n <= 1024:
        num_iter = 1
        block_size_n = min(256, triton.next_power_of_2(n))
        block_size_n = max(32, block_size_n)
        block_size_m = min(8, triton.next_power_of_2(m))
        num_warps = 4

    grid = (triton.cdiv(m, block_size_m), triton.cdiv(n, block_size_n * num_iter))
    quant_kernel[grid](
        x,
        a_q_raw,
        a_scale_raw,
        *x.stride(),
        *a_q_raw.stride(),
        *a_scale_raw.stride(),
        M=m,
        N=n,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        NUM_ITER=num_iter,
        BLOCK_SIZE_M=block_size_m,
        BLOCK_SIZE_N=block_size_n,
        NUM_STAGES=1,
        num_warps=num_warps,
        waves_per_eu=0,
        num_stages=1,
    )

    shuffle_grid = (triton.cdiv(m, 32), triton.cdiv(scale_n_valid, 8))
    shuffle_kernel[shuffle_grid](
        a_scale_raw,
        a_scale_shuffled_raw,
        *a_scale_raw.stride(),
        M=m,
        N_VALID=scale_n_valid,
        N_PAD=scale_n_pad,
        BLOCK_M=32,
        BLOCK_N=8,
    )
    return a_q_raw.view(dtypes.fp4x2), a_scale_shuffled_raw.view(dtypes.fp8_e8m0)


def _maybe_quant_mxfp4_specialized(torch, dtypes, x, plan):
    global _SPECIALIZED_QUANT_ERROR, _SPECIALIZED_QUANT_INFO_PRINTED
    if not _specialized_quant_runtime_enabled(torch, plan):
        return None
    try:
        result = _quant_mxfp4_specialized(torch, dtypes, x)
    except Exception:
        if _SPECIALIZED_QUANT_ERROR is None:
            _SPECIALIZED_QUANT_ERROR = True
            try:
                import traceback

                print("[mxfp4 quant] falling back after error:")
                traceback.print_exc()
            except Exception:
                pass
        return None
    if not _SPECIALIZED_QUANT_INFO_PRINTED:
        _SPECIALIZED_QUANT_INFO_PRINTED = True
        try:
            print("[mxfp4 quant] using specialized quant+shuffle path")
        except Exception:
            pass
    return result


def _run_gemm_asm(aiter, dtypes, a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, plan):
    out = _alloc_gemm_output(a_q, dtypes, m, n)
    aiter.gemm_a4w4_asm(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        out,
        plan["kernel_name"],
        bpreshuffle=True,
        log2_k_split=plan["log2_k_split"],
    )
    return out[:m]


def custom_kernel(data: input_t) -> output_t:
    torch, aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle = _get_runtime()
    a, b, _b_q, b_shuffle, b_scale_sh = data
    a = a.contiguous()
    b = b.contiguous()
    m, k = a.shape
    n, _ = b.shape
    plan = _select_gemm_plan(m, n, k)

    specialized = _maybe_quant_mxfp4_specialized(torch, dtypes, a, plan)
    if specialized is None:
        a_q, a_scale_sh = _quant_mxfp4(a, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
    else:
        a_q, a_scale_sh = specialized

    if plan.get("kind") == "asm":
        return _run_gemm_asm(
            aiter,
            dtypes,
            a_q,
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            m,
            n,
            plan,
        )
    return aiter.gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
scrolls · 372 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