Skip to content
KernelIndex
Search⌘K

submission 629301

francochengcc_30299 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a67d31c6dd6d7bed9f0aa94a5605b66c0a7aac9cee87f386d7fcd788dfba64fa
license declaredunknown
license concludedunknown
authorsfrancochengcc_30299
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.
shared-memory__shared__ float abs_vals[64];
split-k_A4W4_TUNED_HEADER = "cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n"
stages = 1NUM_STAGES=1,

Kernel source

submission.py1097 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

os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")


_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"
_SMALLM_TRITON_ENV_VAR = "MXFP4_ENABLE_SMALLM_TRITON"
_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,
    },
    (64, 7168, 2048): {
        "kind": "unified",
        "specialized_quant_candidate": True,
    },
    (256, 3072, 1536): {
        "kind": "unified",
        "specialized_quant_candidate": True,
    },
}
_NATIVE_QUANT_SHAPES = {
    (4, 512),
    (16, 7168),
    (32, 512),
    (64, 2048),
    (256, 1536),
}
_DIRECT_QUANT_CONFIGS = {
    (16, 7168): {
        "num_iter": 1,
        "block_size_m": 16,
        "block_size_n": 128,
        "num_warps": 4,
    },
    (64, 2048): {
        "num_iter": 2,
        "block_size_m": 64,
        "block_size_n": 128,
        "num_warps": 4,
    },
    (256, 1536): {
        "num_iter": 2,
        "block_size_m": 64,
        "block_size_n": 128,
        "num_warps": 4,
    },
}
_RUNTIME = None
_SPECIALIZED_QUANT_RUNTIME = None
_NATIVE_SHUFFLE_RUNTIME = None
_SMALLM_TRITON_RUNTIME = None
_SPECIALIZED_QUANT_WORKSPACES = {}
_GEMM_OUTPUT_WORKSPACES = {}
_SMALLM_TRITON_WEIGHT_WORKSPACES = {}
_SPECIALIZED_QUANT_ERROR = None
_SPECIALIZED_QUANT_INFO_PRINTED = False
_NATIVE_QUANT_ERROR = None
_NATIVE_QUANT_INFO_PRINTED = False
_NATIVE_SHUFFLE_ERROR = None
_NATIVE_SHUFFLE_INFO_PRINTED = False
_SMALLM_TRITON_ERROR = None
_SMALLM_TRITON_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)
    cache_key = ((out_rows, n), str(getattr(a_q, "device", "")), str(getattr(dtypes, "bf16", "bf16")))
    out = _GEMM_OUTPUT_WORKSPACES.get(cache_key)
    if out is None:
        out = a_q.new_empty((out_rows, n), dtype=dtypes.bf16)
        _GEMM_OUTPUT_WORKSPACES[cache_key] = out
    return out


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 + 31) // 32) * 32
    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 _get_specialized_quant_launch_config(triton, m: int, n: int):
    direct = _DIRECT_QUANT_CONFIGS.get((m, n))
    if direct is not None:
        return direct
    if m <= 32:
        return {
            "num_iter": 1,
            "block_size_m": triton.next_power_of_2(m),
            "block_size_n": 32,
            "num_warps": 1,
        }

    config = {
        "num_iter": 4,
        "block_size_m": 64,
        "block_size_n": 64,
        "num_warps": 4,
    }
    if n <= 16384:
        config["block_size_m"] = 32
        config["block_size_n"] = 128

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


def _get_specialized_shuffle_launch_config(m: int, scale_n_valid: int):
    return {
        "block_m": 32,
        "block_n": 8,
        "num_warps": 1,
    }


def _native_shuffle_enabled(m: int, scale_n_valid: int):
    del m, scale_n_valid
    return False


def _native_quant_enabled(m: int, n: int):
    return (m, n) in _NATIVE_QUANT_SHAPES


def _smallm_triton_enabled(torch, m: int, n: int, k: int, plan):
    if os.environ.get(_SMALLM_TRITON_ENV_VAR, "0") == "0":
        return False
    if plan.get("kind") != "asm":
        return False
    if (m, n, k) != (16, 2112, 7168):
        return False
    if getattr(getattr(torch, "version", None), "hip", None) is None:
        return False
    return hasattr(torch, "Tensor")


