Skip to content
KernelIndex
Search⌘K

submission 655171

fchange3413 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-655171?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
155.0µs
#239 of 782
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cdda5b4407f1a64a6ba499c6df29d5d59621e847d00af9c9a9da6f93ca4118b7
license declaredunknown
license concludedunknown
authorsfchange3413
imported2026-08-15

Techniques

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

fp4"[moe-mxfp4] using native load_inline fused quant+sort replacement for stage1"
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"

Kernel source

submission.py774 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X

import importlib
import importlib.util
import glob
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"),
]
_NATIVE_QUANT_SHAPES = {
    (4, 512),
    (16, 7168),
    (128, 7168),
    (32, 512),
    (64, 2048),
    (256, 1536),
}
_NATIVE_SHUFFLE_SHAPES = {
    (4, 16),
    (16, 224),
    (32, 16),
    (64, 64),
    (256, 48),
}
_A4W4_TUNED_OVERRIDE = None
_FMOE_TUNED_HEADER = (
    "cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,q_dtype_w,"
    "q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,err1,us2,kernelName2,"
    "err2,us,run_1stage,tflops,bw,_tag\n"
)
_FMOE_TUNED_ROWS = [
    (
        256, 16, 7168, 256, 257, 9,
        "ActivationType.Silu", "torch.bfloat16", "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
        "QuantType.per_1x32", 1, 0, 32, 4, "0.0", "", "0.0", "0.0", "", "0.0", "0.0", 0, "0.0", "0.0", "",
    ),
    (
        256, 128, 7168, 256, 257, 9,
        "ActivationType.Silu", "torch.bfloat16", "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
        "QuantType.per_1x32", 1, 0, 32, 4, "0.0", "", "0.0", "0.0", "", "0.0", "0.0", 0, "0.0", "0.0", "",
    ),
    (
        256, 16, 7168, 512, 33, 9,
        "ActivationType.Silu", "torch.bfloat16", "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
        "QuantType.per_1x32", 1, 0, 32, 2, "0.0", "", "0.0", "0.0", "", "0.0", "0.0", 0, "0.0", "0.0", "",
    ),
    (
        256, 128, 7168, 512, 33, 9,
        "ActivationType.Silu", "torch.bfloat16", "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
        "QuantType.per_1x32", 1, 0, 64, 2, "0.0", "", "0.0", "0.0", "", "0.0", "0.0", 0, "0.0", "0.0", "",
    ),
    (
        256, 512, 7168, 512, 33, 9,
        "ActivationType.Silu", "torch.bfloat16", "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
        "QuantType.per_1x32", 1, 0, 32, 0,
        "0.0",
        "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "0.0",
        "0.0",
        "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
        "0.0",
        "0.0",
        0,
        "0.0",
        "0.0",
        "",
    ),
    (
        256, 512, 7168, 256, 257, 9,
        "ActivationType.Silu", "torch.bfloat16", "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
        "QuantType.per_1x32", 1, 0, 32, 0,
        "0.0",
        "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "0.0",
        "0.0",
        "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
        "0.0",
        "0.0",
        0,
        "0.0",
        "0.0",
        "",
    ),
]
_FMOE_TUNED_OVERRIDE = None
_RUNTIME = None
_NATIVE_RUNTIME = None
_PATCHED_QUANT = False
_ORIG_DYNAMIC_MXFP4_QUANT = None
_ORIG_FUSED_DYNAMIC_MXFP4_QUANT_MOE_SORT = None
_ORIG_E8M0_SHUFFLE = None
_ORIG_GET_QUANT = None
_NATIVE_QUANT_ERROR = False
_NATIVE_SHUFFLE_ERROR = False
_NATIVE_QUANT_INFO_PRINTED = False
_QUANT_WORKSPACES = {}
_SHUFFLE_WORKSPACES = {}


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


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")
    return None


def _find_aiter_fmoe_config_paths():
    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:
        config_dir = os.path.join(locations[0], "configs")
        paths = []
        default_path = os.path.join(config_dir, "tuned_fmoe.csv")
        if os.path.exists(default_path):
            paths.append(default_path)
        model_glob = os.path.join(config_dir, "model_configs", "*tuned_fmoe*.csv")
        for path in sorted(glob.glob(model_glob)):
            if "untuned" not in path:
                paths.append(path)
        return paths
    return []


def _ensure_a4w4_tuned_override():
    global _A4W4_TUNED_OVERRIDE
    override_path = _A4W4_TUNED_OVERRIDE or "/tmp/moe_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
    desired = os.pathsep.join([default_path, override_path])
    if os.environ.get("AITER_CONFIG_GEMM_A4W4") != desired:
        os.environ["AITER_CONFIG_GEMM_A4W4"] = desired
    return override_path


