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