submission 589134
Purple rain · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 505 lines, June 9 Researcher Reciprocity License v1.0.
submission_amd_moe_mxfp4_hip_fused.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-589134?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:1427acc0ad3bd553a730bdd60ed139b3001a0ec63307792f27ffcf2aa7f3aa15
license declaredunknown
license concludedunknown
authorsPurple rain
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float smem[];Kernel source
submission_amd_moe_mxfp4_hip_fused.py505 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import os
import threading
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
_HIP_LOCK = threading.Lock()
_HIP_MODULE = None
_HIP_BUILD_ERROR = None
_MODE_ENV = "MOE_MXFP4_HIP_FUSED_MODE" # base | hip
_DEFAULT_MODE = "base"
CPP_WRAPPER = r"""
#include <torch/extension.h>
torch::Tensor hip_moe_forward(
torch::Tensor hidden_states,
torch::Tensor gate_up_weight,
torch::Tensor down_weight,
torch::Tensor gate_up_weight_scale,
torch::Tensor down_weight_scale,
torch::Tensor topk_weights,
torch::Tensor topk_ids,
int64_t d_expert);
"""
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <hip/hip_fp16.h>
#include <cstdint>
#include <vector>
namespace {
constexpr int kThreads = 256;
constexpr int kScaleBlock = 32;
__constant__ float kFp4Lut[16] = {
0.0f, 0.5f, 1.0f, 1.5f,
2.0f, 3.0f, 4.0f, 6.0f,
-0.0f, -0.5f, -1.0f, -1.5f,
-2.0f, -3.0f, -4.0f, -6.0f
};
__device__ __forceinline__ float fp4_nibble_to_float(uint8_t nibble) {
return kFp4Lut[nibble & 0x0F];
}
__device__ __forceinline__ float e8m0_to_float(uint8_t e8m0_val) {
// Match aiter.utility.fp4_utils.e8m0_to_f32:
// normal: exponent bits in fp32
// 0x00 : smallest positive value used by the kernel path
// 0xFF : NaN payload
uint32_t bits = static_cast<uint32_t>(e8m0_val) << 23;
if (e8m0_val == 0) {
bits = 0x00400000u;
} else if (e8m0_val == 0xFF) {
bits = 0x7F800001u;
}
return __uint_as_float(bits);
}
__device__ __forceinline__ float fp4x2_at(const uint8_t* packed_row, int col) {
uint8_t packed = packed_row[col >> 1];
uint8_t nib = (col & 1) ? static_cast<uint8_t>(packed >> 4)
: static_cast<uint8_t>(packed & 0x0F);
return fp4_nibble_to_float(nib);
}
__device__ __forceinline__ float silu(float x) {
return x / (1.0f + expf(-x));
}
__device__ __forceinline__ float warp_sum(float v) {
for (int offset = warpSize / 2; offset > 0; offset >>= 1) {
v += __shfl_down(v, offset);
}
return v;
}
__global__ void fused_moe_kernel(
const hip_bfloat16* hidden_states, // [M, d_hidden]
const uint8_t* gate_up_weight, // [E, 2*d_expert_pad, d_hidden_pad/2]
const uint8_t* down_weight, // [E, d_hidden_pad, d_expert_pad/2]
const uint8_t* gate_up_weight_scale, // [E, 2*d_expert_pad, d_hidden_pad/32]
const uint8_t* down_weight_scale, // [E, d_hidden_pad, d_expert_pad/32]
const float* topk_weights, // [M, topk]
const int32_t* topk_ids, // [M, topk]
float* output_accum, // [M, d_hidden]
int32_t m,
int32_t topk,
int32_t num_experts,
int32_t d_hidden,
int32_t d_hidden_pad,
int32_t d_expert,
int32_t d_expert_pad,
int32_t gateup_k_pack,
int32_t down_k_pack,
int32_t gateup_scale_k,
int32_t down_scale_k) {
int32_t token = static_cast<int32_t>(blockIdx.x);
int32_t k_idx = static_cast<int32_t>(blockIdx.y);
if (token >= m || k_idx >= topk) {
return;
}
int32_t tid = static_cast<int32_t>(threadIdx.x);
int32_t lane = tid % warpSize;
int32_t warp_id = tid / warpSize;
int32_t warps_per_block = static_cast<int32_t>(blockDim.x) / warpSize;
int64_t route_offset = static_cast<int64_t>(token) * topk + k_idx;
int32_t expert = topk_ids[route_offset];
if (expert < 0 || expert >= num_experts) {
return;
}
float route_weight = topk_weights[route_offset];
extern __shared__ float smem[];
float* x_sh = smem; // [d_hidden]
float* inter_sh = smem + d_hidden; // [d_expert]
const hip_bfloat16* x_row =
hidden_states + static_cast<int64_t>(token) * d_hidden;
for (int32_t col = tid; col < d_hidden; col += static_cast<int32_t>(blockDim.x)) {
x_sh[col] = static_cast<float>(x_row[col]);
}
__syncthreads();
const int64_t gateup_expert_stride =
static_cast<int64_t>(2 * d_expert_pad) * gateup_k_pack;
const int64_t gateup_scale_expert_stride =
static_cast<int64_t>(2 * d_expert_pad) * gateup_scale_k;
const int64_t down_expert_stride =
static_cast<int64_t>(d_hidden_pad) * down_k_pack;
const int64_t down_scale_expert_stride =
static_cast<int64_t>(d_hidden_pad) * down_scale_k;
const int64_t gateup_base = static_cast<int64_t>(expert) * gateup_expert_stride;
const int64_t gateup_scale_base =
static_cast<int64_t>(expert) * gateup_scale_expert_stride;
for (int32_t j = warp_id; j < d_expert; j += warps_per_block) {
const uint8_t* gate_row = gate_up_weight + gateup_base + static_cast<int64_t>(j) * gateup_k_pack;
const uint8_t* up_row =
gate_up_weight + gateup_base + static_cast<int64_t>(d_expert_pad + j) * gateup_k_pack;
const uint8_t* gate_scale_row =
gate_up_weight_scale + gateup_scale_base + static_cast<int64_t>(j) * gateup_scale_k;
const uint8_t* up_scale_row =
gate_up_weight_scale + gateup_scale_base + static_cast<int64_t>(d_expert_pad + j) * gateup_scale_k;
float gate_sum = 0.0f;
float up_sum = 0.0f;
for (int32_t col = lane; col < d_hidden; col += warpSize) {
float x = x_sh[col];
float gate_w =
fp4x2_at(gate_row, col) * e8m0_to_float(gate_scale_row[col / kScaleBlock]);
float up_w =
fp4x2_at(up_row, col) * e8m0_to_float(up_scale_row[col / kScaleBlock]);
gate_sum += x * gate_w;
up_sum += x * up_w;
}
gate_sum = warp_sum(gate_sum);
up_sum = warp_sum(up_sum);
if (lane == 0) {
inter_sh[j] = silu(gate_sum) * up_sum;
}
}
__syncthreads();
const int64_t down_base = static_cast<int64_t>(expert) * down_expert_stride;
const int64_t down_scale_base =
static_cast<int64_t>(expert) * down_scale_expert_stride;
for (int32_t out_col = warp_id; out_col < d_hidden; out_col += warps_per_block) {
float out_sum = 0.0f;
if (out_col < d_hidden_pad) {
const uint8_t* down_row =
down_weight + down_base + static_cast<int64_t>(out_col) * down_k_pack;
const uint8_t* down_scale_row =
down_weight_scale + down_scale_base + static_cast<int64_t>(out_col) * down_scale_k;
for (int32_t j = lane; j < d_expert; j += warpSize) {
float w =
fp4x2_at(down_row, j) * e8m0_to_float(down_scale_row[j / kScaleBlock]);
out_sum += inter_sh[j] * w;
}
}
out_sum = warp_sum(out_sum);
if (lane == 0) {
atomicAdd(
output_accum + static_cast<int64_t>(token) * d_hidden + out_col,
route_weight * out_sum);
}
}
}
__global__ void cast_float_to_bf16_kernel(
const float* in,
hip_bfloat16* out,
int64_t total) {
int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx < total) {
out[idx] = static_cast<hip_bfloat16>(in[idx]);
}
}
inline void check_hip_error(const char* where) {
hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}
} // namespace
torch::Tensor hip_moe_forward(
torch::Tensor hidden_states,
torch::Tensor gate_up_weight,
torch::Tensor down_weight,
torch::Tensor gate_up_weight_scale,
torch::Tensor down_weight_scale,
torch::Tensor topk_weights,
torch::Tensor topk_ids,
int64_t d_expert) {
TORCH_CHECK(hidden_states.is_cuda(), "hidden_states must be CUDA/HIP tensor");
TORCH_CHECK(gate_up_weight.is_cuda(), "gate_up_weight must be CUDA/HIP tensor");
TORCH_CHECK(down_weight.is_cuda(), "down_weight must be CUDA/HIP tensor");
TORCH_CHECK(gate_up_weight_scale.is_cuda(), "gate_up_weight_scale must be CUDA/HIP tensor");
TORCH_CHECK(down_weight_scale.is_cuda(), "down_weight_scale must be CUDA/HIP tensor");
TORCH_CHECK(topk_weights.is_cuda(), "topk_weights must be CUDA/HIP tensor");
TORCH_CHECK(topk_ids.is_cuda(), "topk_ids must be CUDA/HIP tensor");
TORCH_CHECK(hidden_states.scalar_type() == torch::kBFloat16, "hidden_states must be bfloat16");
TORCH_CHECK(gate_up_weight.element_size() == 1, "gate_up_weight must be 1-byte packed fp4x2");
TORCH_CHECK(down_weight.element_size() == 1, "down_weight must be 1-byte packed fp4x2");
TORCH_CHECK(gate_up_weight_scale.element_size() == 1, "gate_up_weight_scale must be 1-byte e8m0");
TORCH_CHECK(down_weight_scale.element_size() == 1, "down_weight_scale must be 1-byte e8m0");
TORCH_CHECK(topk_weights.scalar_type() == torch::kFloat32, "topk_weights must be float32");
TORCH_CHECK(topk_ids.scalar_type() == torch::kInt32, "topk_ids must be int32");
TORCH_CHECK(hidden_states.dim() == 2, "hidden_states must be [M, d_hidden]");
TORCH_CHECK(gate_up_weight.dim() == 3, "gate_up_weight must be [E, 2*d_expert_pad, d_hidden_pad/2]");
TORCH_CHECK(down_weight.dim() == 3, "down_weight must be [E, d_hidden_pad, d_expert_pad/2]");
TORCH_CHECK(gate_up_weight_scale.dim() == 3, "gate_up_weight_scale must be [E, 2*d_expert_pad, d_hidden_pad/32]");
TORCH_CHECK(down_weight_scale.dim() == 3, "down_weight_scale must be [E, d_hidden_pad, d_expert_pad/32]");
TORCH_CHECK(topk_weights.dim() == 2, "topk_weights must be [M, topk]");
TORCH_CHECK(topk_ids.dim() == 2, "topk_ids must be [M, topk]");
hidden_states = hidden_states.contiguous();
gate_up_weight = gate_up_weight.contiguous();
down_weight = down_weight.contiguous();
gate_up_weight_scale = gate_up_weight_scale.contiguous();
down_weight_scale = down_weight_scale.contiguous();
topk_weights = topk_weights.contiguous();
topk_ids = topk_ids.contiguous();
const int64_t m = hidden_states.size(0);
const int64_t d_hidden = hidden_states.size(1);
const int64_t topk = topk_ids.size(1);
const int64_t num_experts = gate_up_weight.size(0);
const int64_t d_hidden_pad = down_weight.size(1);
const int64_t d_expert_pad = down_weight.size(2) * 2;
TORCH_CHECK(topk_weights.size(0) == m, "topk_weights M mismatch");
TORCH_CHECK(topk_ids.size(0) == m, "topk_ids M mismatch");
TORCH_CHECK(topk_weights.size(1) == topk, "topk shape mismatch");
TORCH_CHECK(gate_up_weight.size(0) == num_experts, "gate_up_weight E mismatch");
TORCH_CHECK(down_weight.size(0) == num_experts, "down_weight E mismatch");
TORCH_CHECK(gate_up_weight_scale.size(0) == num_experts, "gate_up_weight_scale E mismatch");
TORCH_CHECK(down_weight_scale.size(0) == num_experts, "down_weight_scale E mismatch");
TORCH_CHECK(gate_up_weight.size(1) == 2 * d_expert_pad, "gate_up_weight row mismatch vs d_expert_pad");
TORCH_CHECK(gate_up_weight_scale.size(1) == 2 * d_expert_pad, "gate_up_weight_scale row mismatch");
TORCH_CHECK(gate_up_weight.size(2) * 2 == d_hidden_pad, "gate_up_weight K mismatch vs d_hidden_pad");
TORCH_CHECK(gate_up_weight_scale.size(2) == d_hidden_pad / kScaleBlock, "gate_up_weight_scale K mismatch");
TORCH_CHECK(down_weight_scale.size(1) == d_hidden_pad, "down_weight_scale row mismatch");
TORCH_CHECK(down_weight_scale.size(2) == d_expert_pad / kScaleBlock, "down_weight_scale K mismatch");
TORCH_CHECK(d_hidden <= d_hidden_pad, "d_hidden must be <= d_hidden_pad");
TORCH_CHECK(d_expert > 0 && d_expert <= d_expert_pad, "d_expert out of valid range");
TORCH_CHECK(d_hidden_pad % kScaleBlock == 0, "d_hidden_pad must be divisible by 32");
TORCH_CHECK(d_expert_pad % kScaleBlock == 0, "d_expert_pad must be divisible by 32");
auto out_accum = torch::zeros({m, d_hidden}, hidden_states.options().dtype(torch::kFloat32));
auto output = torch::empty({m, d_hidden}, hidden_states.options().dtype(torch::kBFloat16));
if (m == 0 || d_hidden == 0 || topk == 0) {
return output.zero_();
}
size_t smem_bytes = static_cast<size_t>(d_hidden + d_expert) * sizeof(float);
hipLaunchKernelGGL(
fused_moe_kernel,
dim3(static_cast<unsigned int>(m), static_cast<unsigned int>(topk)),
dim3(kThreads),
smem_bytes,
0,
reinterpret_cast<const hip_bfloat16*>(hidden_states.data_ptr<at::BFloat16>()),
reinterpret_cast<const uint8_t*>(gate_up_weight.data_ptr()),
reinterpret_cast<const uint8_t*>(down_weight.data_ptr()),
reinterpret_cast<const uint8_t*>(gate_up_weight_scale.data_ptr()),
reinterpret_cast<const uint8_t*>(down_weight_scale.data_ptr()),
topk_weights.data_ptr<float>(),
topk_ids.data_ptr<int32_t>(),
out_accum.data_ptr<float>(),
static_cast<int32_t>(m),
static_cast<int32_t>(topk),
static_cast<int32_t>(num_experts),
static_cast<int32_t>(d_hidden),
static_cast<int32_t>(d_hidden_pad),
static_cast<int32_t>(d_expert),
static_cast<int32_t>(d_expert_pad),
static_cast<int32_t>(gate_up_weight.size(2)),
static_cast<int32_t>(down_weight.size(2)),
static_cast<int32_t>(gate_up_weight_scale.size(2)),
static_cast<int32_t>(down_weight_scale.size(2)));
check_hip_error("fused_moe_kernel");
int64_t total = m * d_hidden;
int32_t blocks = static_cast<int32_t>((total + kThreads - 1) / kThreads);
hipLaunchKernelGGL(
cast_float_to_bf16_kernel,
dim3(static_cast<unsigned int>(blocks)),
dim3(kThreads),
0,
0,
out_accum.data_ptr<float>(),
reinterpret_cast<hip_bfloat16*>(output.data_ptr<at::BFloat16>()),
total);
check_hip_error("cast_float_to_bf16_kernel");
return output;
}
"""
def _get_hip_module():
global _HIP_MODULE, _HIP_BUILD_ERROR
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
with _HIP_LOCK:
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
try:
os.environ.setdefault("CXX", "clang++")
_HIP_MODULE = load_inline(
name="moe_mxfp4_hip_fused_v1",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[HIP_SRC],
functions=["hip_moe_forward"],
verbose=False,
extra_cuda_cflags=["-O3", "-std=c++20"],
)
except Exception as e: # pragma: no cover
_HIP_BUILD_ERROR = e
raise RuntimeError(f"HIP inline build failed: {e}") from e
return _HIP_MODULE
def _get_mode() -> str:
mode = (os.getenv(_MODE_ENV, _DEFAULT_MODE) or _DEFAULT_MODE).strip().lower()
if mode in {"base", "hip"}:
return mode
return _DEFAULT_MODE
def _normalize_scale_layout(
scale: torch.Tensor,
expected_shape: tuple[int, int, int],
) -> torch.Tensor:
"""
Normalize e8m0 scale layout to [E, rows, scale_k].
Runtime may provide either native 3D or flattened 2D buffers.
"""
e, rows, k_blocks = expected_shape
need = e * rows * k_blocks
s = scale.contiguous()
if s.dim() == 3:
if tuple(s.shape) == expected_shape:
return s
flat = s.reshape(-1)
else:
flat = s.reshape(-1)
if flat.numel() < need:
raise RuntimeError(
f"scale buffer too small: got {flat.numel()} elems, need {need} for shape {expected_shape}"
)
return flat[:need].view(expected_shape).contiguous()
def _run_base_fused_moe(data: input_t) -> output_t:
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
(
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
hidden_pad = int(config["d_hidden_pad"] - config["d_hidden"])
intermediate_pad = int(config["d_expert_pad"] - config["d_expert"])
return fused_moe(
hidden_states.contiguous(),
gate_up_weight_shuffled.contiguous(),
down_weight_shuffled.contiguous(),
topk_weights.to(torch.float32).contiguous(),
topk_ids.to(torch.int32).contiguous(),
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=gate_up_weight_scale_shuffled.contiguous(),
w2_scale=down_weight_scale_shuffled.contiguous(),
a1_scale=None,
a2_scale=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
def _run_hip_fused(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
module = _get_hip_module()
n_experts = int(config["n_routed_experts"] + config["n_shared_experts"])
d_hidden_pad = int(config["d_hidden_pad"])
d_expert_pad = int(config["d_expert_pad"])
gate_up_weight_scale = _normalize_scale_layout(
gate_up_weight_scale,
(n_experts, 2 * d_expert_pad, d_hidden_pad // 32),
)
down_weight_scale = _normalize_scale_layout(
down_weight_scale,
(n_experts, d_hidden_pad, d_expert_pad // 32),
)
return module.hip_moe_forward(
hidden_states.contiguous(),
gate_up_weight.contiguous(),
down_weight.contiguous(),
gate_up_weight_scale.contiguous(),
down_weight_scale.contiguous(),
topk_weights.to(torch.float32).contiguous(),
topk_ids.to(torch.int32).contiguous(),
int(config["d_expert"]),
)
def custom_kernel(data: input_t) -> output_t:
if _get_mode() == "hip":
return _run_hip_fused(data)
return _run_base_fused_moe(data)
scrolls · 505 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