def _render_fmoe_tuned_override():
    rows = ["{}".format(",".join(str(col) for col in row)) for row in _FMOE_TUNED_ROWS]
    return _FMOE_TUNED_HEADER + "\n".join(rows) + "\n"


def _ensure_fmoe_tuned_override():
    global _FMOE_TUNED_OVERRIDE
    override_path = _FMOE_TUNED_OVERRIDE or "/tmp/moe_mxfp4_fmoe_tuned_override.csv"
    content = _render_fmoe_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
    _FMOE_TUNED_OVERRIDE = override_path

    config_paths = _find_aiter_fmoe_config_paths()
    desired = os.pathsep.join(config_paths + [override_path]) if config_paths else override_path
    if os.environ.get("AITER_CONFIG_FMOE") != desired:
        os.environ["AITER_CONFIG_FMOE"] = desired
    return override_path


def _native_quant_enabled(m: int, n: int):
    return (m, n) in _NATIVE_QUANT_SHAPES and (n % 64) == 0


def _native_shuffle_enabled(m: int, n_valid: int):
    return (m, n_valid) in _NATIVE_SHUFFLE_SHAPES


def _get_native_workspace(torch, x):
    m, n = x.shape
    scale_n_valid = (n + 31) // 32
    key = (tuple(x.shape), str(x.device), str(x.dtype))
    workspace = _QUANT_WORKSPACES.get(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),
        }
        _QUANT_WORKSPACES[key] = workspace
    return workspace


