Skip to content
KernelIndex
Search⌘K

submission 669381

lgc0338 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 501 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-669381?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
178.3µs
#420 of 782
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2ef3925f95ce9771d2946b1a24e17c3fc10ae932a7bd322a43ef147b3a0d446d
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memoryextern __shared__ char lds_raw[];
vector-width = uint4uint4 b_pf0 = {};

Kernel source

submission.py501 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X

import os
os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')
os.environ.setdefault('CXX', 'clang++')

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# =============================================================================
# Pure-HIP MoE v4: per-block partial buffer + gather reduce (no atomicAdd)
# =============================================================================

HIP_SRC = r"""
#include <hip/hip_runtime.h>

typedef uint8_t fp4x2_t;
typedef fp4x2_t fp4x64_t __attribute__((ext_vector_type(32)));
typedef float fp32x16_t __attribute__((ext_vector_type(16)));

__device__ __forceinline__ uint32_t f2u(float f) { uint32_t u; __builtin_memcpy(&u, &f, 4); return u; }
__device__ __forceinline__ float u2f(uint32_t u) { float f; __builtin_memcpy(&f, &u, 4); return f; }

__device__ __forceinline__ long long b_sh_addr(int expert, int n, int k_byte, int K_half, int N_dim) {
    long long eo = (long long)expert * N_dim * K_half;
    int nb = n >> 4, nl = n & 15, kb = k_byte >> 5, kg = (k_byte & 31) >> 4;
    return eo + (long long)nb * (K_half * 16) + kb * 512 + kg * 256 + nl * 16;
}

__device__ __forceinline__ uint8_t read_sh_scale(
    const uint8_t* s, int expert, int n, int ks, int sn, int N_dim
) {
    long long base = (long long)expert * N_dim * sn;
    int br = n >> 5, rh = (n & 31) >> 4, rl = n & 15;
    int bc = ks >> 3, ch = (ks & 7) >> 2, cl = ks & 3;
    return s[base + br * (sn * 32) + bc * 256 + cl * 64 + rl * 4 + ch * 2 + rh];
}

// ============================================================================
// Sorting Kernel (with inverse mapping for reduce)
// ============================================================================
__global__ void moe_sort_kernel(
    const int* __restrict__ topk_ids,
    const float* __restrict__ topk_weights,
    int* __restrict__ sorted_ids,
    float* __restrict__ sorted_weights,
    int* __restrict__ expert_block_ids,
    int* __restrict__ num_valid_out,
    int* __restrict__ expert_counts,
    int* __restrict__ expert_offsets,
    int* __restrict__ token_to_sorted,
    int M, int topk, int E, int block_m
) {
    if (blockIdx.x == 0) {
        for (int i = threadIdx.x; i < E; i += blockDim.x)
            expert_counts[i] = 0;
        __syncthreads();

        for (int i = threadIdx.x; i < M * topk; i += blockDim.x) {
            int eid = topk_ids[i];
            atomicAdd(&expert_counts[eid], 1);
        }
        __syncthreads();

        if (threadIdx.x == 0) {
            int offset = 0;
            for (int e = 0; e < E; e++) {
                expert_offsets[e] = offset;
                int padded = ((expert_counts[e] + block_m - 1) / block_m) * block_m;
                offset += padded;
            }
            expert_offsets[E] = offset;
            num_valid_out[0] = M * topk;

            int blk = 0;
            for (int e = 0; e < E; e++) {
                int padded = ((expert_counts[e] + block_m - 1) / block_m) * block_m;
                for (int b = 0; b < padded / block_m; b++)
                    expert_block_ids[blk++] = e;
            }
        }
        __syncthreads();

        int total_padded = expert_offsets[E];
        for (int i = threadIdx.x; i < total_padded; i += blockDim.x) {
            sorted_ids[i] = M * topk;
            sorted_weights[i] = 0.0f;
        }
        __syncthreads();

        for (int i = threadIdx.x; i < E; i += blockDim.x)
            expert_counts[i] = 0;
        __syncthreads();

        for (int i = threadIdx.x; i < M * topk; i += blockDim.x) {
            int eid = topk_ids[i];
            int slot = atomicAdd(&expert_counts[eid], 1);
            int pos = expert_offsets[eid] + slot;
            sorted_ids[pos] = i;
            sorted_weights[pos] = topk_weights[i];
            token_to_sorted[i] = pos;
        }
    }
}

// ============================================================================
// Fused MoE v4: partial buffer output (no atomicAdd)
// ============================================================================
#define NUM_WF 4
#define K_TILE 64
#define TILE_M 32
#define THREADS (64 * NUM_WF)

__global__ __launch_bounds__(THREADS)
void fused_moe_kernel(
    const uint16_t* __restrict__ A_bf16,
    const uint8_t* __restrict__ W1_sh,
    const uint8_t* __restrict__ W1_scale_sh,
    const uint8_t* __restrict__ W2_sh,
    const uint8_t* __restrict__ W2_scale_sh,
    float* __restrict__ partial_buf,
    const int* __restrict__ sorted_ids,
    const int* __restrict__ sorted_expert_ids,
    const int* __restrict__ num_valid_ptr,
    const float* __restrict__ sorted_weights,
    int K, int N1, int d_expert_pad, int N2, int topk, int d_hidden, int block_m
) {
    extern __shared__ char lds_raw[];
    float* red = (float*)lds_raw;
    float* s1out = red + NUM_WF * 64 * 32;

    const int block_id = blockIdx.x;
    const int warp_id = threadIdx.x / 64;
    const int tid = threadIdx.x & 63;
    const int wf_row = tid & 31;
    const int k_half = tid >> 5;
    const int expert_id = sorted_expert_ids[block_id];
    const int token_base = block_id * block_m;
    const int num_valid = *num_valid_ptr;
    const int K_half = K / 2;
    const int K2_half = d_expert_pad / 2;
    const int sn1 = K / 32;
    const int sn2 = d_expert_pad / 32;

    int orig_token = -1;
    int my_sorted = token_base + wf_row;
    // sorted_ids is sized max_blk*bm, so my_sorted is always in bounds
    int tok = sorted_ids[my_sorted];
    if (tok < num_valid) {  // sentinel M*topk rejected; valid flat indices pass
        orig_token = tok / topk;
    }

    // ============ STAGE 1: splitK across wavefronts ============
    const int K_per_wf = K / NUM_WF;
    const int k_start = warp_id * K_per_wf;
    const int k_end = k_start + K_per_wf;

    for (int n_outer = 0; n_outer < N1; n_outer += 64) {
        int my_n0 = n_outer + wf_row;
        int my_n1 = n_outer + 32 + wf_row;
        fp32x16_t c0 = {}, c1 = {};

        uint4 b_pf0 = {};
        uint8_t bs_pf0 = 127;
        if (my_n0 < N1) {
            b_pf0 = *(const uint4*)(W1_sh + b_sh_addr(expert_id, my_n0, k_start/2+k_half*16, K_half, N1));
            bs_pf0 = read_sh_scale(W1_scale_sh, expert_id, my_n0, k_start/32+k_half, sn1, N1);
        }

        for (int k = k_start; k < k_end; k += K_TILE) {
            fp4x64_t b0 = {};
            { const uint8_t* p = (const uint8_t*)&b_pf0; for(int i=0;i<16;i++) b0[i]=p[i]; }
            uint8_t sb0 = bs_pf0;

            int nk = k + K_TILE;
            if (nk < k_end && my_n0 < N1) {
                b_pf0 = *(const uint4*)(W1_sh + b_sh_addr(expert_id, my_n0, nk/2+k_half*16, K_half, N1));
                bs_pf0 = read_sh_scale(W1_scale_sh, expert_id, my_n0, nk/32+k_half, sn1, N1);
            }

            fp4x64_t a_reg = {};
            uint8_t scale_a = 127;
            if (orig_token >= 0) {
                int ak = k + k_half * 32;
                if (ak + 32 <= K) {
                    uint4 ad[4];
                    for(int i=0;i<4;i++) ad[i]=((const uint4*)(A_bf16+(long long)orig_token*K+ak))[i];
                    const uint16_t* au = (const uint16_t*)ad;
                    float amax=0, vals[32];
                    for(int i=0;i<32;i++){vals[i]=u2f((uint32_t)au[i]<<16); amax=fmaxf(amax,fabsf(vals[i]));}
                    uint32_t abu=(f2u(amax)+0x200000u)&0xFF800000u;
                    float ar=u2f(abu);
                    float sub=(ar==0.0f)?-127.0f:fminf(fmaxf(floorf(log2f(ar))-2.0f,-127.0f),127.0f);
                    scale_a=(uint8_t)((int)sub+127);
                    float qs=exp2f(sub);
                    uint32_t pk[4]={};
                    for(int j=0;j<4;j++){
                        pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8],vals[j*8+1],qs,0);
                        pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8+2],vals[j*8+3],qs,1);
                        pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8+4],vals[j*8+5],qs,2);
                        pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8+6],vals[j*8+7],qs,3);
                    }
                    const uint8_t* p=(const uint8_t*)pk;
                    for(int i=0;i<16;i++) a_reg[i]=p[i];
                }
            }

            c0=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg,b0,c0,4,4,0,(uint32_t)scale_a,0,(uint32_t)sb0);

            fp4x64_t b1={};
            uint8_t sb1=127;
            if(my_n1<N1){
                uint4 raw=*(const uint4*)(W1_sh+b_sh_addr(expert_id,my_n1,k/2+k_half*16,K_half,N1));
                const uint8_t* p=(const uint8_t*)&raw; for(int i=0;i<16;i++) b1[i]=p[i];
                sb1=read_sh_scale(W1_scale_sh,expert_id,my_n1,k/32+k_half,sn1,N1);
            }
            c1=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg,b1,c1,4,4,0,(uint32_t)scale_a,0,(uint32_t)sb1);
        }

        int lb = warp_id*64*32 + tid*32;
        for(int i=0;i<16;i++) red[lb+i]=c0[i];
        for(int i=0;i<16;i++) red[lb+16+i]=c1[i];
        __syncthreads();

        if(warp_id==0){
            for(int p=0;p<32;p++){
                float sum=red[tid*32+p];
                for(int w=1;w<NUM_WF;w++) sum+=red[w*64*32+tid*32+p];
                int nc=(p<16)?(n_outer+wf_row):(n_outer+32+wf_row);
                int mi=k_half*4+((p&15)/4)*8+((p&15)%4);
                if(mi<TILE_M && nc<N1) s1out[mi*N1+nc]=sum;
            }
        }
        __syncthreads();
    }

    // ============ SiLU + quantize intermediate ============
    // ifp4 aliases s1out in LDS — must sync between reads and writes
    // to prevent cross-wavefront race conditions
    unsigned char* ifp4=(unsigned char*)s1out;
    unsigned char* iscale=ifp4+TILE_M*(d_expert_pad/2);

    int nblk=TILE_M*(d_expert_pad/32);
    // Process in batches of THREADS, with sync between read and write phases
    for(int batch=0; batch<nblk; batch+=THREADS){
        int b=batch+threadIdx.x;
        float lv[32]; uint32_t lpk[4]={}; uint8_t le8=127;
        int lm=-1, lg=-1, lns=0;
        if(b<nblk){
            lm=b/(d_expert_pad/32); lg=b%(d_expert_pad/32); lns=lg*32;
            float amax=0;
            for(int i=0;i<32;i++){
                float gate=s1out[lm*N1+lns+i], up=s1out[lm*N1+d_expert_pad+lns+i];
                float silu=gate/(1.0f+exp2f(-1.44269504089f*gate));
                lv[i]=silu*up;
                amax=fmaxf(amax,fabsf(lv[i]));
            }
            uint32_t abu=(f2u(amax)+0x200000u)&0xFF800000u;
            float ar=u2f(abu);
            float sub=(ar==0.0f)?-127.0f:fminf(fmaxf(floorf(log2f(ar))-2.0f,-127.0f),127.0f);
            le8=(uint8_t)((int)sub+127);
            float qs=exp2f(sub);
            for(int j=0;j<4;j++){
                lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8],lv[j*8+1],qs,0);
                lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8+2],lv[j*8+3],qs,1);
                lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8+4],lv[j*8+5],qs,2);
                lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8+6],lv[j*8+7],qs,3);
            }
        }
        __syncthreads(); // all reads from s1out done
        if(b<nblk){
            iscale[lm*(d_expert_pad/32)+lg]=le8;
            const uint8_t* p=(const uint8_t*)lpk;
            for(int i=0;i<16;i++) ifp4[lm*(d_expert_pad/2)+lns/2+i]=p[i];
        }
        __syncthreads(); // all writes done before next batch
    }

    // ============ STAGE 2: write to partial_buf (no atomicAdd) ============
    for(int n_outer=0;n_outer<d_hidden;n_outer+=NUM_WF*64){
        int my_n0=n_outer+warp_id*64+wf_row;
        int my_n1=n_outer+warp_id*64+32+wf_row;
        fp32x16_t c0={},c1={};

        for(int k=0;k+K_TILE<=d_expert_pad;k+=K_TILE){
            fp4x64_t a2={};uint8_t sa2=127;
            int am=wf_row;
            if(am<TILE_M){
                int off=am*(d_expert_pad/2)+(k+k_half*32)/2;
                for(int i=0;i<16;i++) a2[i]=ifp4[off+i];
                sa2=iscale[am*(d_expert_pad/32)+(k+k_half*32)/32];
            }

            fp4x64_t b0={};uint8_t sb0=127;
            if(my_n0<N2){
                uint4 raw=*(const uint4*)(W2_sh+b_sh_addr(expert_id,my_n0,k/2+k_half*16,K2_half,N2));
                const uint8_t* p=(const uint8_t*)&raw; for(int i=0;i<16;i++) b0[i]=p[i];
                sb0=read_sh_scale(W2_scale_sh,expert_id,my_n0,k/32+k_half,sn2,N2);
            }
            c0=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a2,b0,c0,4,4,0,(uint32_t)sa2,0,(uint32_t)sb0);

            fp4x64_t b1={};uint8_t sb1=127;
            if(my_n1<N2){
                uint4 raw=*(const uint4*)(W2_sh+b_sh_addr(expert_id,my_n1,k/2+k_half*16,K2_half,N2));
                const uint8_t* p=(const uint8_t*)&raw; for(int i=0;i<16;i++) b1[i]=p[i];
                sb1=read_sh_scale(W2_scale_sh,expert_id,my_n1,k/32+k_half,sn2,N2);
            }
            c1=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a2,b1,c1,4,4,0,(uint32_t)sa2,0,(uint32_t)sb1);
        }

        // Write to partial buffer: partial_buf[sorted_pos * d_hidden + col]
        // Note: spos can exceed num_valid (valid tokens in late experts have high offsets)
        // Padding rows have A=0 → MFMA output=0 → writes 0*w=0 (harmless)
        for(int p=0;p<16;p++){
            int mi=k_half*4+(p/4)*8+(p%4);
            int spos=token_base+mi;
            if(mi<TILE_M){
                float w=sorted_weights[spos];
                long long row_off = (long long)spos * d_hidden;
                if(my_n0<d_hidden) partial_buf[row_off + my_n0] = c0[p] * w;
                if(my_n1<d_hidden) partial_buf[row_off + my_n1] = c1[p] * w;
            }
        }
    }
}

// ============================================================================
// Gather-based reduce kernel (zero atomics)
// ============================================================================
__global__ __launch_bounds__(256)
void moe_reduce_kernel(
    const float* __restrict__ partial_buf,
    uint16_t* __restrict__ output,
    const int* __restrict__ token_to_sorted,
    int M, int topk, int d_hidden
) {
    int token = blockIdx.x;
    int col = blockIdx.y * 256 + threadIdx.x;
    if (token >= M || col >= d_hidden) return;

    float sum = 0.0f;
    #pragma unroll
    for (int s = 0; s < 9; s++) {
        if (s < topk) {
            int pos = token_to_sorted[token * topk + s];
            sum += partial_buf[(long long)pos * d_hidden + col];
        }
    }

    // RNE bf16 conversion
    uint32_t bits = f2u(sum);
    bits += (0x7FFFu + ((bits >> 16) & 1u));
    output[(long long)token * d_hidden + col] = (uint16_t)(bits >> 16);
}

// C++ wrappers
void launch_moe_sort(
    torch::Tensor topk_ids, torch::Tensor topk_weights,
    torch::Tensor sorted_ids, torch::Tensor sorted_weights,
    torch::Tensor expert_block_ids, torch::Tensor num_valid,
    torch::Tensor expert_counts, torch::Tensor expert_offsets,
    torch::Tensor token_to_sorted,
    int M, int topk, int E, int block_m
) {
    moe_sort_kernel<<<1, 256>>>(
        (const int*)topk_ids.data_ptr(), (const float*)topk_weights.data_ptr(),
        (int*)sorted_ids.data_ptr(), (float*)sorted_weights.data_ptr(),
        (int*)expert_block_ids.data_ptr(), (int*)num_valid.data_ptr(),
        (int*)expert_counts.data_ptr(), (int*)expert_offsets.data_ptr(),
        (int*)token_to_sorted.data_ptr(),
        M, topk, E, block_m);
}

void launch_fused_moe(
    torch::Tensor A, torch::Tensor W1, torch::Tensor W1s,
    torch::Tensor W2, torch::Tensor W2s, torch::Tensor partial_buf,
    torch::Tensor si, torch::Tensor sei, torch::Tensor nv, torch::Tensor sw,
    int K, int N1, int dep, int N2, int topk, int dh, int bm, int nb
) {
    int lds = NUM_WF*64*32*4 + 32*N1*4;
    fused_moe_kernel<<<nb, THREADS, lds>>>(
        (const uint16_t*)A.data_ptr(),
        (const uint8_t*)W1.data_ptr(),(const uint8_t*)W1s.data_ptr(),
        (const uint8_t*)W2.data_ptr(),(const uint8_t*)W2s.data_ptr(),
        (float*)partial_buf.data_ptr(),
        (const int*)si.data_ptr(),(const int*)sei.data_ptr(),
        (const int*)nv.data_ptr(),(const float*)sw.data_ptr(),
        K,N1,dep,N2,topk,dh,bm);
}

void launch_moe_reduce(
    torch::Tensor partial_buf, torch::Tensor output,
    torch::Tensor token_to_sorted,
    int M, int topk, int d_hidden
) {
    dim3 grid(M, (d_hidden + 255) / 256);
    dim3 block(256);
    moe_reduce_kernel<<<grid, block>>>(
        (const float*)partial_buf.data_ptr(),
        (uint16_t*)output.data_ptr(),
        (const int*)token_to_sorted.data_ptr(),
        M, topk, d_hidden);
}
"""

