Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
152.7µs
#212 of 782
2026-03-20

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.

fp4MoE 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