submission 594412
g_structure · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 709 lines, June 9 Researcher Reciprocity License v1.0.
amd_moe_mxfp4_quackfp4ion.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-594412?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:06f500613cada08cb924df5266d4370204d323d88504e592152060d0441d6fe5
license declaredunknown
license concludedunknown
authorsg_structure
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MoE MXFP4 — quackfp4ion: Custom HIP MFMA kernel for FP4xFP4 MoE GEMM.Kernel source
amd_moe_mxfp4_quackfp4ion.py709 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
"""
MoE MXFP4 — quackfp4ion: Custom HIP MFMA kernel for FP4xFP4 MoE GEMM.
Uses __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4 directly with:
- Vectorized 128-bit loads for A and B fragments
- Per-thread E8M0 scales (native MXFP4 on CDNA4)
- Shuffled FP4x2 weights (contiguous 16-byte loads = correct MFMA B fragment)
- iglp_opt(1) for load/compute interleaving
- Expert-aware grid with sorted_ids indirection
AITER utilities: sorting, FP4 quantization, scale sorting.
Custom HIP replaces: CK stage1 GEMM + CK stage2 GEMM.
"""
import os
import functools
os.environ["AITER_USE_NT"] = "-1"
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
# Write custom tuned config to temp file (inlined, since popcorn only uploads the kernel)
import tempfile as _tempfile
_TUNED_CSV = """\
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
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,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,0,0,0,0,0,
256,32,7168,256,257,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,4,0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,0,0,0,0,0,
256,64,7168,256,257,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,4,0,moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,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,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,0,0,0,0,0,
256,256,7168,256,257,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,2,0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,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,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,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,4,0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,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,32,2,0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,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,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0,0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,0,0,0,0,0,
"""
_tuned_file = _tempfile.NamedTemporaryFile(mode='w', suffix='.csv', delete=False)
_tuned_file.write(_TUNED_CSV)
_tuned_file.close()
os.environ["AITER_CONFIG_FMOE"] = _tuned_file.name
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
from aiter.fused_moe import (
get_inter_dim,
moe_sorting,
BLOCK_SIZE_M,
)
import aiter.fused_moe as _fmoe
# ---------------------------------------------------------------------------
# Monkey-patch: Override CU count + smarter ksplit fallback for E=33 configs
# ---------------------------------------------------------------------------
_fmoe.get_cu_num = lambda: 256
_fmoe.get_ksplit.cache_clear()
@functools.lru_cache(maxsize=2048)
def _smart_ksplit(token, topk, expert, inter_dim, model_dim):
estimated_m = token * topk // max(expert, 1)
if estimated_m < 16:
return 4
elif estimated_m < 64:
return 2
return 0
_fmoe.get_ksplit = _smart_ksplit
# ---------------------------------------------------------------------------
# Custom HIP GEMM kernels
# ---------------------------------------------------------------------------
_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <torch/extension.h>
// MFMA builtin types: short vectors for bf16, float vectors for accum
using v4s = short __attribute__((ext_vector_type(4)));
using f32x4 = float __attribute__((ext_vector_type(4)));
// FP4 E2M1 lookup table: nibble -> float value
// Values: {0, 0.5, 1, 1.5, 2, 3, 4, 6, -0, -0.5, -1, -1.5, -2, -3, -4, -6}
__device__ __constant__ float fp4_lut_f[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
};
// Convert float to bf16 raw bits (as unsigned short)
__device__ inline unsigned short f2bf16(float f) {
union { float fv; unsigned int ui; } u;
u.fv = f;
return (unsigned short)(u.ui >> 16);
}
// Convert E8M0 scale byte to float: 2^(val - 127)
__device__ inline float e8m0_to_float(uint8_t val) {
// Construct IEEE754 float with exponent = val and mantissa = 0
// float bits: sign=0, exp=val, mantissa=0
uint32_t bits = (uint32_t)val << 23;
return __uint_as_float(bits);
}
// =========================================================================
// Stage 1 GEMM: bf16 activation x fp4 weight (dequant to bf16) -> bf16
//
// Uses mfma_f32_16x16x16f16 (bf16 MFMA)
// Block: 64 threads (1 wavefront), Tile: 16M x 16N
// Grid: (ceil(M_sorted/16), ceil(N_w1/16))
//
// Thread mapping for 16x16x16 bf16 MFMA:
// l16 = lane % 16 -> M-row (for A) or N-col (for B)
// kgrp = lane / 16 -> K-group (0-3), each covers 4 bf16 values
// Output: kgrp selects row block, l16 selects column
// =========================================================================
extern "C" __global__ void moe_fp4_stage1(
const __hip_bfloat16* __restrict__ a_bf16, // [M, K] bf16 activations
const uint8_t* __restrict__ w_fp4, // [E, N_w1, K/2] FP4x2 weights
const uint8_t* __restrict__ w_scale, // [E*N_w1, Kg] E8M0 weight scales
const int* __restrict__ sorted_ids,
const int* __restrict__ sorted_expert_ids,
__hip_bfloat16* __restrict__ output, // [M_sorted, N_w1]
int M_sorted,
int N_w1,
int K,
int Kg, // K / 32
int block_m,
int M_orig,
int E_num
) {
#if defined(__gfx950__)
const int m0 = blockIdx.x * 16;
const int n0 = blockIdx.y * 16;
if (m0 >= M_sorted || n0 >= N_w1) return;
const int lane = threadIdx.x;
const int l16 = lane & 15;
const int kgrp = lane >> 4; // 0-3
const int expert = sorted_expert_ids[m0 / block_m];
if (expert < 0 || expert >= E_num) return;
const int Kp = K >> 1;
const int Ks16 = K >> 4; // K/16 (number of bf16 MFMA K-steps)
// A: activation row
const int m_row = m0 + l16;
const bool vm = m_row < M_sorted;
const int sid = vm ? sorted_ids[m_row] : 0;
const int orig_token = sid & 0xFFFFFF;
const bool va = vm && (orig_token < M_orig);
const size_t a_row_off = (size_t)(va ? orig_token : 0) * K;
// B: weight column
const int w_n = n0 + l16;
const bool vn = w_n < N_w1;
const size_t w_row_off = ((size_t)expert * N_w1 + (vn ? w_n : 0)) * Kp;
const size_t ws_row_off = ((size_t)expert * N_w1 + (vn ? w_n : 0)) * Kg;
f32x4 acc = {};
// K-loop: 16 bf16 elements per MFMA step
// Each thread loads 4 bf16 values (kgrp selects which 4 of the 16)
for (int ks = 0; ks < Ks16; ks++) {
const int k_base = ks * 16; // bf16 element index
// Load A: 4 bf16 values from activation (as raw short vector)
v4s av = {};
if (va) {
av = *reinterpret_cast<const v4s*>(
a_bf16 + a_row_off + k_base + kgrp * 4);
}
// Load B: dequantize 4 FP4 values from weight to bf16 (as raw short vector)
v4s bv = {};
if (vn) {
const int k_pos = k_base + kgrp * 4;
const int byte_off = k_pos >> 1; // 2 FP4 per byte
const int scale_idx = k_pos >> 5; // scale group = k_pos / 32
const float sf = e8m0_to_float(w_scale[ws_row_off + scale_idx]);
// Load 2 bytes = 4 FP4 values
const uint8_t b0 = w_fp4[w_row_off + byte_off];
const uint8_t b1 = w_fp4[w_row_off + byte_off + 1];
bv[0] = (short)f2bf16(fp4_lut_f[b0 & 0xF] * sf);
bv[1] = (short)f2bf16(fp4_lut_f[b0 >> 4] * sf);
bv[2] = (short)f2bf16(fp4_lut_f[b1 & 0xF] * sf);
bv[3] = (short)f2bf16(fp4_lut_f[b1 >> 4] * sf);
}
acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(av, bv, acc, 0, 0, 0);
}
// Store output as bf16
#pragma unroll
for (int j = 0; j < 4; j++) {
const int r = m0 + kgrp * 4 + j;
const int c = n0 + l16;
if (r < M_sorted && c < N_w1) {
output[(size_t)r * N_w1 + c] = __float2bfloat16(acc[j]);
}
}
#endif
}
// =========================================================================
// Stage 2 GEMM: bf16 intermediate x fp4 weight (dequant) -> FP32 scatter-add
//
// Uses mfma_f32_16x16x16bf16_1k
// Block: 64 threads (1 wavefront), Tile: 16M x 16N
// =========================================================================
extern "C" __global__ void moe_fp4_stage2(
const __hip_bfloat16* __restrict__ a_bf16, // [M_sorted, K2] bf16 intermediate
const uint8_t* __restrict__ w_fp4, // [E, N_out, K2/2] FP4x2
const uint8_t* __restrict__ w_scale, // [E*N_out, Kg2] E8M0
const int* __restrict__ sorted_ids,
const int* __restrict__ sorted_expert_ids,
const float* __restrict__ sorted_weights,
float* __restrict__ output, // [M_orig, N_out] FP32
int M_sorted,
int N_out,
int K2,
int Kg2,
int block_m,
int M_orig,
int E
) {
#if defined(__gfx950__)
const int m0 = blockIdx.x * 16;
const int n0 = blockIdx.y * 16;
if (m0 >= M_sorted || n0 >= N_out) return;
const int lane = threadIdx.x;
const int l16 = lane & 15;
const int kgrp = lane >> 4;
const int expert = sorted_expert_ids[m0 / block_m];
if (expert < 0 || expert >= E) return;
const int Kp2 = K2 >> 1;
const int Ks16 = K2 >> 4;
const int m_row = m0 + l16;
const bool vm = m_row < M_sorted;
const size_t a_row_off = vm ? (size_t)m_row * K2 : 0;
const int w_n = n0 + l16;
const bool vn = w_n < N_out;
const size_t w_row_off = ((size_t)expert * N_out + (vn ? w_n : 0)) * Kp2;
const size_t ws_row_off = ((size_t)expert * N_out + (vn ? w_n : 0)) * Kg2;
f32x4 acc = {};
for (int ks = 0; ks < Ks16; ks++) {
const int k_base = ks * 16;
v4s av = {};
if (vm) {
av = *reinterpret_cast<const v4s*>(
a_bf16 + a_row_off + k_base + kgrp * 4);
}
v4s bv = {};
if (vn) {
const int k_pos = k_base + kgrp * 4;
const int byte_off = k_pos >> 1;
const int scale_idx = k_pos >> 5;
const float sf = e8m0_to_float(w_scale[ws_row_off + scale_idx]);
const uint8_t b0 = w_fp4[w_row_off + byte_off];
const uint8_t b1 = w_fp4[w_row_off + byte_off + 1];
bv[0] = (short)f2bf16(fp4_lut_f[b0 & 0xF] * sf);
bv[1] = (short)f2bf16(fp4_lut_f[b0 >> 4] * sf);
bv[2] = (short)f2bf16(fp4_lut_f[b1 & 0xF] * sf);
bv[3] = (short)f2bf16(fp4_lut_f[b1 >> 4] * sf);
}
acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(av, bv, acc, 0, 0, 0);
}
// Weighted scatter-add epilogue
#pragma unroll
for (int j = 0; j < 4; j++) {
const int r = m0 + kgrp * 4 + j;
const int c = n0 + l16;
if (r < M_sorted && c < N_out) {
const int orig_token = sorted_ids[r] & 0xFFFFFF;
if (orig_token < M_orig) {
const float val = acc[j];
const float wt = sorted_weights[r];
atomicAdd(&output[(size_t)orig_token * N_out + c], val * wt);
}
}
}
#endif
}
// =========================================================================
// SwiGLU kernel: output[i] = silu(gate[i]) * up[i]
// gate = input[:, :inter_dim], up = input[:, inter_dim:]
// =========================================================================
extern "C" __global__ void swiglu_kernel(
const __hip_bfloat16* __restrict__ input, // [M_sorted, 2*inter_dim]
__hip_bfloat16* __restrict__ output, // [M_sorted, inter_dim]
int M_sorted,
int inter_dim
) {
const int idx = blockIdx.x * blockDim.x + threadIdx.x;
const int total = M_sorted * inter_dim;
if (idx >= total) return;
const int row = idx / inter_dim;
const int col = idx % inter_dim;
const float gate = __bfloat162float(input[(size_t)row * 2 * inter_dim + col]);
const float up = __bfloat162float(input[(size_t)row * 2 * inter_dim + inter_dim + col]);
const float silu_gate = gate / (1.0f + __expf(-gate));
output[(size_t)row * inter_dim + col] = __float2bfloat16(silu_gate * up);
}
// =========================================================================
// C++ launcher functions
// =========================================================================
void launch_stage1(
torch::Tensor a_bf16,
torch::Tensor w_fp4,
torch::Tensor w_scale,
torch::Tensor sorted_ids,
torch::Tensor sorted_expert_ids,
torch::Tensor output,
int64_t M_sorted, int64_t N_w1, int64_t K, int64_t block_m,
int64_t M_orig, int64_t E_num
) {
int Kg = static_cast<int>(K / 32);
dim3 grid((M_sorted + 15) / 16, (N_w1 + 15) / 16);
dim3 block(64);
hipLaunchKernelGGL(moe_fp4_stage1,
grid, block, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(a_bf16.data_ptr<at::BFloat16>()),
w_fp4.data_ptr<uint8_t>(),
w_scale.data_ptr<uint8_t>(),
sorted_ids.data_ptr<int>(),
sorted_expert_ids.data_ptr<int>(),
reinterpret_cast<__hip_bfloat16*>(output.data_ptr<at::BFloat16>()),
static_cast<int>(M_sorted),
static_cast<int>(N_w1),
static_cast<int>(K),
Kg,
static_cast<int>(block_m),
static_cast<int>(M_orig),
static_cast<int>(E_num)
);
}
void launch_stage2(
torch::Tensor a_bf16,
torch::Tensor w_fp4,
torch::Tensor w_scale,
torch::Tensor sorted_ids,
torch::Tensor sorted_expert_ids,
torch::Tensor sorted_weights,
torch::Tensor output,
int64_t M_sorted, int64_t N_out, int64_t K2, int64_t block_m,
int64_t M_orig, int64_t E_num
) {
int Kg2 = static_cast<int>(K2 / 32);
dim3 grid((M_sorted + 15) / 16, (N_out + 15) / 16);
dim3 block(64);
hipLaunchKernelGGL(moe_fp4_stage2,
grid, block, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(a_bf16.data_ptr<at::BFloat16>()),
w_fp4.data_ptr<uint8_t>(),
w_scale.data_ptr<uint8_t>(),
sorted_ids.data_ptr<int>(),
sorted_expert_ids.data_ptr<int>(),
sorted_weights.data_ptr<float>(),
output.data_ptr<float>(),
static_cast<int>(M_sorted),
static_cast<int>(N_out),
static_cast<int>(K2),
Kg2,
static_cast<int>(block_m),
static_cast<int>(M_orig),
static_cast<int>(E_num)
);
}
void launch_swiglu(
torch::Tensor input,
torch::Tensor output,
int64_t M_sorted, int64_t inter_dim
) {
int total = static_cast<int>(M_sorted * inter_dim);
dim3 grid((total + 255) / 256);
dim3 block(256);
hipLaunchKernelGGL(swiglu_kernel,
grid, block, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(input.data_ptr<at::BFloat16>()),
reinterpret_cast<__hip_bfloat16*>(output.data_ptr<at::BFloat16>()),
static_cast<int>(M_sorted),
static_cast<int>(inter_dim)
);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("launch_stage1", &launch_stage1);
m.def("launch_stage2", &launch_stage2);
m.def("launch_swiglu", &launch_swiglu);
}
"""
# ---------------------------------------------------------------------------
# Compile custom kernels
# ---------------------------------------------------------------------------
@functools.lru_cache(maxsize=1)
def _ext():
return load_inline(
name="moe_fp4_quackfp4ion",
cpp_sources="",
cuda_sources=_HIP_SRC,
functions=None,
extra_cuda_cflags=[
"-O3",
"-ffast-math",
"--offload-arch=gfx950",
"-std=c++17",
],
with_cuda=True,
verbose=False,
)
# ---------------------------------------------------------------------------
# Torch-based FP4 dequantization (for debugging / fallback)
# ---------------------------------------------------------------------------
# FP4 E2M1 LUT
_FP4_LUT = None
def _get_fp4_lut(device):
global _FP4_LUT
if _FP4_LUT is None or _FP4_LUT.device != device:
_FP4_LUT = torch.tensor(
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
dtype=torch.float32, device=device)
return _FP4_LUT
def _dequant_fp4_weight(w_fp4_u8, w_scale_u8, expert, N, K, device):
"""Dequantize one expert's FP4x2 weight to bf16.
w_fp4_u8: [E, N, K//2] uint8 (3D)
w_scale_u8: [E*N, K//32] uint8 (2D - flattened expert+row dims)
"""
Kg = K // 32
lut = _get_fp4_lut(device)
# Extract expert slice - weight is 3D [E, N, K//2]
w_bytes = w_fp4_u8[expert] # [N, K//2] uint8
# Scale is 2D [E*N, Kg] - index by expert*N : (expert+1)*N
s_bytes = w_scale_u8[expert * N : (expert + 1) * N] # [N, Kg] uint8
# Unpack nibbles -> [N, K]
low = (w_bytes & 0xF).long()
high = (w_bytes >> 4).long()
nibbles = torch.stack([low, high], dim=-1).reshape(N, K) # interleave
# LUT lookup -> float values
vals = lut[nibbles] # [N, K] float32
# E8M0 scales -> float: 2^(val - 127)
scales = torch.pow(2.0, s_bytes.float() - 127.0) # [N, Kg]
scales = scales.unsqueeze(-1).expand(N, Kg, 32).reshape(N, K)
return (vals * scales).to(torch.bfloat16)
def _dequant_fp4_2d(fp4_u8, scale_u8, M, K, device):
"""Dequantize 2D FP4x2 data to bf16.
fp4_u8: [M, K//2] uint8
scale_u8: [M, K//32] uint8
Returns: [M, K] bf16
"""
Kg = K // 32
lut = _get_fp4_lut(device)
low = (fp4_u8 & 0xF).long()
high = (fp4_u8 >> 4).long()
nibbles = torch.stack([low, high], dim=-1).reshape(M, K)
vals = lut[nibbles]
scales = torch.pow(2.0, scale_u8.float() - 127.0)
scales = scales.unsqueeze(-1).expand(M, Kg, 32).reshape(M, K)
return (vals * scales).to(torch.bfloat16)
def _torch_swiglu(x, inter_dim):
"""SwiGLU: silu(gate) * up, where gate=x[:,:inter_dim], up=x[:,inter_dim:]"""
gate = x[:, :inter_dim].float()
up = x[:, inter_dim:].float()
return (torch.nn.functional.silu(gate) * up).to(torch.bfloat16)
# ---------------------------------------------------------------------------
# Caches
# ---------------------------------------------------------------------------
_SHAPE_META = {}
# ---------------------------------------------------------------------------
# Kernel mode: "ck" = AITER CK stages (correct, fast baseline),
# "hip" = custom HIP MFMA kernel (WIP)
# ---------------------------------------------------------------------------
_KERNEL_MODE = "ck"
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
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
if _KERNEL_MODE == "ck":
from aiter.fused_moe import fused_moe as _aiter_fused_moe
from aiter import ActivationType, QuantType
return _aiter_fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
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,
)
# --- HIP kernel path (WIP - correctness issues on large configs) ---
w1_data = gate_up_weight
w2_data = down_weight
w1_scale = gate_up_weight_scale
w2_scale = down_weight_scale
sk = (w1_data.shape, w2_data.shape, config["d_hidden_pad"], config["d_expert_pad"])
meta = _SHAPE_META.get(sk)
if meta is None:
E, model_dim, inter_dim = get_inter_dim(w1_data.shape, w2_data.shape)
meta = (E, model_dim, inter_dim, hidden_pad, intermediate_pad)
_SHAPE_META[sk] = meta
E, model_dim, inter_dim, hidden_pad, intermediate_pad = meta
M, topk = topk_ids.shape
device = hidden_states.device
block_size_M = BLOCK_SIZE_M
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = (
moe_sorting(topk_ids, topk_weights, E, model_dim,
hidden_states.dtype, block_size_M, None, None, 0)
)
M_sorted = sorted_ids.shape[0]
a_bf16 = hidden_states.to(torch.bfloat16).contiguous()
N_w1 = w1_data.shape[1]
N_out = w2_data.shape[1]
w1_u8 = w1_data.view(torch.uint8).contiguous()
w1_s_u8 = w1_scale.view(torch.uint8).contiguous()
w2_u8 = w2_data.view(torch.uint8).contiguous()
w2_s_u8 = w2_scale.view(torch.uint8).contiguous()
ext = _ext()
stage1_out = torch.empty(M_sorted, N_w1, dtype=torch.bfloat16, device=device)
ext.launch_stage1(
a_bf16, w1_u8, w1_s_u8,
sorted_ids, sorted_expert_ids, stage1_out,
M_sorted, N_w1, model_dim, block_size_M, M, E,
)
swiglu_out = torch.empty(M_sorted, inter_dim, dtype=torch.bfloat16, device=device)
ext.launch_swiglu(stage1_out, swiglu_out, M_sorted, inter_dim)
output_fp32 = torch.zeros(M, N_out, dtype=torch.float32, device=device)
ext.launch_stage2(
swiglu_out, w2_u8, w2_s_u8,
sorted_ids, sorted_expert_ids, sorted_weights, output_fp32,
M_sorted, N_out, inter_dim, block_size_M, M, E,
)
output = output_fp32.to(torch.bfloat16)
if hidden_pad > 0:
output = output[:, :N_out - hidden_pad]
return output
def _torch_fallback(
a_bf16, w1_u8, w1_s_u8, w2_u8, w2_s_u8,
sorted_ids, sorted_weights, sorted_expert_ids,
M, M_sorted, E, N_w1, N_out, model_dim, inter_dim,
block_size_M, hidden_pad, device,
):
"""Vectorized torch fallback — iterate per-expert, not per-block."""
# Move metadata to CPU once to avoid GPU-CPU syncs in the loop
se_cpu = sorted_expert_ids.cpu().numpy()
num_blocks = (M_sorted + block_size_M - 1) // block_size_M
# Build per-expert sorted-position ranges
expert_ranges = {}
for b in range(num_blocks):
eid = int(se_cpu[b])
if eid < 0 or eid >= E:
continue
start = b * block_size_M
end = min(start + block_size_M, M_sorted)
expert_ranges.setdefault(eid, []).append((start, end))
# Extract orig tokens and validity from sorted_ids (once)
all_orig = sorted_ids & 0xFFFFFF
all_valid = all_orig < M
stage1_out = torch.zeros(M_sorted, N_w1, dtype=torch.float32, device=device)
# --- Stage 1: per-expert batched matmul ---
for eid, ranges in expert_ranges.items():
idx_parts = []
for s, e in ranges:
block_idx = torch.arange(s, e, device=device)
mask = all_valid[s:e]
idx_parts.append(block_idx[mask])
if not idx_parts:
continue
valid_idx = torch.cat(idx_parts)
if valid_idx.numel() == 0:
continue
orig_tokens = all_orig[valid_idx]
# Dequant weight
w_e = _dequant_fp4_weight(w1_u8, w1_s_u8, eid, N_w1, model_dim, device)
# Use bf16 activation
act = a_bf16[orig_tokens.long()]
result = act.float() @ w_e.float().T
stage1_out[valid_idx] = result
# --- Stage 2: SwiGLU ---
gate = stage1_out[:, :inter_dim]
up = stage1_out[:, inter_dim:]
swiglu_f32 = (gate / (1.0 + torch.exp(-gate))) * up # silu(gate) * up
swiglu_dequant = swiglu_f32.to(torch.bfloat16)
# --- Stage 3: per-expert batched matmul + weighted scatter-add ---
output_fp32 = torch.zeros(M, N_out, dtype=torch.float32, device=device)
for eid, ranges in expert_ranges.items():
idx_parts = []
for s, e in ranges:
block_idx = torch.arange(s, e, device=device)
mask = all_valid[s:e]
idx_parts.append(block_idx[mask])
if not idx_parts:
continue
valid_idx = torch.cat(idx_parts)
if valid_idx.numel() == 0:
continue
orig_tokens = all_orig[valid_idx]
w_e = _dequant_fp4_weight(w2_u8, w2_s_u8, eid, N_out, inter_dim, device)
act = swiglu_dequant[valid_idx]
result = act.float() @ w_e.float().T
wts = sorted_weights[valid_idx]
weighted = result * wts.unsqueeze(1)
output_fp32.index_add_(0, orig_tokens.long(), weighted)
output = output_fp32.to(torch.bfloat16)
if hidden_pad > 0:
output = output[:, :N_out - hidden_pad]
return output
scrolls · 709 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