CPP_SRC = """
void launch_moe_sort(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    int,int,int,int);
void launch_fused_moe(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    torch::Tensor,int,int,int,int,int,int,int,int);
void launch_moe_reduce(torch::Tensor,torch::Tensor,torch::Tensor,int,int,int);
"""

try:
    _hip = load_inline(
        name='moe_pure_hip_v4',
        cpp_sources=[CPP_SRC],
        cuda_sources=[HIP_SRC],
        functions=['launch_moe_sort', 'launch_fused_moe', 'launch_moe_reduce'],
        verbose=True,
        extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
    )
    _OK = True
except Exception as e:
    import sys; print(f"[ERR] {e}", file=sys.stderr)
    _OK = False


def _u8(t):
    if t.dtype == torch.uint8: return t
    return torch.tensor([],dtype=torch.uint8,device=t.device).set_(
        t.untyped_storage(),t.storage_offset(),t.shape,t.stride())

_cache = {}

@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    (hs, _,_,_,_, w1s,w2s,w1ss,w2ss, tw,ti, cfg) = data

    M = hs.shape[0]
    E = cfg["n_routed_experts"] + cfg["n_shared_experts"]
    topk = cfg["total_top_k"]
    dep = cfg["d_expert_pad"]
    dhp = cfg["d_hidden_pad"]
    dh = cfg["d_hidden"]
    N1 = 2 * dep
    bm = 32

    # LDS check: for large dep, s1out exceeds 160KB LDS
    # Stage1 needs red(32KB) + s1out(32*N1*4). Must fit in 160KB.
    lds_need = 4*64*32*4 + 32*N1*4
    if True:  # TODO: HIP kernel has correctness issues, AITER for all shapes for now
        # For dep > 512, use AITER (HIP kernel needs LDS redesign for large dep)
        from aiter import ActivationType, QuantType
        from aiter.fused_moe import fused_moe
        return fused_moe(hs,w1s,w2s,tw,ti,expert_mask=None,
            activation=ActivationType.Silu,quant_type=QuantType.per_1x32,
            doweight_stage1=False,w1_scale=w1ss,w2_scale=w2ss,
            a1_scale=None,a2_scale=None,
            hidden_pad=dhp-dh,intermediate_pad=dep-cfg["d_expert"])

    dev = hs.device
    max_pad = M*topk + E*bm - topk
    max_blk = (max_pad+bm-1)//bm

    buf_rows = max_blk * bm  # round up to cover all blocks
    key = (M, E, topk, dh, dep)
    if key not in _cache:
        _cache[key] = {
            'sei': torch.zeros(max_blk, dtype=torch.int32, device=dev),
            'nv': torch.empty(1, dtype=torch.int32, device=dev),
            'ec': torch.zeros(E, dtype=torch.int32, device=dev),
            'eo': torch.zeros(E+1, dtype=torch.int32, device=dev),
            't2s': torch.empty(M*topk, dtype=torch.int32, device=dev),
            'pbuf': torch.empty(buf_rows * dh, dtype=torch.float32, device=dev),
            'out': torch.empty(M, dh, dtype=torch.bfloat16, device=dev),
        }
    c = _cache[key]

    # Must re-init per call: routing changes each iteration, stale data corrupts results
    si = torch.full((buf_rows,), M*topk, dtype=torch.int32, device=dev)
    sw = torch.zeros(buf_rows, dtype=torch.float32, device=dev)

    _hip.launch_moe_sort(ti, tw, si, sw, c['sei'], c['nv'],
                         c['ec'], c['eo'], c['t2s'], M, topk, E, bm)

    nb = max_blk

    _hip.launch_fused_moe(
        hs, _u8(w1s), _u8(w1ss), _u8(w2s), _u8(w2ss), c['pbuf'],
        si, c['sei'], c['nv'], sw,
        dhp, N1, dep, dhp, topk, dh, bm, nb)

    _hip.launch_moe_reduce(c['pbuf'], c['out'], c['t2s'], M, topk, dh)

    return c['out']
