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
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 MI355Ximport os- os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')- os.environ.setdefault('CXX', 'clang++')-import torch- from torch.utils.cpp_extension import load_inlinefrom 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