Skip to content
KernelIndex
Search⌘K

submission 671802

lgc0338 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:18902c28f1cac164a5603bb34ac05c5d267c8f0bc4ad8b6e07bca383f5d53038
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15

Kernel source

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

import os
import torch
from task import input_t, output_t

# Per-shape env var tuning
def _set_env(M, E, dep):
    """Set AITER env vars for optimal per-shape performance."""
    # E=33 bs=16: SplitK helps (-36% from 96→62 μs in earlier tests)
    if E <= 33 and M <= 16:
        os.environ['AITER_KSPLIT'] = '2'
    else:
        os.environ.pop('AITER_KSPLIT', None)

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

    from aiter import ActivationType, QuantType
    from aiter.fused_moe import fused_moe

    M = hs.shape[0]
    E = cfg["n_routed_experts"] + cfg["n_shared_experts"]
    dhp = cfg["d_hidden_pad"]
    dh = cfg["d_hidden"]
    dep = cfg["d_expert_pad"]
    de = cfg["d_expert"]

    _set_env(M, E, dep)

    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-de)
scrolls · 41 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 669381.

⋯ 1 unchanged lines
#!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)
- # =============================================================================
+ # Per-shape env var tuning
+ def _set_env(M, E, dep):
+ """Set AITER env vars for optimal per-shape performance."""
+ # E=33 bs=16: SplitK helps (-36% from 96→62 μs in earlier tests)
+ if E <= 33 and M <= 16:
+ os.environ['AITER_KSPLIT'] = '2'
+ else:
+ os.environ.pop('AITER_KSPLIT', None)
- 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
+ from aiter import ActivationType, QuantType
+ from aiter.fused_moe import fused_moe
+
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
+ dep = cfg["d_expert_pad"]
+ de = cfg["d_expert"]
- # 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"])
+ _set_env(M, E, dep)
- 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']
+ 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-de)
scrolls · 522 diff lines total

Best evidence level for this revision: reported

JSON