def _get_smallm_triton_runtime():
    global _SMALLM_TRITON_RUNTIME
    if _SMALLM_TRITON_RUNTIME is not None:
        return _SMALLM_TRITON_RUNTIME

    import triton
    from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
        _gemm_a16wfp4_preshuffle_kernel,
        _get_config,
    )
    from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
        _gemm_afp4wfp4_reduce_kernel,
    )

    _SMALLM_TRITON_RUNTIME = (
        triton,
        _gemm_a16wfp4_preshuffle_kernel,
        _gemm_afp4wfp4_reduce_kernel,
        _get_config,
    )
    return _SMALLM_TRITON_RUNTIME


def _run_smallm_triton_preshuffle(torch, dtypes, x, b_shuffle, b_scale_sh, m: int, n: int):
    triton, preshuffle_kernel, reduce_kernel, get_config = _get_smallm_triton_runtime()
    b_shuffle_u8, b_scale_sh_u8 = _get_smallm_triton_weight_views(torch, b_shuffle, b_scale_sh, n)
    packed_k = int(b_shuffle_u8.shape[1] // 16)
    config, _ = get_config(m, n, packed_k, True)
    config = dict(config)

    num_ksplit = int(config.get("NUM_KSPLIT", 1))
    block_size_k = int(config["BLOCK_SIZE_K"])
    config["BLOCK_SIZE_N"] = max(32, int(config["BLOCK_SIZE_N"]))

    if block_size_k >= 2 * packed_k:
        block_size_k = triton.next_power_of_2(2 * packed_k)
        config["BLOCK_SIZE_K"] = block_size_k
        config["SPLITK_BLOCK_SIZE"] = 2 * packed_k
        config["NUM_KSPLIT"] = 1
        num_ksplit = 1
    else:
        config["SPLITK_BLOCK_SIZE"] = 2 * packed_k

    y = torch.empty((m, n), dtype=dtypes.bf16, device=x.device)
    y_pp = None
    if num_ksplit > 1:
        y_pp = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=x.device)

    grid = (
        config["NUM_KSPLIT"]
        * triton.cdiv(m, config["BLOCK_SIZE_M"])
        * triton.cdiv(n, config["BLOCK_SIZE_N"]),
    )

    preshuffle_kernel[grid](
        x,
        b_shuffle_u8,
        y if config["NUM_KSPLIT"] == 1 else y_pp,
        b_scale_sh_u8,
        m,
        n,
        packed_k,
        x.stride(0),
        x.stride(1),
        b_shuffle_u8.stride(0),
        b_shuffle_u8.stride(1),
        0 if config["NUM_KSPLIT"] == 1 else y_pp.stride(0),
        y.stride(0) if config["NUM_KSPLIT"] == 1 else y_pp.stride(1),
        y.stride(1) if config["NUM_KSPLIT"] == 1 else y_pp.stride(2),
        b_scale_sh_u8.stride(0),
        b_scale_sh_u8.stride(1),
        PREQUANT=True,
        **config,
    )

    if config["NUM_KSPLIT"] > 1:
        reduce_block_m = 16
        reduce_block_n = 64
        actual_ksplit = triton.cdiv(packed_k, (config["SPLITK_BLOCK_SIZE"] // 2))
        grid_reduce = (
            triton.cdiv(m, reduce_block_m),
            triton.cdiv(n, reduce_block_n),
        )
        reduce_kernel[grid_reduce](
            y_pp,
            y,
            m,
            n,
            y_pp.stride(0),
            y_pp.stride(1),
            y_pp.stride(2),
            y.stride(0),
            y.stride(1),
            reduce_block_m,
            reduce_block_n,
            actual_ksplit,
            triton.next_power_of_2(config["NUM_KSPLIT"]),
        )
    return y


def _get_smallm_triton_weight_views(torch, b_shuffle, b_scale_sh, n: int):
    key = (
        int(getattr(b_shuffle, "data_ptr", lambda: 0)()),
        int(getattr(b_scale_sh, "data_ptr", lambda: 0)()),
        tuple(b_shuffle.shape),
        tuple(b_scale_sh.shape),
        str(getattr(b_shuffle, "device", "")),
    )
    cached = _SMALLM_TRITON_WEIGHT_WORKSPACES.get(key)
    if cached is not None:
        return cached

    b_shuffle_u8 = b_shuffle.view(torch.uint8).reshape(n // 16, b_shuffle.shape[1] * 16)

    b_scale_asm_u8 = b_scale_sh.view(torch.uint8)
    scale_rows, scale_cols = b_scale_asm_u8.shape
    raw_scales = (
        b_scale_asm_u8.view(scale_rows // 32, scale_cols // 8, 4, 16, 2, 2, 1)
        .permute(0, 5, 3, 1, 4, 2, 6)
        .contiguous()
        .view(scale_rows, scale_cols)
    )[:n]
    b_scale_triton_u8 = (
        raw_scales.view(n // 32, 2, 16, scale_cols // 8, 2, 4, 1)
        .permute(0, 3, 5, 2, 4, 1, 6)
        .contiguous()
        .view(n // 32, scale_cols * 32)
    )
    cached = (b_shuffle_u8, b_scale_triton_u8)
    _SMALLM_TRITON_WEIGHT_WORKSPACES[key] = cached
    return cached


def _maybe_run_smallm_triton(torch, dtypes, a, b_shuffle, b_scale_sh, m: int, n: int, k: int, plan):
    global _SMALLM_TRITON_ERROR, _SMALLM_TRITON_INFO_PRINTED
    if not _smallm_triton_enabled(torch, m, n, k, plan):
        return None

    try:
        y = _run_smallm_triton_preshuffle(
            torch,
            dtypes,
            a,
            b_shuffle,
            b_scale_sh,
            m,
            n,
        )
    except Exception:
        if _SMALLM_TRITON_ERROR is None:
            _SMALLM_TRITON_ERROR = True
            try:
                import traceback

                print("[mxfp4 smallm triton] falling back after error:")
                traceback.print_exc()
            except Exception:
                pass
        return None
    if not _SMALLM_TRITON_INFO_PRINTED:
        _SMALLM_TRITON_INFO_PRINTED = True
        try:
            print("[mxfp4 smallm triton] using uint8-view preshuffle kernel")
        except Exception:
            pass
    return y


def _get_native_shuffle_module(torch):
    global _NATIVE_SHUFFLE_RUNTIME
    if _NATIVE_SHUFFLE_RUNTIME is not None:
        return _NATIVE_SHUFFLE_RUNTIME

    from torch.utils.cpp_extension import load_inline

    cpp_src = """
    void mxfp4_quant_shuffle_small(
        torch::Tensor x,
        torch::Tensor x_fp4,
        torch::Tensor bs_raw,
        torch::Tensor bs_shuffled,
        int64_t m,
        int64_t n,
        int64_t n_pad);
    void mxfp4_quant_shuffle_small_noraw(
        torch::Tensor x,
        torch::Tensor x_fp4,
        torch::Tensor bs_shuffled,
        int64_t m,
        int64_t n,
        int64_t n_pad);
    void mxfp4_shuffle_e8m0(
        torch::Tensor src,
        torch::Tensor dst,
        int64_t m,
        int64_t n_valid,
        int64_t n_pad);
    """
    hip_src = r"""
    #include <torch/extension.h>
    #include <hip/hip_runtime.h>
    #include <hip/amd_detail/amd_hip_bf16.h>
    #include <cstdint>
    #include <cmath>

    __device__ __forceinline__ uint8_t float_to_e8m0(float x) {
      if (x <= 0.0f) {
        return 0;
      }
      uint32_t u = __float_as_uint(x);
      uint32_t exponent = (u >> 23) & 0xFF;
      bool round_case = ((u & 0x400000) > 0) &&
                        (((u & 0x200000) > 0) || ((u & 0x1FFFFF) > 0) || (exponent > 0));
      if (round_case && exponent < 0xFF) {
        exponent += 1;
      }
      return static_cast<uint8_t>(exponent);
    }

    __device__ __forceinline__ float e8m0_to_float(uint8_t x) {
      if (x == 0) {
        return __uint_as_float(0x00400000u);
      }
      if (x == 0xFF) {
        return __uint_as_float(0x7F800001u);
      }
      return __uint_as_float(static_cast<uint32_t>(x) << 23);
    }

    __device__ __forceinline__ uint8_t float_to_mxfp4(float x) {
      constexpr int EXP_BIAS_FP32 = 127;
      constexpr int EXP_BIAS_FP4 = 1;
      constexpr int EBITS_F32 = 8;
      constexpr int EBITS_FP4 = 2;
      constexpr int MBITS_F32 = 23;
      constexpr int MBITS_FP4 = 1;
      constexpr float MAX_NORMAL = 6.0f;
      constexpr float MIN_NORMAL = 1.0f;
      constexpr uint8_t MAX_INT = 0x7;

      uint32_t qx = __float_as_uint(x);
      uint32_t sign = qx & 0x80000000u;
      qx ^= sign;
      float qx_fp32 = __uint_as_float(qx);
      bool saturate_mask = qx_fp32 >= MAX_NORMAL;
      bool denormal_mask = (!saturate_mask) && (qx_fp32 < MIN_NORMAL);
      bool normal_mask = !(saturate_mask || denormal_mask);

      constexpr int denorm_exp = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1;
      constexpr uint32_t denorm_mask_int = static_cast<uint32_t>(denorm_exp) << MBITS_F32;
      float denorm_mask_float = __uint_as_float(denorm_mask_int);

      uint8_t denormal_x = static_cast<uint8_t>(__float_as_uint(qx_fp32 + denorm_mask_float) - denorm_mask_int);
      uint32_t normal_x = qx;
      uint32_t mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1u;
      constexpr int32_t val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1;
      normal_x = static_cast<uint32_t>(static_cast<int32_t>(normal_x) + val_to_add);
      normal_x += mant_odd;
      normal_x = normal_x >> (MBITS_F32 - MBITS_FP4);

      uint8_t out = MAX_INT;
      if (normal_mask) {
        out = static_cast<uint8_t>(normal_x);
      }
      if (denormal_mask) {
        out = denormal_x;
      }
      uint8_t sign_lp = static_cast<uint8_t>(sign >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4));
      return out | sign_lp;
    }

    __global__ void mxfp4_quant_shuffle_small_kernel(
        const __hip_bfloat16* x,
        uint8_t* x_fp4,
        uint8_t* bs_raw,
        uint8_t* bs_shuffled,
        int64_t stride_x_m,
        int64_t stride_x_n,
        int64_t stride_x_fp4_m,
        int64_t stride_x_fp4_n,
        int64_t m,
        int64_t n,
        int64_t n_pad) {
      __shared__ float abs_vals[64];
      __shared__ uint8_t fp4_codes[64];
      __shared__ uint8_t bs_e8m0[2];

      int lane = threadIdx.x;
      int row = blockIdx.y;
      int subgroup = lane / 32;
      int lane_in_group = lane % 32;
      int group_base = subgroup * 32;
      int block_n = blockIdx.x * 2 + subgroup;
      int col = block_n * 32 + lane_in_group;
      if (row >= m || col >= n) {
        return;
      }

      float x_val = __bfloat162float(x[row * stride_x_m + col * stride_x_n]);
      abs_vals[lane] = fabsf(x_val);
      __syncthreads();

      for (int offset = 16; offset > 0; offset >>= 1) {
        if (lane_in_group < offset) {
          abs_vals[group_base + lane_in_group] =
              fmaxf(abs_vals[group_base + lane_in_group], abs_vals[group_base + lane_in_group + offset]);
        }
        __syncthreads();
      }

      if (lane_in_group == 0) {
        float amax = abs_vals[group_base];
        if (amax == 0.0f) {
          bs_e8m0[subgroup] = 0;
        } else {
          uint32_t amax_bits = __float_as_uint(amax);
          amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
          float rounded_amax = __uint_as_float(amax_bits);
          bs_e8m0[subgroup] = float_to_e8m0(rounded_amax / 4.0f);
        }
      }
      __syncthreads();

      float scale = e8m0_to_float(bs_e8m0[subgroup]);
      float qx = x_val / scale;
      fp4_codes[lane] = float_to_mxfp4(qx);
      __syncthreads();

      if (lane_in_group < 16) {
        uint8_t even = fp4_codes[group_base + lane_in_group * 2];
        uint8_t odd = fp4_codes[group_base + lane_in_group * 2 + 1];
        x_fp4[row * stride_x_fp4_m + (block_n * 16 + lane_in_group) * stride_x_fp4_n] =
            static_cast<uint8_t>(even | (odd << 4));
      }

      if (lane_in_group == 0) {
        if (bs_raw != nullptr) {
          bs_raw[row * (n / 32) + block_n] = bs_e8m0[subgroup];
        }
        int64_t g0 = row / 32;
        int64_t rem_m = row % 32;
        int64_t g1 = rem_m / 16;
        int64_t g2 = rem_m % 16;
        int64_t g3 = block_n / 8;
        int64_t rem_n = block_n % 8;
        int64_t g4 = rem_n / 4;
        int64_t g5 = rem_n % 4;
        int64_t dst_off = g1 + g4 * 2 + g2 * 4 + g5 * 64 + g3 * 256 + g0 * (32 * n_pad);
        bs_shuffled[dst_off] = bs_e8m0[subgroup];
      }
    }

    __global__ void mxfp4_shuffle_e8m0_kernel(
        const uint8_t* src,
        uint8_t* dst,
        int64_t stride_src_m,
        int64_t stride_src_n,
        int64_t m,
        int64_t n_valid,
        int64_t n_pad) {
      int row = blockIdx.y * blockDim.y + threadIdx.y;
      int col = blockIdx.x * blockDim.x + threadIdx.x;
      if (row >= m || col >= n_valid) {
        return;
      }

      uint8_t val = src[row * stride_src_m + col * stride_src_n];
      int64_t g0 = row / 32;
      int64_t rem_m = row % 32;
      int64_t g1 = rem_m / 16;
      int64_t g2 = rem_m % 16;
      int64_t g3 = col / 8;
      int64_t rem_n = col % 8;
      int64_t g4 = rem_n / 4;
      int64_t g5 = rem_n % 4;
      int64_t dst_off =
          g1 + g4 * 2 + g2 * 4 + g5 * 64 + g3 * 256 + g0 * (32 * n_pad);
      dst[dst_off] = val;
    }

    void mxfp4_shuffle_e8m0(
        torch::Tensor src,
        torch::Tensor dst,
        int64_t m,
        int64_t n_valid,
        int64_t n_pad) {
      TORCH_CHECK(src.is_cuda(), "src must be CUDA/HIP");
      TORCH_CHECK(dst.is_cuda(), "dst must be CUDA/HIP");
      TORCH_CHECK(src.scalar_type() == torch::kUInt8, "src must be uint8");
      TORCH_CHECK(dst.scalar_type() == torch::kUInt8, "dst must be uint8");

      const dim3 threads(16, 16);
      const dim3 blocks(
          static_cast<unsigned int>((n_valid + threads.x - 1) / threads.x),
          static_cast<unsigned int>((m + threads.y - 1) / threads.y));

      hipLaunchKernelGGL(
          mxfp4_shuffle_e8m0_kernel,
          blocks,
          threads,
          0,
          0,
          src.data_ptr<uint8_t>(),
          dst.data_ptr<uint8_t>(),
          src.stride(0),
          src.stride(1),
          m,
          n_valid,
          n_pad);

      auto err = hipGetLastError();
      TORCH_CHECK(err == hipSuccess, hipGetErrorString(err));
    }

    void mxfp4_quant_shuffle_small(
        torch::Tensor x,
        torch::Tensor x_fp4,
        torch::Tensor bs_raw,
        torch::Tensor bs_shuffled,
        int64_t m,
        int64_t n,
        int64_t n_pad) {
      TORCH_CHECK(x.is_cuda(), "x must be CUDA/HIP");
      TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
      TORCH_CHECK(x_fp4.scalar_type() == torch::kUInt8, "x_fp4 must be uint8");
      TORCH_CHECK(bs_raw.scalar_type() == torch::kUInt8, "bs_raw must be uint8");
      TORCH_CHECK(bs_shuffled.scalar_type() == torch::kUInt8, "bs_shuffled must be uint8");
      TORCH_CHECK(
          m == 4 || m == 16 || m == 32 || m == 64 || m == 256,
          "native quant probe only supports m=4, m=16, m=32, m=64, or m=256");
      TORCH_CHECK((n % 64) == 0, "n must be divisible by 64");

      const dim3 threads(64);
      const dim3 blocks(
          static_cast<unsigned int>(n / 64),
          static_cast<unsigned int>(m));

      hipLaunchKernelGGL(
          mxfp4_quant_shuffle_small_kernel,
          blocks,
          threads,
          0,
          0,
          reinterpret_cast<const __hip_bfloat16*>(x.data_ptr()),
          x_fp4.data_ptr<uint8_t>(),
          bs_raw.data_ptr<uint8_t>(),
          bs_shuffled.data_ptr<uint8_t>(),
          x.stride(0),
          x.stride(1),
          x_fp4.stride(0),
          x_fp4.stride(1),
          m,
          n,
          n_pad);

      auto err = hipGetLastError();
      TORCH_CHECK(err == hipSuccess, hipGetErrorString(err));
    }

    void mxfp4_quant_shuffle_small_noraw(
        torch::Tensor x,
        torch::Tensor x_fp4,
        torch::Tensor bs_shuffled,
        int64_t m,
        int64_t n,
        int64_t n_pad) {
      TORCH_CHECK(x.is_cuda(), "x must be CUDA/HIP");
      TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
      TORCH_CHECK(x_fp4.scalar_type() == torch::kUInt8, "x_fp4 must be uint8");
      TORCH_CHECK(bs_shuffled.scalar_type() == torch::kUInt8, "bs_shuffled must be uint8");
      TORCH_CHECK(
          m == 4 || m == 16 || m == 32 || m == 64 || m == 256,
          "native quant probe only supports m=4, m=16, m=32, m=64, or m=256");
      TORCH_CHECK((n % 64) == 0, "n must be divisible by 64");

      const dim3 threads(64);
      const dim3 blocks(
          static_cast<unsigned int>(n / 64),
          static_cast<unsigned int>(m));

      hipLaunchKernelGGL(
          mxfp4_quant_shuffle_small_kernel,
          blocks,
          threads,
          0,
          0,
          reinterpret_cast<const __hip_bfloat16*>(x.data_ptr()),
          x_fp4.data_ptr<uint8_t>(),
          nullptr,
          bs_shuffled.data_ptr<uint8_t>(),
          x.stride(0),
          x.stride(1),
          x_fp4.stride(0),
          x_fp4.stride(1),
          m,
          n,
          n_pad);

      auto err = hipGetLastError();
      TORCH_CHECK(err == hipSuccess, hipGetErrorString(err));
    }

    """

    _NATIVE_SHUFFLE_RUNTIME = load_inline(
        name="mxfp4_native_shuffle_v4",
        cpp_sources=[cpp_src],
        cuda_sources=[hip_src],
        functions=[
            "mxfp4_quant_shuffle_small",
            "mxfp4_quant_shuffle_small_noraw",
            "mxfp4_shuffle_e8m0",
        ],
        extra_cflags=["-std=c++20"],
        extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20"],
        verbose=False,
    )
    return _NATIVE_SHUFFLE_RUNTIME


def _maybe_native_shuffle(torch, src, dst, m: int, n_valid: int, n_pad: int):
    global _NATIVE_SHUFFLE_ERROR, _NATIVE_SHUFFLE_INFO_PRINTED
    if not _native_shuffle_enabled(m, n_valid):
        return False
    try:
        module = _get_native_shuffle_module(torch)
        module.mxfp4_shuffle_e8m0(src, dst, m, n_valid, n_pad)
    except Exception:
        if _NATIVE_SHUFFLE_ERROR is None:
            _NATIVE_SHUFFLE_ERROR = True
            try:
                import traceback

                print("[mxfp4 native shuffle] falling back after error:")
                traceback.print_exc()
            except Exception:
                pass
        return False
    if not _NATIVE_SHUFFLE_INFO_PRINTED:
        _NATIVE_SHUFFLE_INFO_PRINTED = True
        try:
            print("[mxfp4 native shuffle] using load_inline HIP path")
        except Exception:
            pass
    return True


def _maybe_native_quant_and_shuffle(
    torch,
    x,
    a_q_raw,
    a_scale_raw,
    a_scale_shuffled_raw,
    m: int,
    n: int,
    n_pad: int,
):
    global _NATIVE_QUANT_ERROR, _NATIVE_QUANT_INFO_PRINTED
    if not _native_quant_enabled(m, n):
        return False
    try:
        module = _get_native_shuffle_module(torch)
        module.mxfp4_quant_shuffle_small_noraw(
            x,
            a_q_raw,
            a_scale_shuffled_raw,
            m,
            n,
            n_pad,
        )
    except Exception:
        if _NATIVE_QUANT_ERROR is None:
            _NATIVE_QUANT_ERROR = True
            try:
                import traceback

                print("[mxfp4 native quant] falling back after error:")
                traceback.print_exc()
            except Exception:
                pass
        return False
    if not _NATIVE_QUANT_INFO_PRINTED:
        _NATIVE_QUANT_INFO_PRINTED = True
        try:
            print("[mxfp4 native quant] using load_inline HIP path")
        except Exception:
            pass
    return True


def _quant_mxfp4_specialized_with_workspace(torch, dtypes, x, workspace):
    triton, quant_kernel, shuffle_kernel = _get_specialized_quant_runtime()
    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"]

    native_quant_used = _maybe_native_quant_and_shuffle(
        torch,
        x,
        a_q_raw,
        a_scale_raw,
        a_scale_shuffled_raw,
        m,
        n,
        scale_n_pad,
    )
    if not native_quant_used:
        launch = _get_specialized_quant_launch_config(triton, m, n)
        num_iter = launch["num_iter"]
        block_size_m = launch["block_size_m"]
        block_size_n = launch["block_size_n"]
        num_warps = launch["num_warps"]

        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,
        )

    if not native_quant_used and not _maybe_native_shuffle(torch, a_scale_raw, a_scale_shuffled_raw, m, scale_n_valid, scale_n_pad):
        shuffle_launch = _get_specialized_shuffle_launch_config(m, scale_n_valid)
        shuffle_block_m = shuffle_launch["block_m"]
        shuffle_block_n = shuffle_launch["block_n"]
        shuffle_grid = (triton.cdiv(m, shuffle_block_m), triton.cdiv(scale_n_valid, shuffle_block_n))
        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=shuffle_block_m,
            BLOCK_N=shuffle_block_n,
            num_warps=shuffle_launch["num_warps"],
            num_stages=1,
        )
    return a_q_raw.view(dtypes.fp4x2), a_scale_shuffled_raw.view(dtypes.fp8_e8m0)


def _quant_mxfp4_specialized(torch, dtypes, x):
    workspace = _get_specialized_quant_workspace(torch, x)
    return _quant_mxfp4_specialized_with_workspace(torch, dtypes, x, workspace)


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)

    smallm_triton = _maybe_run_smallm_triton(
        torch,
        dtypes,
        a,
        b_shuffle,
        b_scale_sh,
        m,
        n,
        k,
        plan,
    )
    if smallm_triton is not None:
        return smallm_triton

    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 · 1097 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