def _get_shuffle_workspace(torch, src):
    m, n_valid = src.shape
    m_pad = ((m + 31) // 32) * 32
    n_pad = ((n_valid + 7) // 8) * 8
    key = (tuple(src.shape), str(src.device), str(src.dtype))
    workspace = _SHUFFLE_WORKSPACES.get(key)
    if workspace is None:
        workspace = {
            "dst_raw": torch.full((m_pad, n_pad), 127, dtype=torch.uint8, device=src.device),
            "m_pad": m_pad,
            "n_pad": n_pad,
        }
        _SHUFFLE_WORKSPACES[key] = workspace
    return workspace


def _get_native_module(torch):
    global _NATIVE_RUNTIME
    if _NATIVE_RUNTIME is not None:
        return _NATIVE_RUNTIME

    from torch.utils.cpp_extension import load_inline

    rocm_arch = os.environ.get("PYTORCH_ROCM_ARCH", "gfx950").split(";")[0]
    cpp_src = """
    void mxfp4_quant_small_raw(
        torch::Tensor x,
        torch::Tensor x_fp4,
        torch::Tensor bs_raw,
        int64_t m,
        int64_t n);
    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__ 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_small_raw_kernel(
        const __hip_bfloat16* x,
        uint8_t* x_fp4,
        uint8_t* bs_raw,
        int64_t stride_x_m,
        int64_t stride_x_n,
        int64_t stride_x_fp4_m,
        int64_t stride_x_fp4_n,
        int64_t stride_bs_m,
        int64_t stride_bs_n,
        int64_t m,
        int64_t n) {
      __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();

      uint32_t scale_bits = static_cast<uint32_t>(bs_e8m0[subgroup]) << 23;
      float scale = __uint_as_float(scale_bits);
      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) {
        bs_raw[row * stride_bs_m + block_n * stride_bs_n] = 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_quant_small_raw(
        torch::Tensor x,
        torch::Tensor x_fp4,
        torch::Tensor bs_raw,
        int64_t m,
        int64_t n) {
      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(
          m == 4 || m == 16 || m == 32 || m == 64 || m == 128 || m == 256 || m == 512,
          "native quant only supports m=4,16,32,64,128,256,512");
      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_small_raw_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>(),
          x.stride(0),
          x.stride(1),
          x_fp4.stride(0),
          x_fp4.stride(1),
          bs_raw.stride(0),
          bs_raw.stride(1),
          m,
          n);

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

    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));
    }
    """

    _NATIVE_RUNTIME = load_inline(
        name="moe_mxfp4_native_quant_v1",
        cpp_sources=[cpp_src],
        cuda_sources=[hip_src],
        functions=["mxfp4_quant_small_raw", "mxfp4_shuffle_e8m0"],
        extra_cflags=["-std=c++20"],
        extra_cuda_cflags=[f"--offload-arch={rocm_arch}", "-std=c++20"],
        verbose=False,
    )
    return _NATIVE_RUNTIME


def _native_dynamic_mxfp4_quant(torch, dtypes, x):
    workspace = _get_native_workspace(torch, x)
    module = _get_native_module(torch)
    m, n = x.shape
    module.mxfp4_quant_small_raw(
        x,
        workspace["a_q_raw"],
        workspace["a_scale_raw"],
        m,
        n,
    )
    return (
        workspace["a_q_raw"].view(dtypes.fp4x2),
        workspace["a_scale_raw"].view(dtypes.fp8_e8m0),
    )


def _native_e8m0_shuffle(torch, src):
    src_u8 = src.view(torch.uint8) if src.dtype != torch.uint8 else src
    workspace = _get_shuffle_workspace(torch, src_u8)
    module = _get_native_module(torch)
    workspace["dst_raw"].fill_(127)
    module.mxfp4_shuffle_e8m0(
        src_u8,
        workspace["dst_raw"],
        src_u8.shape[0],
        src_u8.shape[1],
        workspace["n_pad"],
    )
    if src.dtype == torch.uint8:
        return workspace["dst_raw"]
    return workspace["dst_raw"].view(src.dtype)


def _install_specialized_quant_hooks(
    torch,
    dtypes,
    quant_mod,
    fused_quant_mod,
    fp4_utils,
    fused_moe_mod,
):
    global _PATCHED_QUANT, _ORIG_DYNAMIC_MXFP4_QUANT, _ORIG_E8M0_SHUFFLE
    global _ORIG_FUSED_DYNAMIC_MXFP4_QUANT_MOE_SORT
    global _ORIG_GET_QUANT, _NATIVE_QUANT_ERROR, _NATIVE_SHUFFLE_ERROR
    global _NATIVE_QUANT_INFO_PRINTED
    if _PATCHED_QUANT:
        return

    _ORIG_DYNAMIC_MXFP4_QUANT = quant_mod.dynamic_mxfp4_quant
    _ORIG_FUSED_DYNAMIC_MXFP4_QUANT_MOE_SORT = (
        fused_quant_mod.fused_dynamic_mxfp4_quant_moe_sort
    )
    _ORIG_E8M0_SHUFFLE = fp4_utils.e8m0_shuffle
    _ORIG_GET_QUANT = getattr(fused_moe_mod, "get_quant", None)

    def patched_dynamic_mxfp4_quant(x, *args, **kwargs):
        global _NATIVE_QUANT_ERROR
        if (
            not _NATIVE_QUANT_ERROR
            and getattr(x, "is_cuda", False)
            and getattr(x, "dtype", None) == torch.bfloat16
            and getattr(x, "dim", lambda: 0)() == 2
        ):
            m, n = x.shape
            if _native_quant_enabled(m, n):
                try:
                    return _native_dynamic_mxfp4_quant(torch, dtypes, x)
                except Exception:
                    _NATIVE_QUANT_ERROR = True
        return _ORIG_DYNAMIC_MXFP4_QUANT(x, *args, **kwargs)

    def patched_e8m0_shuffle(src, *args, **kwargs):
        global _NATIVE_SHUFFLE_ERROR
        if (
            not _NATIVE_SHUFFLE_ERROR
            and getattr(src, "is_cuda", False)
            and getattr(src, "dim", lambda: 0)() == 2
        ):
            m, n_valid = src.shape
            if _native_shuffle_enabled(m, n_valid):
                try:
                    return _native_e8m0_shuffle(torch, src)
                except Exception:
                    _NATIVE_SHUFFLE_ERROR = True
        return _ORIG_E8M0_SHUFFLE(src, *args, **kwargs)

    def patched_fused_dynamic_mxfp4_quant_moe_sort(
        x,
        sorted_ids,
        num_valid_ids,
        token_num,
        topk,
        block_size=32,
        scaling_mode="even",
    ):
        global _NATIVE_QUANT_ERROR, _NATIVE_QUANT_INFO_PRINTED
        if (
            not _NATIVE_QUANT_ERROR
            and topk == 1
            and getattr(x, "is_cuda", False)
            and getattr(x, "dtype", None) == torch.bfloat16
            and getattr(x, "dim", lambda: 0)() == 2
        ):
            m, n = x.shape
            if m == token_num and _native_quant_enabled(m, n):
                try:
                    a_q, a_scale = _native_dynamic_mxfp4_quant(torch, dtypes, x)
                    a_scale = fp4_utils.moe_mxfp4_sort(
                        a_scale,
                        sorted_ids=sorted_ids,
                        num_valid_ids=num_valid_ids,
                        token_num=token_num,
                        block_size=block_size,
                    )
                    if not _NATIVE_QUANT_INFO_PRINTED:
                        _NATIVE_QUANT_INFO_PRINTED = True
                        print(
                            "[moe-mxfp4] using native load_inline fused quant+sort replacement for stage1"
                        )
                    return a_q, a_scale
                except Exception:
                    _NATIVE_QUANT_ERROR = True
        return _ORIG_FUSED_DYNAMIC_MXFP4_QUANT_MOE_SORT(
            x,
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=token_num,
            topk=topk,
            block_size=block_size,
            scaling_mode=scaling_mode,
        )

    def patched_get_quant(quant_type):
        base = _ORIG_GET_QUANT(quant_type)
        quant_name = getattr(quant_type, "name", "")
        if quant_name != "per_1x32" and "per_1x32" not in str(quant_type):
            return base

        def wrapped_quant(x, *args, **kwargs):
            global _NATIVE_QUANT_ERROR, _NATIVE_QUANT_INFO_PRINTED
            num_rows = kwargs.get("num_rows")
            if (
                not _NATIVE_QUANT_ERROR
                and (num_rows is None)
                and getattr(x, "is_cuda", False)
                and getattr(x, "dtype", None) == torch.bfloat16
                and getattr(x, "dim", lambda: 0)() == 2
            ):
                m, n = x.shape
                if _native_quant_enabled(m, n):
                    try:
                        result = _native_dynamic_mxfp4_quant(torch, dtypes, x)
                        if not _NATIVE_QUANT_INFO_PRINTED:
                            _NATIVE_QUANT_INFO_PRINTED = True
                            print("[moe-mxfp4] using native load_inline quant for per_1x32")
                        return result
                    except Exception:
                        _NATIVE_QUANT_ERROR = True
            return base(x, *args, **kwargs)

        return wrapped_quant

    quant_mod.dynamic_mxfp4_quant = patched_dynamic_mxfp4_quant
    fused_quant_mod.fused_dynamic_mxfp4_quant_moe_sort = (
        patched_fused_dynamic_mxfp4_quant_moe_sort
    )
    if hasattr(fp4_utils, "dynamic_mxfp4_quant"):
        fp4_utils.dynamic_mxfp4_quant = patched_dynamic_mxfp4_quant
    fp4_utils.e8m0_shuffle = patched_e8m0_shuffle
    if _ORIG_GET_QUANT is not None:
        fused_moe_mod.get_quant = patched_get_quant
    if hasattr(fused_moe_mod, "dynamic_mxfp4_quant"):
        fused_moe_mod.dynamic_mxfp4_quant = patched_dynamic_mxfp4_quant
    if hasattr(fused_moe_mod, "fused_dynamic_mxfp4_quant_moe_sort"):
        fused_moe_mod.fused_dynamic_mxfp4_quant_moe_sort = (
            patched_fused_dynamic_mxfp4_quant_moe_sort
        )
    if hasattr(fused_moe_mod, "e8m0_shuffle"):
        fused_moe_mod.e8m0_shuffle = patched_e8m0_shuffle
    _PATCHED_QUANT = True


def _get_runtime():
    global _RUNTIME
    if _RUNTIME is None:
        _ensure_a4w4_tuned_override()
        _ensure_fmoe_tuned_override()
        import torch
        import aiter

        from aiter import ActivationType, QuantType, dtypes

        quant_mod = importlib.import_module("aiter.ops.triton.quant")
        fused_quant_mod = importlib.import_module("aiter.ops.triton.quant.fused_mxfp4_quant")
        fp4_utils = importlib.import_module("aiter.utility.fp4_utils")
        fused_moe_mod = importlib.import_module("aiter.fused_moe")

        _install_specialized_quant_hooks(
            torch,
            dtypes,
            quant_mod,
            fused_quant_mod,
            fp4_utils,
            fused_moe_mod,
        )
        _RUNTIME = (
            torch,
            ActivationType,
            QuantType,
            fused_moe_mod.fused_moe,
        )
    return _RUNTIME


def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        gate_up_weight,
        down_weight,
        gate_up_weight_scale,
        down_weight_scale,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data
    del gate_up_weight
    del down_weight
    del gate_up_weight_scale
    del down_weight_scale

    _, ActivationType, QuantType, fused_moe = _get_runtime()

    hidden_states = hidden_states.contiguous()
    topk_weights = topk_weights.contiguous()
    topk_ids = topk_ids.contiguous()

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    return fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )
scrolls · 774 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