Skip to content
KernelIndex
Search⌘K

submission 620883

fchange · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:20db7f3ed3970b66784e586f201c42b352b5e410cc2fda82f34bfbcce5fbe76b
license declaredunknown
license concludedunknown
authorsfchange
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[32];
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.py993 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"
_ASM_64X1024 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_64x1024E"
_ASM_64X1024_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_64x1024.co"
_ASM_64X512 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x512E"
_ASM_64X512_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_64x512.co"
_ASM_128X512 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_128x512E"
_ASM_128X512_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_128x512.co"
_ASM_224X256 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_224x256E"
_ASM_224X256_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_224x256.co"
_ASM_256X256 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x256E"
_ASM_256X256_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_256x256.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"
_SPECIALIZED_QUANT_VARIANT_ENV_VAR = "MXFP4_SPECIALIZED_QUANT_VARIANT"
_LARGE_SHAPE_VARIANT_ENV_VAR = "MXFP4_LARGE_SHAPE_VARIANT"
_SPECIALIZED_QUANT_DEFAULT_VARIANT = "combo_m4m16m32cppqshuf"
_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,
    },
}
_LARGE_SHAPE_VARIANTS = {
    "baseline": {},
    "64x1024_256x256s1": {
        (64, 7168, 2048): {
            "kind": "asm",
            "kernel_name": _ASM_64X1024,
            "co_name": _ASM_64X1024_CO,
            "log2_k_split": None,
            "specialized_quant_candidate": False,
        },
        (256, 3072, 1536): {
            "kind": "asm",
            "kernel_name": _ASM_256X256,
            "co_name": _ASM_256X256_CO,
            "log2_k_split": 1,
            "specialized_quant_candidate": False,
        },
    },
    "64x1024_256x256s2": {
        (64, 7168, 2048): {
            "kind": "asm",
            "kernel_name": _ASM_64X1024,
            "co_name": _ASM_64X1024_CO,
            "log2_k_split": None,
            "specialized_quant_candidate": False,
        },
        (256, 3072, 1536): {
            "kind": "asm",
            "kernel_name": _ASM_256X256,
            "co_name": _ASM_256X256_CO,
            "log2_k_split": 2,
            "specialized_quant_candidate": False,
        },
    },
    "128x512s1_both": {
        (64, 7168, 2048): {
            "kind": "asm",
            "kernel_name": _ASM_128X512,
            "co_name": _ASM_128X512_CO,
            "log2_k_split": 1,
            "specialized_quant_candidate": False,
        },
        (256, 3072, 1536): {
            "kind": "asm",
            "kernel_name": _ASM_128X512,
            "co_name": _ASM_128X512_CO,
            "log2_k_split": 1,
            "specialized_quant_candidate": False,
        },
    },
    "128x512s2_both": {
        (64, 7168, 2048): {
            "kind": "asm",
            "kernel_name": _ASM_128X512,
            "co_name": _ASM_128X512_CO,
            "log2_k_split": 2,
            "specialized_quant_candidate": False,
        },
        (256, 3072, 1536): {
            "kind": "asm",
            "kernel_name": _ASM_128X512,
            "co_name": _ASM_128X512_CO,
            "log2_k_split": 2,
            "specialized_quant_candidate": False,
        },
    },
}
_RUNTIME = None
_SPECIALIZED_QUANT_RUNTIME = None
_NATIVE_SHUFFLE_RUNTIME = None
_SPECIALIZED_QUANT_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


def _select_gemm_plan(m: int, n: int, k: int):
    shape = (m, n, k)
    variant = os.environ.get(_LARGE_SHAPE_VARIANT_ENV_VAR, "baseline")
    variant_plans = _LARGE_SHAPE_VARIANTS.get(variant, _LARGE_SHAPE_VARIANTS["baseline"])
    if shape in variant_plans:
        return variant_plans[shape]
    return _SHAPE_PLANS.get(shape, _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 _get_specialized_quant_launch_config(triton, m: int, n: int):
    variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
    if variant == "combo_best" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 256,
            "num_warps": 4,
        }
    if variant == "combo_best" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16sh16" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 256,
            "num_warps": 4,
        }
    if variant == "combo_m16sh16" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16n128sh16" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16n128sh16" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16n128cppshuf" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16n128cppshuf" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16n128cppqshuf" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16n128cppqshuf" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16m32cppqshuf" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m16m32cppqshuf" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m4m16m32cppqshuf" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "combo_m4m16m32cppqshuf" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "lg64x64" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 4,
            "block_size_m": 64,
            "block_size_n": 64,
            "num_warps": 4,
        }
    if variant == "lg64x128i2" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 2,
            "block_size_m": 64,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "lg32x128w8" and (m, n) in ((64, 2048), (256, 1536)):
        return {
            "num_iter": 4,
            "block_size_m": 32,
            "block_size_n": 128,
            "num_warps": 8,
        }
    if variant == "m16n128w4" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 128,
            "num_warps": 4,
        }
    if variant == "m16n256w4" and (m, n) == (16, 7168):
        return {
            "num_iter": 1,
            "block_size_m": 16,
            "block_size_n": 256,
            "num_warps": 4,
        }
    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):
    variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
    if variant in ("m16sh16", "combo_m16sh16", "combo_m16n128sh16") and (m, scale_n_valid) == (16, 224):
        return {
            "block_m": 16,
            "block_n": 16,
            "num_warps": 1,
        }
    return {
        "block_m": 32,
        "block_n": 8,
        "num_warps": 1,
    }


def _native_shuffle_enabled(m: int, scale_n_valid: int):
    variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
    return variant == "combo_m16n128cppshuf" and (m, scale_n_valid) == (16, 224)


def _native_quant_enabled(m: int, n: int):
    variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
    if variant == "combo_m16n128cppqshuf":
        return (m, n) == (16, 7168)
    if variant == "combo_m16m32cppqshuf":
        return (m, n) in ((16, 7168), (32, 512))
    if variant == "combo_m4m16m32cppqshuf":
        return (m, n) in ((4, 512), (16, 7168), (32, 512))
    return False


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_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_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[32];
      __shared__ uint8_t fp4_codes[32];
      __shared__ uint8_t bs_e8m0;

      int lane = threadIdx.x;
      int row = blockIdx.y;
      int block_n = blockIdx.x;
      int col = block_n * 32 + lane;
      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 < offset) {
          abs_vals[lane] = fmaxf(abs_vals[lane], abs_vals[lane + offset]);
        }
        __syncthreads();
      }

      if (lane == 0) {
        float amax = abs_vals[0];
        if (amax == 0.0f) {
          bs_e8m0 = 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 = float_to_e8m0(rounded_amax / 4.0f);
        }
      }
      __syncthreads();

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

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

      if (lane == 0) {
        int64_t g1 = row / 16;
        int64_t g2 = row % 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;
        bs_shuffled[dst_off] = bs_e8m0;
      }
    }

    __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_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, "native quant probe only supports m=4, m=16, or m=32");
      TORCH_CHECK((n % 32) == 0, "n must be divisible by 32");

      const dim3 threads(32);
      const dim3 blocks(
          static_cast<unsigned int>(n / 32),
          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_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_v1",
        cpp_sources=[cpp_src],
        cuda_sources=[hip_src],
        functions=["mxfp4_quant_shuffle_small", "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_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(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(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"]

    native_quant_used = _maybe_native_quant_and_shuffle(torch, x, a_q_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 _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 · 993 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