submission 723815
Op Gup · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 714 lines, June 9 Researcher Reciprocity License v1.0.
solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-723815?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:9a72349cdfcbf0ed03d82fe00171861227e6c2090f31ee75877cedcba72831b5
license declaredunknown
license concludedunknown
authorsOp Gup
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
solution.py — DeepSeek-R1 MXFP4 MoE fused kernel, single-file submission.shared-memory
__shared__ int32_t smem[MAX_EXPERTS];Kernel source
solution.py714 lines
"""
solution.py — DeepSeek-R1 MXFP4 MoE fused kernel, single-file submission.
Compiles HIP kernels at first import via torch.utils.cpp_extension.load_inline.
Compiled .so is cached in /tmp/moe_mxfp4_cache/ so subsequent runs skip compilation.
Usage:
from solution import fused_moe_mxfp4
output = fused_moe_mxfp4(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)
"""
import os
import torch
from torch import Tensor
from typing import Dict, Any, Optional, Tuple
import torch.nn.functional as F
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
# ── Tile shapes — tune these for MI355X ──────────────────────────────────────
TILE_M = int(os.environ.get("TILE_M", "32"))
TILE_N = int(os.environ.get("TILE_N", "128"))
TILE_K = int(os.environ.get("TILE_K", "128"))
TILE_M_S2 = int(os.environ.get("TILE_M_S2", "32"))
TILE_N_S2 = int(os.environ.get("TILE_N_S2", "128"))
TILE_K_S2 = int(os.environ.get("TILE_K_S2", "128"))
GPU_ARCH = os.environ.get("GPU_ARCH", "gfx942") # MI355X = gfx942
# ══════════════════════════════════════════════════════════════════════════════
# Kernel source strings
# ══════════════════════════════════════════════════════════════════════════════
_ROUTING_SRC = r"""
#include <hip/hip_runtime.h>
#include <cstdint>
constexpr int WARP_SIZE_R = 64;
constexpr int THREADS_R = 512;
constexpr int MAX_EXPERTS = 512;
__global__ void kernel_count_tokens(
const int32_t* __restrict__ topk_ids,
int32_t* __restrict__ expert_cnt,
int32_t M, int32_t total_top_k, int32_t n_total_experts
) {
__shared__ int32_t smem[MAX_EXPERTS];
for (int i = threadIdx.x; i < n_total_experts; i += blockDim.x) smem[i] = 0;
__syncthreads();
int total = M * total_top_k;
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < total;
idx += gridDim.x * blockDim.x) {
int eid = topk_ids[idx];
if (eid >= 0 && eid < n_total_experts) atomicAdd(&smem[eid], 1);
}
__syncthreads();
for (int i = threadIdx.x; i < n_total_experts; i += blockDim.x)
atomicAdd(&expert_cnt[i], smem[i]);
}
__global__ void kernel_prefix_sum(
const int32_t* __restrict__ expert_cnt,
int32_t* __restrict__ expert_start,
int32_t n_total_experts
) {
if (threadIdx.x != 0) return;
expert_start[0] = 0;
for (int i = 0; i < n_total_experts; ++i)
expert_start[i+1] = expert_start[i] + expert_cnt[i];
}
__global__ void kernel_scatter_tokens(
const int32_t* __restrict__ topk_ids,
const int32_t* __restrict__ expert_start,
int32_t* __restrict__ cursors,
int32_t* __restrict__ sorted_ids,
int32_t M, int32_t total_top_k, int32_t n_total_experts
) {
int total = M * total_top_k;
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < total;
idx += gridDim.x * blockDim.x) {
int token_idx = idx / total_top_k;
int eid = topk_ids[idx];
if (eid < 0 || eid >= n_total_experts) continue;
int slot = atomicAdd(&cursors[eid], 1);
sorted_ids[expert_start[eid] + slot] = token_idx;
}
}
"""
_STAGE1_SRC = rf"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
#include <cmath>
#define TILE_M {TILE_M}
#define TILE_K {TILE_K}
#define BLOCK_S1 128
using fp4x2_t = uint8_t;
using e8m0_t = uint8_t;
__device__ __forceinline__ float e8m0_to_f32(e8m0_t s) {{
return __builtin_amdgcn_ldexp(1.0f, (int)s - 127);
}}
__device__ __forceinline__ float fp4_to_f32(uint8_t n) {{
int sign = (n>>3)&1, exp = (n>>1)&3, mant = n&1;
float v = (exp==0) ? mant*0.5f : (1.f+mant*0.5f)*__builtin_amdgcn_ldexp(1.f,exp-1);
return sign ? -v : v;
}}
__device__ __forceinline__ float silu(float x) {{
return x / (1.f + __expf(-x));
}}
__global__ void __launch_bounds__(BLOCK_S1, 4)
kernel_stage1(
const __hip_bfloat16* __restrict__ hidden,
const fp4x2_t* __restrict__ gate_up_w,
const e8m0_t* __restrict__ gate_up_scale,
const int32_t* __restrict__ sorted_token_ids,
const int32_t* __restrict__ expert_start,
const int32_t* __restrict__ expert_cnt,
__hip_bfloat16* __restrict__ intermediate,
int32_t M, int32_t d_hidden, int32_t d_hidden_pad,
int32_t d_expert_pad, int32_t n_total_experts
) {{
const int eid = blockIdx.x;
const int token_tile = blockIdx.y;
if (eid >= n_total_experts) return;
const int n_tok = expert_cnt[eid];
if (n_tok == 0) return;
const int tok_start = expert_start[eid];
const int local_idx = token_tile * TILE_M + threadIdx.x;
if (local_idx >= n_tok) return;
const int token_id = sorted_token_ids[tok_start + local_idx];
const __hip_bfloat16* x = hidden + (int64_t)token_id * d_hidden;
// Weight strides
const int64_t w_stride_e = (int64_t)2 * d_expert_pad * (d_hidden_pad/2);
const int64_t s_stride_e = (int64_t)2 * d_expert_pad * (d_hidden_pad/32);
const fp4x2_t* we = gate_up_w + eid * w_stride_e;
const e8m0_t* se = gate_up_scale + eid * s_stride_e;
__hip_bfloat16* out = intermediate + (int64_t)(tok_start + local_idx) * d_expert_pad;
for (int n = 0; n < d_expert_pad; ++n) {{
float gate_acc = 0.f, up_acc = 0.f;
const fp4x2_t* wg = we + (int64_t)n * (d_hidden_pad/2);
const fp4x2_t* wu = we + (int64_t)(d_expert_pad+n)* (d_hidden_pad/2);
const e8m0_t* sg = se + (int64_t)n * (d_hidden_pad/32);
const e8m0_t* su = se + (int64_t)(d_expert_pad+n)* (d_hidden_pad/32);
for (int k = 0; k < d_hidden; k += 32) {{
float sg_f = e8m0_to_f32(sg[k/32]);
float su_f = e8m0_to_f32(su[k/32]);
for (int kk = 0; kk < 32 && (k+kk) < d_hidden; kk += 2) {{
float x0 = __bfloat162float(x[k+kk]);
float x1 = __bfloat162float(x[k+kk+1]);
fp4x2_t pg = wg[(k+kk)/2];
fp4x2_t pu = wu[(k+kk)/2];
gate_acc += x0*fp4_to_f32(pg&0xF)*sg_f + x1*fp4_to_f32(pg>>4)*sg_f;
up_acc += x0*fp4_to_f32(pu&0xF)*su_f + x1*fp4_to_f32(pu>>4)*su_f;
}}
}}
out[n] = __float2bfloat16(silu(gate_acc) * up_acc);
}}
}}
"""
_STAGE2_SRC = rf"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
#include <cmath>
#define TILE_M_S2 {TILE_M_S2}
#define BLOCK_S2 128
using fp4x2_t = uint8_t;
using e8m0_t = uint8_t;
__device__ __forceinline__ float e8m0_to_f32_s2(e8m0_t s) {{
return __builtin_amdgcn_ldexp(1.0f, (int)s - 127);
}}
__device__ __forceinline__ float fp4_to_f32_s2(uint8_t n) {{
int sign=(n>>3)&1, exp=(n>>1)&3, mant=n&1;
float v=(exp==0)?mant*0.5f:(1.f+mant*0.5f)*__builtin_amdgcn_ldexp(1.f,exp-1);
return sign?-v:v;
}}
// bf16 atomic add via CAS loop
__device__ __forceinline__ void atomic_add_bf16(
__hip_bfloat16* addr, float val
) {{
uint32_t* base = (uint32_t*)((uintptr_t)addr & ~3ULL);
bool hi = ((uintptr_t)addr & 2) != 0;
uint32_t old = *base, assumed, next;
do {{
assumed = old;
float cur = hi
? __bfloat162float((__hip_bfloat16)(assumed >> 16))
: __bfloat162float((__hip_bfloat16)(assumed & 0xFFFF));
uint16_t upd = ((__hip_bfloat16_raw)__float2bfloat16(cur + val)).x;
next = hi ? ((assumed & 0xFFFF) | ((uint32_t)upd << 16))
: ((assumed & 0xFFFF0000) | upd);
old = atomicCAS(base, assumed, next);
}} while (old != assumed);
}}
__global__ void __launch_bounds__(BLOCK_S2, 4)
kernel_stage2(
const __hip_bfloat16* __restrict__ intermediate,
const fp4x2_t* __restrict__ down_w,
const e8m0_t* __restrict__ down_scale,
const int32_t* __restrict__ sorted_token_ids,
const int32_t* __restrict__ expert_start,
const int32_t* __restrict__ expert_cnt,
const float* __restrict__ topk_weights,
const int32_t* __restrict__ topk_ids,
__hip_bfloat16* __restrict__ output,
int32_t M, int32_t d_hidden, int32_t d_hidden_pad,
int32_t d_expert_pad, int32_t n_total_experts,
int32_t n_routed_experts, int32_t total_top_k
) {{
const int eid = blockIdx.x;
const int token_tile = blockIdx.y;
if (eid >= n_total_experts) return;
const int n_tok = expert_cnt[eid];
if (n_tok == 0) return;
const int tok_start = expert_start[eid];
const int local_idx = token_tile * TILE_M_S2 + threadIdx.x;
if (local_idx >= n_tok) return;
const int active_idx = tok_start + local_idx;
const int token_id = sorted_token_ids[active_idx];
const bool is_shared = (eid >= n_routed_experts);
// Find routing weight
float rw = 1.0f;
if (!is_shared) {{
for (int k = 0; k < total_top_k; ++k) {{
if (topk_ids[token_id * total_top_k + k] == eid) {{
rw = topk_weights[token_id * total_top_k + k];
break;
}}
}}
}}
const __hip_bfloat16* interm = intermediate + (int64_t)active_idx * d_expert_pad;
const int64_t w_stride_e = (int64_t)d_hidden_pad * (d_expert_pad/2);
const int64_t s_stride_e = (int64_t)d_hidden_pad * (d_expert_pad/32);
const fp4x2_t* we = down_w + (int64_t)eid * w_stride_e;
const e8m0_t* se = down_scale + (int64_t)eid * s_stride_e;
__hip_bfloat16* out = output + (int64_t)token_id * d_hidden;
for (int n = 0; n < d_hidden; ++n) {{
float acc = 0.f;
const fp4x2_t* wr = we + (int64_t)n * (d_expert_pad/2);
const e8m0_t* sr = se + (int64_t)n * (d_expert_pad/32);
for (int k = 0; k < d_expert_pad; k += 32) {{
float sf = e8m0_to_f32_s2(sr[k/32]);
for (int kk = 0; kk < 32 && (k+kk) < d_expert_pad; kk += 2) {{
float a0 = __bfloat162float(interm[k+kk]);
float a1 = __bfloat162float(interm[k+kk+1]);
fp4x2_t pw = wr[(k+kk)/2];
acc += a0*fp4_to_f32_s2(pw&0xF)*sf + a1*fp4_to_f32_s2(pw>>4)*sf;
}}
}}
atomic_add_bf16(&out[n], rw * acc);
}}
}}
"""
_CPP_BINDING = rf"""
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <cstdint>
#include <hip/hip_runtime.h>
// ── Kernel declarations ────────────────────────────────────────────────────
void kernel_count_tokens(
const int32_t*, int32_t*, int32_t, int32_t, int32_t); // declared in routing
void kernel_prefix_sum(
const int32_t*, int32_t*, int32_t);
void kernel_scatter_tokens(
const int32_t*, const int32_t*, int32_t*, int32_t*,
int32_t, int32_t, int32_t);
void kernel_stage1(
const __hip_bfloat16*, const uint8_t*, const uint8_t*,
const int32_t*, const int32_t*, const int32_t*,
__hip_bfloat16*,
int32_t, int32_t, int32_t, int32_t, int32_t);
void kernel_stage2(
const __hip_bfloat16*, const uint8_t*, const uint8_t*,
const int32_t*, const int32_t*, const int32_t*,
const float*, const int32_t*, __hip_bfloat16*,
int32_t, int32_t, int32_t, int32_t, int32_t, int32_t, int32_t);
// ── routing_dispatch ──────────────────────────────────────────────────────
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>
routing_dispatch(torch::Tensor topk_ids, int64_t n_total_experts, int64_t M) {{
auto dev = topk_ids.device();
auto oi = torch::TensorOptions().dtype(torch::kInt32).device(dev);
int64_t ttk = topk_ids.size(1);
auto sorted = torch::empty({{M * ttk}}, oi);
auto start = torch::zeros({{n_total_experts + 1}}, oi);
auto cnt = torch::zeros({{n_total_experts}}, oi);
auto cursors = torch::zeros({{n_total_experts}}, oi);
int total = (int)(M * ttk);
int blocks = min((total + 511)/512, 256);
hipLaunchKernelGGL(kernel_count_tokens, dim3(blocks), dim3(512), 0, 0,
topk_ids.data_ptr<int32_t>(), cnt.data_ptr<int32_t>(),
(int32_t)M, (int32_t)ttk, (int32_t)n_total_experts);
hipLaunchKernelGGL(kernel_prefix_sum, dim3(1), dim3(64), 0, 0,
cnt.data_ptr<int32_t>(), start.data_ptr<int32_t>(), (int32_t)n_total_experts);
hipLaunchKernelGGL(kernel_scatter_tokens, dim3(blocks), dim3(512), 0, 0,
topk_ids.data_ptr<int32_t>(), start.data_ptr<int32_t>(),
cursors.data_ptr<int32_t>(), sorted.data_ptr<int32_t>(),
(int32_t)M, (int32_t)ttk, (int32_t)n_total_experts);
return {{sorted, start, cnt}};
}}
// ── gemm_stage1 ───────────────────────────────────────────────────────────
torch::Tensor gemm_stage1(
torch::Tensor hidden, torch::Tensor gate_up_w, torch::Tensor gate_up_s,
torch::Tensor sorted, torch::Tensor start, torch::Tensor cnt,
int64_t n_active, int64_t n_total_experts,
int64_t d_hidden, int64_t d_hidden_pad, int64_t d_expert_pad
) {{
auto interm = torch::empty({{n_active, d_expert_pad}},
torch::TensorOptions().dtype(torch::kBFloat16).device(hidden.device()));
int M = (int)hidden.size(0);
int tile_m = {TILE_M};
int max_tt = (M + tile_m - 1) / tile_m;
hipLaunchKernelGGL(kernel_stage1,
dim3((int)n_total_experts, max_tt), dim3(128), 0, 0,
(const __hip_bfloat16*)hidden.data_ptr(),
gate_up_w.data_ptr<uint8_t>(), gate_up_s.data_ptr<uint8_t>(),
sorted.data_ptr<int32_t>(), start.data_ptr<int32_t>(), cnt.data_ptr<int32_t>(),
(__hip_bfloat16*)interm.data_ptr(),
M, (int32_t)d_hidden, (int32_t)d_hidden_pad,
(int32_t)d_expert_pad, (int32_t)n_total_experts);
return interm;
}}
// ── gemm_stage2 ───────────────────────────────────────────────────────────
void gemm_stage2(
torch::Tensor interm, torch::Tensor down_w, torch::Tensor down_s,
torch::Tensor sorted, torch::Tensor start, torch::Tensor cnt,
torch::Tensor topk_weights, torch::Tensor topk_ids,
torch::Tensor output,
int64_t n_total_experts, int64_t n_routed_experts,
int64_t d_hidden, int64_t d_hidden_pad, int64_t d_expert_pad
) {{
int M = (int)topk_ids.size(0);
int ttk = (int)topk_ids.size(1);
int tile_m = {TILE_M_S2};
int max_tt = (M + tile_m - 1) / tile_m;
hipLaunchKernelGGL(kernel_stage2,
dim3((int)n_total_experts, max_tt), dim3(128), 0, 0,
(const __hip_bfloat16*)interm.data_ptr(),
down_w.data_ptr<uint8_t>(), down_s.data_ptr<uint8_t>(),
sorted.data_ptr<int32_t>(), start.data_ptr<int32_t>(), cnt.data_ptr<int32_t>(),
topk_weights.data_ptr<float>(), topk_ids.data_ptr<int32_t>(),
(__hip_bfloat16*)output.data_ptr(),
M, (int32_t)d_hidden, (int32_t)d_hidden_pad,
(int32_t)d_expert_pad, (int32_t)n_total_experts,
(int32_t)n_routed_experts, ttk);
}}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {{
m.def("routing_dispatch", &routing_dispatch);
m.def("gemm_stage1", &gemm_stage1);
m.def("gemm_stage2", &gemm_stage2);
}}
"""
# ══════════════════════════════════════════════════════════════════════════════
# Runtime compilation
# ══════════════════════════════════════════════════════════════════════════════
_ext = None
def _load_ext():
global _ext
if _ext is not None:
return _ext
from torch.utils.cpp_extension import load_inline
import tempfile
rocm_path = os.environ.get("ROCM_PATH", "/opt/rocm")
cache_dir = os.environ.get("MOE_CACHE_DIR", "/tmp/moe_mxfp4_cache")
os.makedirs(cache_dir, exist_ok=True)
extra_flags = [
f"--offload-arch={GPU_ARCH}",
"-O3",
"-std=c++17",
"-ffast-math",
f"-DTILE_M={TILE_M}",
f"-DTILE_N={TILE_N}",
f"-DTILE_K={TILE_K}",
f"-DTILE_M_S2={TILE_M_S2}",
f"-DTILE_N_S2={TILE_N_S2}",
f"-DTILE_K_S2={TILE_K_S2}",
]
_ext = load_inline(
name="moe_mxfp4_ext",
cpp_sources=_CPP_BINDING,
cuda_sources=_ROUTING_SRC + "\n" + _STAGE1_SRC + "\n" + _STAGE2_SRC,
extra_cuda_cflags=extra_flags,
extra_cflags=["-O3", "-std=c++17"],
build_directory=cache_dir,
verbose=os.environ.get("MOE_VERBOSE", "0") == "1",
)
return _ext
# ══════════════════════════════════════════════════════════════════════════════
# Config helper
# ══════════════════════════════════════════════════════════════════════════════
def _pad(x: int, align: int = 256) -> int:
return ((x + align - 1) // align) * align
class _Cfg:
__slots__ = ("d_hidden","d_expert","d_hidden_pad","d_expert_pad",
"n_routed_experts","n_shared_experts","n_experts_per_token",
"total_top_k","bs","n_total_experts")
def __init__(self, c):
self.d_hidden = int(c["d_hidden"])
self.d_expert = int(c["d_expert"])
self.d_hidden_pad = int(c["d_hidden_pad"])
self.d_expert_pad = int(c["d_expert_pad"])
self.n_routed_experts = int(c["n_routed_experts"])
self.n_shared_experts = int(c["n_shared_experts"])
self.n_experts_per_token = int(c["n_experts_per_token"])
self.total_top_k = int(c["total_top_k"])
self.bs = int(c.get("bs", c.get("batch_size", 0)))
self.n_total_experts = self.n_routed_experts + self.n_shared_experts
assert self.total_top_k == self.n_experts_per_token + self.n_shared_experts
assert self.d_hidden_pad % 256 == 0
assert self.d_expert_pad % 256 == 0
# ══════════════════════════════════════════════════════════════════════════════
# Public API
# ══════════════════════════════════════════════════════════════════════════════
def fused_moe_mxfp4(
hidden_states: Tensor,
gate_up_weight: Tensor,
down_weight: Tensor,
gate_up_weight_scale: Tensor,
down_weight_scale: Tensor,
gate_up_weight_shuffled: Tensor,
down_weight_shuffled: Tensor,
gate_up_weight_scale_shuffled: Tensor,
down_weight_scale_shuffled: Tensor,
topk_weights: Tensor,
topk_ids: Tensor,
config: Dict[str, Any],
output: Optional[Tensor] = None,
) -> Tensor:
"""
DeepSeek-R1 style MXFP4 MoE forward pass.
Matches the challenge input tuple exactly.
Returns output [M, d_hidden] bf16.
"""
cfg = _Cfg(config)
M = hidden_states.shape[0]
assert hidden_states.dtype == torch.bfloat16
assert hidden_states.shape == (M, cfg.d_hidden)
assert topk_weights.shape == (M, cfg.total_top_k)
assert topk_ids.shape == (M, cfg.total_top_k)
assert gate_up_weight_scale_shuffled.data_ptr() % 128 == 0, \
"gate_up_weight_scale_shuffled must be 128B-aligned"
assert down_weight_scale_shuffled.data_ptr() % 128 == 0, \
"down_weight_scale_shuffled must be 128B-aligned"
if output is None:
output = torch.zeros(M, cfg.d_hidden, dtype=torch.bfloat16,
device=hidden_states.device)
else:
output.zero_()
try:
ext = _load_ext()
return _hip_forward(ext, hidden_states,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, output, cfg, M)
except Exception as e:
# Fallback to pure-PyTorch reference if HIP ext fails to load/run
if os.environ.get("MOE_VERBOSE", "0") == "1":
print(f"[moe_mxfp4] HIP ext unavailable ({e}), using reference path")
return _reference_forward(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
topk_weights, topk_ids, output, cfg)
def custom_kernel(data: input_t) -> output_t:
config = data[11]
hidden_pad = int(config["d_hidden_pad"] - config["d_hidden"])
intermediate_pad = int(config["d_expert_pad"] - config["d_expert"])
return fused_moe(
data[0],
data[5],
data[6],
data[9],
data[10],
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=data[7],
w2_scale=data[8],
a1_scale=None,
a2_scale=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
def _hip_forward(ext, hidden, gu_shuf, dw_shuf, gus_shuf, ds_shuf,
topk_weights, topk_ids, output, cfg: _Cfg, M: int) -> Tensor:
# 1. Routing
sorted_ids, expert_start, expert_cnt = ext.routing_dispatch(
topk_ids, cfg.n_total_experts, M
)
# 2. Actual active token count (scalar D2H — cheap)
n_active = int(expert_cnt.sum().item())
# 3. Stage 1: a4w4 gate+up GEMM + SwiGLU → intermediate [n_active, d_ep]
intermediate = ext.gemm_stage1(
hidden, gu_shuf, gus_shuf,
sorted_ids, expert_start, expert_cnt,
n_active, cfg.n_total_experts,
cfg.d_hidden, cfg.d_hidden_pad, cfg.d_expert_pad,
)
# 4. Stage 2: a4w4 down GEMM + weighted scatter-add → output
ext.gemm_stage2(
intermediate, dw_shuf, ds_shuf,
sorted_ids, expert_start, expert_cnt,
topk_weights, topk_ids, output,
cfg.n_total_experts, cfg.n_routed_experts,
cfg.d_hidden, cfg.d_hidden_pad, cfg.d_expert_pad,
)
return output
# ══════════════════════════════════════════════════════════════════════════════
# Pure-PyTorch reference (correctness fallback, not for benchmarking)
# ══════════════════════════════════════════════════════════════════════════════
def _dequant(w_raw: Tensor, scale: Tensor, rows: int, cols: int) -> Tensor:
b = w_raw.reshape(-1).view(torch.uint8)
lo = (b & 0x0F).to(torch.int8)
hi = ((b >> 4) & 0x0F).to(torch.int8)
nib = torch.stack([lo, hi], -1).reshape(-1)
sign = ((nib >> 3) & 1).float()
exp = ((nib >> 1) & 3).float()
mant = (nib & 1).float()
val = torch.where(exp > 0,
(1.0 + mant * 0.5) * (2.0 ** (exp - 1)),
mant * 0.5) * (1 - 2 * sign)
sf = (2.0 ** (scale.reshape(-1).float() - 127.0))
val = val.reshape(-1, 32) * sf.unsqueeze(-1)
return val.reshape(rows, cols).float()
def _reference_forward(hidden, gu_w, dw_w, gu_s, dw_s,
topk_weights, topk_ids, output, cfg: _Cfg) -> Tensor:
M = hidden.shape[0]
x = hidden.float()
d_ep = cfg.d_expert_pad
d_hp = cfg.d_hidden_pad
for i in range(M):
for k in range(cfg.total_top_k):
eid = int(topk_ids[i, k].item())
is_shared = eid >= cfg.n_routed_experts
weight = 1.0 if is_shared else float(topk_weights[i, k].item())
W_gu = _dequant(gu_w[eid], gu_s[eid], 2*d_ep, d_hp)
W_gate = W_gu[:d_ep, :cfg.d_hidden]
W_up = W_gu[d_ep:, :cfg.d_hidden]
xi = x[i]
gate = xi @ W_gate.T
up = xi @ W_up.T
inter = gate.sigmoid() * gate * up # SwiGLU
W_down = _dequant(dw_w[eid], dw_s[eid], d_hp, d_ep)
W_down = W_down[:cfg.d_hidden, :cfg.d_expert]
eout = inter[:cfg.d_expert] @ W_down.T
output[i] = (output[i].float() + weight * eout).to(torch.bfloat16)
return output
# ══════════════════════════════════════════════════════════════════════════════
# Quick self-test
# ══════════════════════════════════════════════════════════════════════════════
if __name__ == "__main__":
import math
def _make(bs, E, d_h, d_e, top_k, device="cuda"):
pad = lambda x: ((x+255)//256)*256
d_hp, d_ep = pad(d_h), pad(d_e)
n_shared, n_routed = 1, E-1
n_ept = top_k - n_shared
cfg = dict(d_hidden=d_h, d_expert=d_e, d_hidden_pad=d_hp,
d_expert_pad=d_ep, n_routed_experts=n_routed,
n_shared_experts=n_shared, n_experts_per_token=n_ept,
total_top_k=top_k, bs=bs)
dev = torch.device(device)
h = torch.randn(bs, d_h, dtype=torch.bfloat16, device=dev)
gu = torch.randint(0,256,(E,2*d_ep,d_hp//2), dtype=torch.uint8,device=dev)
dw = torch.randint(0,256,(E,d_hp,d_ep//2), dtype=torch.uint8,device=dev)
gus = torch.full((E,2*d_ep,d_hp//32),127, dtype=torch.uint8,device=dev)
ds = torch.full((E,d_hp,d_ep//32), 127, dtype=torch.uint8,device=dev)
gush = gu.clone().reshape(-1)
dsh = dw.clone().reshape(-1)
gss = gus.clone().reshape(-1)
dss = ds.clone().reshape(-1)
ti_r = torch.stack([torch.randperm(n_routed,device=dev)[:n_ept]
for _ in range(bs)]).int()
ti_s = torch.arange(n_routed,n_routed+n_shared,device=dev
).unsqueeze(0).expand(bs,-1).int()
ti = torch.cat([ti_r, ti_s], 1)
tw = torch.cat([F.softmax(torch.randn(bs,n_ept,device=dev),-1),
torch.ones(bs,n_shared,device=dev)], 1)
return (h, gu, dw, gus, ds, gush, dsh, gss, dss, tw, ti, cfg)
shapes = [
(16, 257, 7168, 256, 9),
(128, 257, 7168, 256, 9),
(512, 257, 7168, 256, 9),
(16, 33, 7168, 512, 9),
(128, 33, 7168, 512, 9),
(512, 33, 7168, 512, 9),
(512, 33, 7168, 2048, 9),
]
aiter_ref = [152.7, 239.0, 336.5, 106.2, 141.1, 225.0, 380.4]
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"\n{'bs':>5} {'E':>5} {'d_h':>6} {'d_e':>6} {'mean_µs':>10} {'AITER':>8} {'speedup':>8}")
print("-"*55)
times, geomeans_a = [], []
for (bs, E, d_h, d_e, tk), aref in zip(shapes, aiter_ref):
args = _make(bs, E, d_h, d_e, tk, device)
cfg = args[-1]
# Warmup
for _ in range(5): fused_moe_mxfp4(*args[:-1], cfg)
if device == "cuda": torch.cuda.synchronize()
# Time
ts = []
for _ in range(50):
s = torch.cuda.Event(enable_timing=True)
e = torch.cuda.Event(enable_timing=True)
s.record(); fused_moe_mxfp4(*args[:-1], cfg); e.record()
torch.cuda.synchronize()
ts.append(s.elapsed_time(e) * 1000)
mean_us = sum(ts)/len(ts)
speedup = aref / mean_us
times.append(mean_us)
geomeans_a.append(aref)
print(f"{bs:>5} {E:>5} {d_h:>6} {d_e:>6} {mean_us:>10.1f} {aref:>8.1f} {speedup:>7.2f}x")
gm_ours = math.exp(sum(math.log(t) for t in times) / len(times))
gm_aiter = math.exp(sum(math.log(t) for t in geomeans_a) / len(geomeans_a))
print(f"\nGeomean: {gm_ours:.1f}µs AITER: {gm_aiter:.1f}µs Speedup: {gm_aiter/gm_ours:.2f}x")
scrolls · 714 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