scrolls · 501 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 654068.

#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
- """
- Python-level extreme optimization on CK baseline:
- 1. @torch.inference_mode() — disable autograd
- 2. Pre-allocated persistent buffers — avoid torch.empty per call
- 3. splitk for E=33 small batch decode
- 4. Try doweight_stage1=True
- 5. moe_sorting_dispatch_policy variations
- """
+ import os
+ os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')
+ os.environ.setdefault('CXX', 'clang++')
import torch
+ from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
- from aiter import ActivationType, QuantType
- from aiter.fused_moe import fused_moe
+ # =============================================================================
+ # Pure-HIP MoE v4: per-block partial buffer + gather reduce (no atomicAdd)
+ # =============================================================================
+ HIP_SRC = r"""
+ #include <hip/hip_runtime.h>
+
+ typedef uint8_t fp4x2_t;
+ typedef fp4x2_t fp4x64_t __attribute__((ext_vector_type(32)));
+ typedef float fp32x16_t __attribute__((ext_vector_type(16)));
+
+ __device__ __forceinline__ uint32_t f2u(float f) { uint32_t u; __builtin_memcpy(&u, &f, 4); return u; }
+ __device__ __forceinline__ float u2f(uint32_t u) { float f; __builtin_memcpy(&f, &u, 4); return f; }
+
+ __device__ __forceinline__ long long b_sh_addr(int expert, int n, int k_byte, int K_half, int N_dim) {
+ long long eo = (long long)expert * N_dim * K_half;
+ int nb = n >> 4, nl = n & 15, kb = k_byte >> 5, kg = (k_byte & 31) >> 4;
+ return eo + (long long)nb * (K_half * 16) + kb * 512 + kg * 256 + nl * 16;
+ }
+
+ __device__ __forceinline__ uint8_t read_sh_scale(
+ const uint8_t* s, int expert, int n, int ks, int sn, int N_dim
+ ) {
+ long long base = (long long)expert * N_dim * sn;
+ int br = n >> 5, rh = (n & 31) >> 4, rl = n & 15;
+ int bc = ks >> 3, ch = (ks & 7) >> 2, cl = ks & 3;
+ return s[base + br * (sn * 32) + bc * 256 + cl * 64 + rl * 4 + ch * 2 + rh];
+ }
+
+ // ============================================================================
+ // Sorting Kernel (with inverse mapping for reduce)
+ // ============================================================================
+ __global__ void moe_sort_kernel(
+ const int* __restrict__ topk_ids,
+ const float* __restrict__ topk_weights,
+ int* __restrict__ sorted_ids,
+ float* __restrict__ sorted_weights,
+ int* __restrict__ expert_block_ids,
+ int* __restrict__ num_valid_out,
+ int* __restrict__ expert_counts,
+ int* __restrict__ expert_offsets,
+ int* __restrict__ token_to_sorted,
+ int M, int topk, int E, int block_m
+ ) {
+ if (blockIdx.x == 0) {
+ for (int i = threadIdx.x; i < E; i += blockDim.x)
+ expert_counts[i] = 0;
+ __syncthreads();
+
+ for (int i = threadIdx.x; i < M * topk; i += blockDim.x) {
+ int eid = topk_ids[i];
+ atomicAdd(&expert_counts[eid], 1);
+ }
+ __syncthreads();
+
+ if (threadIdx.x == 0) {
+ int offset = 0;
+ for (int e = 0; e < E; e++) {
+ expert_offsets[e] = offset;
+ int padded = ((expert_counts[e] + block_m - 1) / block_m) * block_m;
+ offset += padded;
+ }
+ expert_offsets[E] = offset;
+ num_valid_out[0] = M * topk;
+
+ int blk = 0;
+ for (int e = 0; e < E; e++) {
+ int padded = ((expert_counts[e] + block_m - 1) / block_m) * block_m;
+ for (int b = 0; b < padded / block_m; b++)
+ expert_block_ids[blk++] = e;
+ }
+ }
+ __syncthreads();
+
+ int total_padded = expert_offsets[E];
+ for (int i = threadIdx.x; i < total_padded; i += blockDim.x) {
+ sorted_ids[i] = M * topk;
+ sorted_weights[i] = 0.0f;
+ }
+ __syncthreads();
+
+ for (int i = threadIdx.x; i < E; i += blockDim.x)
+ expert_counts[i] = 0;
+ __syncthreads();
+
+ for (int i = threadIdx.x; i < M * topk; i += blockDim.x) {
+ int eid = topk_ids[i];
+ int slot = atomicAdd(&expert_counts[eid], 1);
+ int pos = expert_offsets[eid] + slot;
+ sorted_ids[pos] = i;
+ sorted_weights[pos] = topk_weights[i];
+ token_to_sorted[i] = pos;
+ }
+ }
+ }
+
+ // ============================================================================
+ // Fused MoE v4: partial buffer output (no atomicAdd)
+ // ============================================================================
+ #define NUM_WF 4
+ #define K_TILE 64
+ #define TILE_M 32
+ #define THREADS (64 * NUM_WF)
+
+ __global__ __launch_bounds__(THREADS)
+ void fused_moe_kernel(
+ const uint16_t* __restrict__ A_bf16,
+ const uint8_t* __restrict__ W1_sh,
+ const uint8_t* __restrict__ W1_scale_sh,
+ const uint8_t* __restrict__ W2_sh,
+ const uint8_t* __restrict__ W2_scale_sh,
+ float* __restrict__ partial_buf,
+ const int* __restrict__ sorted_ids,
+ const int* __restrict__ sorted_expert_ids,
+ const int* __restrict__ num_valid_ptr,
+ const float* __restrict__ sorted_weights,
+ int K, int N1, int d_expert_pad, int N2, int topk, int d_hidden, int block_m
+ ) {
+ extern __shared__ char lds_raw[];
+ float* red = (float*)lds_raw;
+ float* s1out = red + NUM_WF * 64 * 32;
+
+ const int block_id = blockIdx.x;
+ const int warp_id = threadIdx.x / 64;
+ const int tid = threadIdx.x & 63;
+ const int wf_row = tid & 31;
+ const int k_half = tid >> 5;
+ const int expert_id = sorted_expert_ids[block_id];
+ const int token_base = block_id * block_m;
+ const int num_valid = *num_valid_ptr;
+ const int K_half = K / 2;
+ const int K2_half = d_expert_pad / 2;
+ const int sn1 = K / 32;
+ const int sn2 = d_expert_pad / 32;
+
+ int orig_token = -1;
+ int my_sorted = token_base + wf_row;
+ // sorted_ids is sized max_blk*bm, so my_sorted is always in bounds
+ int tok = sorted_ids[my_sorted];
+ if (tok < num_valid) { // sentinel M*topk rejected; valid flat indices pass
+ orig_token = tok / topk;
+ }
+
+ // ============ STAGE 1: splitK across wavefronts ============
+ const int K_per_wf = K / NUM_WF;
+ const int k_start = warp_id * K_per_wf;
+ const int k_end = k_start + K_per_wf;
+
+ for (int n_outer = 0; n_outer < N1; n_outer += 64) {
+ int my_n0 = n_outer + wf_row;
+ int my_n1 = n_outer + 32 + wf_row;
+ fp32x16_t c0 = {}, c1 = {};
+
+ uint4 b_pf0 = {};
+ uint8_t bs_pf0 = 127;
+ if (my_n0 < N1) {
+ b_pf0 = *(const uint4*)(W1_sh + b_sh_addr(expert_id, my_n0, k_start/2+k_half*16, K_half, N1));
+ bs_pf0 = read_sh_scale(W1_scale_sh, expert_id, my_n0, k_start/32+k_half, sn1, N1);
+ }
+
+ for (int k = k_start; k < k_end; k += K_TILE) {
+ fp4x64_t b0 = {};
+ { const uint8_t* p = (const uint8_t*)&b_pf0; for(int i=0;i<16;i++) b0[i]=p[i]; }
+ uint8_t sb0 = bs_pf0;
+
+ int nk = k + K_TILE;
+ if (nk < k_end && my_n0 < N1) {
+ b_pf0 = *(const uint4*)(W1_sh + b_sh_addr(expert_id, my_n0, nk/2+k_half*16, K_half, N1));
+ bs_pf0 = read_sh_scale(W1_scale_sh, expert_id, my_n0, nk/32+k_half, sn1, N1);
+ }
+
+ fp4x64_t a_reg = {};
+ uint8_t scale_a = 127;
+ if (orig_token >= 0) {
+ int ak = k + k_half * 32;
+ if (ak + 32 <= K) {
+ uint4 ad[4];
+ for(int i=0;i<4;i++) ad[i]=((const uint4*)(A_bf16+(long long)orig_token*K+ak))[i];
+ const uint16_t* au = (const uint16_t*)ad;
+ float amax=0, vals[32];
+ for(int i=0;i<32;i++){vals[i]=u2f((uint32_t)au[i]<<16); amax=fmaxf(amax,fabsf(vals[i]));}
+ uint32_t abu=(f2u(amax)+0x200000u)&0xFF800000u;
+ float ar=u2f(abu);
+ float sub=(ar==0.0f)?-127.0f:fminf(fmaxf(floorf(log2f(ar))-2.0f,-127.0f),127.0f);
+ scale_a=(uint8_t)((int)sub+127);
+ float qs=exp2f(sub);
+ uint32_t pk[4]={};
+ for(int j=0;j<4;j++){
+ pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8],vals[j*8+1],qs,0);
+ pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8+2],vals[j*8+3],qs,1);
+ pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8+4],vals[j*8+5],qs,2);
+ pk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[j],vals[j*8+6],vals[j*8+7],qs,3);
+ }
+ const uint8_t* p=(const uint8_t*)pk;
+ for(int i=0;i<16;i++) a_reg[i]=p[i];
+ }
+ }
+
+ c0=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg,b0,c0,4,4,0,(uint32_t)scale_a,0,(uint32_t)sb0);
+
+ fp4x64_t b1={};
+ uint8_t sb1=127;
+ if(my_n1<N1){
+ uint4 raw=*(const uint4*)(W1_sh+b_sh_addr(expert_id,my_n1,k/2+k_half*16,K_half,N1));
+ const uint8_t* p=(const uint8_t*)&raw; for(int i=0;i<16;i++) b1[i]=p[i];
+ sb1=read_sh_scale(W1_scale_sh,expert_id,my_n1,k/32+k_half,sn1,N1);
+ }
+ c1=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg,b1,c1,4,4,0,(uint32_t)scale_a,0,(uint32_t)sb1);
+ }
+
+ int lb = warp_id*64*32 + tid*32;
+ for(int i=0;i<16;i++) red[lb+i]=c0[i];
+ for(int i=0;i<16;i++) red[lb+16+i]=c1[i];
+ __syncthreads();
+
+ if(warp_id==0){
+ for(int p=0;p<32;p++){
+ float sum=red[tid*32+p];
+ for(int w=1;w<NUM_WF;w++) sum+=red[w*64*32+tid*32+p];
+ int nc=(p<16)?(n_outer+wf_row):(n_outer+32+wf_row);
+ int mi=k_half*4+((p&15)/4)*8+((p&15)%4);
+ if(mi<TILE_M && nc<N1) s1out[mi*N1+nc]=sum;
+ }
+ }
+ __syncthreads();
+ }
+
+ // ============ SiLU + quantize intermediate ============
+ // ifp4 aliases s1out in LDS — must sync between reads and writes
+ // to prevent cross-wavefront race conditions
+ unsigned char* ifp4=(unsigned char*)s1out;
+ unsigned char* iscale=ifp4+TILE_M*(d_expert_pad/2);
+
+ int nblk=TILE_M*(d_expert_pad/32);
+ // Process in batches of THREADS, with sync between read and write phases
+ for(int batch=0; batch<nblk; batch+=THREADS){
+ int b=batch+threadIdx.x;
+ float lv[32]; uint32_t lpk[4]={}; uint8_t le8=127;
+ int lm=-1, lg=-1, lns=0;
+ if(b<nblk){
+ lm=b/(d_expert_pad/32); lg=b%(d_expert_pad/32); lns=lg*32;
+ float amax=0;
+ for(int i=0;i<32;i++){
+ float gate=s1out[lm*N1+lns+i], up=s1out[lm*N1+d_expert_pad+lns+i];
+ float silu=gate/(1.0f+exp2f(-1.44269504089f*gate));
+ lv[i]=silu*up;
+ amax=fmaxf(amax,fabsf(lv[i]));
+ }
+ uint32_t abu=(f2u(amax)+0x200000u)&0xFF800000u;
+ float ar=u2f(abu);
+ float sub=(ar==0.0f)?-127.0f:fminf(fmaxf(floorf(log2f(ar))-2.0f,-127.0f),127.0f);
+ le8=(uint8_t)((int)sub+127);
+ float qs=exp2f(sub);
+ for(int j=0;j<4;j++){
+ lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8],lv[j*8+1],qs,0);
+ lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8+2],lv[j*8+3],qs,1);
+ lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8+4],lv[j*8+5],qs,2);
+ lpk[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(lpk[j],lv[j*8+6],lv[j*8+7],qs,3);
+ }
+ }
+ __syncthreads(); // all reads from s1out done
+ if(b<nblk){
+ iscale[lm*(d_expert_pad/32)+lg]=le8;
+ const uint8_t* p=(const uint8_t*)lpk;
+ for(int i=0;i<16;i++) ifp4[lm*(d_expert_pad/2)+lns/2+i]=p[i];
+ }
+ __syncthreads(); // all writes done before next batch
+ }
+
+ // ============ STAGE 2: write to partial_buf (no atomicAdd) ============
+ for(int n_outer=0;n_outer<d_hidden;n_outer+=NUM_WF*64){
+ int my_n0=n_outer+warp_id*64+wf_row;
+ int my_n1=n_outer+warp_id*64+32+wf_row;
+ fp32x16_t c0={},c1={};
+
+ for(int k=0;k+K_TILE<=d_expert_pad;k+=K_TILE){
+ fp4x64_t a2={};uint8_t sa2=127;
+ int am=wf_row;
+ if(am<TILE_M){
+ int off=am*(d_expert_pad/2)+(k+k_half*32)/2;
+ for(int i=0;i<16;i++) a2[i]=ifp4[off+i];
+ sa2=iscale[am*(d_expert_pad/32)+(k+k_half*32)/32];
+ }
+
+ fp4x64_t b0={};uint8_t sb0=127;
+ if(my_n0<N2){
+ uint4 raw=*(const uint4*)(W2_sh+b_sh_addr(expert_id,my_n0,k/2+k_half*16,K2_half,N2));
+ const uint8_t* p=(const uint8_t*)&raw; for(int i=0;i<16;i++) b0[i]=p[i];
+ sb0=read_sh_scale(W2_scale_sh,expert_id,my_n0,k/32+k_half,sn2,N2);
+ }
+ c0=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a2,b0,c0,4,4,0,(uint32_t)sa2,0,(uint32_t)sb0);
+
+ fp4x64_t b1={};uint8_t sb1=127;
+ if(my_n1<N2){
+ uint4 raw=*(const uint4*)(W2_sh+b_sh_addr(expert_id,my_n1,k/2+k_half*16,K2_half,N2));
+ const uint8_t* p=(const uint8_t*)&raw; for(int i=0;i<16;i++) b1[i]=p[i];
+ sb1=read_sh_scale(W2_scale_sh,expert_id,my_n1,k/32+k_half,sn2,N2);
+ }
+ c1=__builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a2,b1,c1,4,4,0,(uint32_t)sa2,0,(uint32_t)sb1);
+ }
+
+ // Write to partial buffer: partial_buf[sorted_pos * d_hidden + col]
+ // Note: spos can exceed num_valid (valid tokens in late experts have high offsets)
+ // Padding rows have A=0 → MFMA output=0 → writes 0*w=0 (harmless)
+ for(int p=0;p<16;p++){
+ int mi=k_half*4+(p/4)*8+(p%4);
+ int spos=token_base+mi;
+ if(mi<TILE_M){
+ float w=sorted_weights[spos];
+ long long row_off = (long long)spos * d_hidden;
+ if(my_n0<d_hidden) partial_buf[row_off + my_n0] = c0[p] * w;
+ if(my_n1<d_hidden) partial_buf[row_off + my_n1] = c1[p] * w;
+ }
+ }
+ }
+ }
+
+ // ============================================================================
+ // Gather-based reduce kernel (zero atomics)
+ // ============================================================================
+ __global__ __launch_bounds__(256)
+ void moe_reduce_kernel(
+ const float* __restrict__ partial_buf,
+ uint16_t* __restrict__ output,
+ const int* __restrict__ token_to_sorted,
+ int M, int topk, int d_hidden
+ ) {
+ int token = blockIdx.x;
+ int col = blockIdx.y * 256 + threadIdx.x;
+ if (token >= M || col >= d_hidden) return;
+
+ float sum = 0.0f;
+ #pragma unroll
+ for (int s = 0; s < 9; s++) {
+ if (s < topk) {
+ int pos = token_to_sorted[token * topk + s];
+ sum += partial_buf[(long long)pos * d_hidden + col];
+ }
+ }
+
+ // RNE bf16 conversion
+ uint32_t bits = f2u(sum);
+ bits += (0x7FFFu + ((bits >> 16) & 1u));
+ output[(long long)token * d_hidden + col] = (uint16_t)(bits >> 16);
+ }
+
+ // C++ wrappers
+ void launch_moe_sort(
+ torch::Tensor topk_ids, torch::Tensor topk_weights,
+ torch::Tensor sorted_ids, torch::Tensor sorted_weights,
+ torch::Tensor expert_block_ids, torch::Tensor num_valid,
+ torch::Tensor expert_counts, torch::Tensor expert_offsets,
+ torch::Tensor token_to_sorted,
+ int M, int topk, int E, int block_m
+ ) {
+ moe_sort_kernel<<<1, 256>>>(
+ (const int*)topk_ids.data_ptr(), (const float*)topk_weights.data_ptr(),
+ (int*)sorted_ids.data_ptr(), (float*)sorted_weights.data_ptr(),
+ (int*)expert_block_ids.data_ptr(), (int*)num_valid.data_ptr(),
+ (int*)expert_counts.data_ptr(), (int*)expert_offsets.data_ptr(),
+ (int*)token_to_sorted.data_ptr(),
+ M, topk, E, block_m);
+ }
+
+ void launch_fused_moe(
+ torch::Tensor A, torch::Tensor W1, torch::Tensor W1s,
+ torch::Tensor W2, torch::Tensor W2s, torch::Tensor partial_buf,
+ torch::Tensor si, torch::Tensor sei, torch::Tensor nv, torch::Tensor sw,
+ int K, int N1, int dep, int N2, int topk, int dh, int bm, int nb
+ ) {
+ int lds = NUM_WF*64*32*4 + 32*N1*4;
+ fused_moe_kernel<<<nb, THREADS, lds>>>(
+ (const uint16_t*)A.data_ptr(),
+ (const uint8_t*)W1.data_ptr(),(const uint8_t*)W1s.data_ptr(),
+ (const uint8_t*)W2.data_ptr(),(const uint8_t*)W2s.data_ptr(),
+ (float*)partial_buf.data_ptr(),
+ (const int*)si.data_ptr(),(const int*)sei.data_ptr(),
+ (const int*)nv.data_ptr(),(const float*)sw.data_ptr(),
+ K,N1,dep,N2,topk,dh,bm);
+ }
+
+ void launch_moe_reduce(
+ torch::Tensor partial_buf, torch::Tensor output,
+ torch::Tensor token_to_sorted,
+ int M, int topk, int d_hidden
+ ) {
+ dim3 grid(M, (d_hidden + 255) / 256);
+ dim3 block(256);
+ moe_reduce_kernel<<<grid, block>>>(
+ (const float*)partial_buf.data_ptr(),
+ (uint16_t*)output.data_ptr(),
+ (const int*)token_to_sorted.data_ptr(),
+ M, topk, d_hidden);
+ }
+ """
+
+ CPP_SRC = """
+ void launch_moe_sort(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
+ torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
+ int,int,int,int);
+ void launch_fused_moe(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
+ torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
+ torch::Tensor,int,int,int,int,int,int,int,int);
+ void launch_moe_reduce(torch::Tensor,torch::Tensor,torch::Tensor,int,int,int);
+ """
+
+ try:
+ _hip = load_inline(
+ name='moe_pure_hip_v4',
+ cpp_sources=[CPP_SRC],
+ cuda_sources=[HIP_SRC],
+ functions=['launch_moe_sort', 'launch_fused_moe', 'launch_moe_reduce'],
+ verbose=True,
+ extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
+ )
+ _OK = True
+ except Exception as e:
+ import sys; print(f"[ERR] {e}", file=sys.stderr)
+ _OK = False
+
+
+ def _u8(t):
+ if t.dtype == torch.uint8: return t
+ return torch.tensor([],dtype=torch.uint8,device=t.device).set_(
+ t.untyped_storage(),t.storage_offset(),t.shape,t.stride())
+
+ _cache = {}
+
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
- (
- hidden_states,
- _guw, _dw, _gus, _ds,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- gate_up_weight_scale_shuffled,
- down_weight_scale_shuffled,
- topk_weights,
- topk_ids,
- config,
- ) = data
+ (hs, _,_,_,_, w1s,w2s,w1ss,w2ss, tw,ti, cfg) = data
- hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
- intermediate_pad = config["d_expert_pad"] - config["d_expert"]
+ M = hs.shape[0]
+ E = cfg["n_routed_experts"] + cfg["n_shared_experts"]
+ topk = cfg["total_top_k"]
+ dep = cfg["d_expert_pad"]
+ dhp = cfg["d_hidden_pad"]
+ dh = cfg["d_hidden"]
+ N1 = 2 * dep
+ bm = 32
- return fused_moe(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- topk_weights,
- topk_ids,
- expert_mask=None,
- activation=ActivationType.Silu,
- quant_type=QuantType.per_1x32,
- doweight_stage1=False,
- w1_scale=gate_up_weight_scale_shuffled,
- w2_scale=down_weight_scale_shuffled,
- a1_scale=None,
- a2_scale=None,
- hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
- )
+ # LDS check: for large dep, s1out exceeds 160KB LDS
+ # Stage1 needs red(32KB) + s1out(32*N1*4). Must fit in 160KB.
+ lds_need = 4*64*32*4 + 32*N1*4
+ if True: # TODO: HIP kernel has correctness issues, AITER for all shapes for now
+ # For dep > 512, use AITER (HIP kernel needs LDS redesign for large dep)
+ from aiter import ActivationType, QuantType
+ from aiter.fused_moe import fused_moe
+ return fused_moe(hs,w1s,w2s,tw,ti,expert_mask=None,
+ activation=ActivationType.Silu,quant_type=QuantType.per_1x32,
+ doweight_stage1=False,w1_scale=w1ss,w2_scale=w2ss,
+ a1_scale=None,a2_scale=None,
+ hidden_pad=dhp-dh,intermediate_pad=dep-cfg["d_expert"])
+
+ dev = hs.device
+ max_pad = M*topk + E*bm - topk
+ max_blk = (max_pad+bm-1)//bm
+
+ buf_rows = max_blk * bm # round up to cover all blocks
+ key = (M, E, topk, dh, dep)
+ if key not in _cache:
+ _cache[key] = {
+ 'sei': torch.zeros(max_blk, dtype=torch.int32, device=dev),
+ 'nv': torch.empty(1, dtype=torch.int32, device=dev),
+ 'ec': torch.zeros(E, dtype=torch.int32, device=dev),
+ 'eo': torch.zeros(E+1, dtype=torch.int32, device=dev),
+ 't2s': torch.empty(M*topk, dtype=torch.int32, device=dev),
+ 'pbuf': torch.empty(buf_rows * dh, dtype=torch.float32, device=dev),
+ 'out': torch.empty(M, dh, dtype=torch.bfloat16, device=dev),
+ }
+ c = _cache[key]
+
+ # Must re-init per call: routing changes each iteration, stale data corrupts results
+ si = torch.full((buf_rows,), M*topk, dtype=torch.int32, device=dev)
+ sw = torch.zeros(buf_rows, dtype=torch.float32, device=dev)
+
+ _hip.launch_moe_sort(ti, tw, si, sw, c['sei'], c['nv'],
+ c['ec'], c['eo'], c['t2s'], M, topk, E, bm)
+
+ nb = max_blk
+
+ _hip.launch_fused_moe(
+ hs, _u8(w1s), _u8(w1ss), _u8(w2s), _u8(w2ss), c['pbuf'],
+ si, c['sei'], c['nv'], sw,
+ dhp, N1, dep, dhp, topk, dh, bm, nb)
+
+ _hip.launch_moe_reduce(c['pbuf'], c['out'], c['t2s'], M, topk, dh)
+
+ return c['out']
scrolls · 540 diff lines total

Best evidence level for this revision: reported

JSON