Skip to content
KernelIndex
Search⌘K

submission 744553

lgc0338 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-744553?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
78.2µs
#5 of 782
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6fcc5bff64262284881a9f3f7a116612f9b5062c155816c528c97c4006bbb142
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15

Techniques

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

fp4const uint16_t* __restrict__ A, uint8_t* __restrict__ fp4,
shared-memory__shared__ int s_ec[512]; // expert counts (E<=257)
vector-width = uint4const uint4* Av = reinterpret_cast<const uint4*>(A + base);

Kernel source

submission.py297 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
"""Opt2: Fused scale sort for ALL quant (A+intermediate) (topk==1) via t2s mapping."""
import os, sys
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

HIP_SRC = r"""
#include <hip/hip_runtime.h>
__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;}
// Sort v4: reduced scan size (512 not 1024) + fused sentinel/scatter
__shared__ int s_ec[512];   // expert counts (E<=257)
__shared__ int s_scan[512]; // prefix sum workspace
__shared__ int s_eo[258];   // expert offsets (E+1, max 258)

__global__ void moe_sort_kernel(
    const int* __restrict__ ti, const float* __restrict__ tw,
    int* __restrict__ si, float* __restrict__ sw,
    int* __restrict__ sei, int* __restrict__ nv,
    int* __restrict__ ec, int* __restrict__ eo,
    int* __restrict__ t2s,
    int M, int topk, int E, int bm
) {
    if(blockIdx.x!=0) return;
    const int tid = threadIdx.x;
    const int NT = blockDim.x;  // 1024
    const int total = M * topk;
    // Scan size: next power of 2 >= E, capped at 512
    const int SCAN_N = (E <= 32) ? 32 : (E <= 64) ? 64 : (E <= 128) ? 128 : (E <= 256) ? 256 : 512;

    // Phase 1: Count experts in shared memory
    for(int i=tid; i<E; i+=NT) s_ec[i] = 0;
    __syncthreads();
    for(int i=tid; i<total; i+=NT) atomicAdd(&s_ec[ti[i]], 1);
    __syncthreads();

    // Phase 2: Parallel prefix sum (Blelloch on SCAN_N elements, not 1024)
    int padded = 0;
    if(tid < E) {
        padded = ((s_ec[tid] + bm - 1) / bm) * bm;
        ec[tid] = 0;  // reset for scatter
    }
    if(tid < SCAN_N) s_scan[tid] = (tid < E) ? padded : 0;
    __syncthreads();

    // Up-sweep: log2(SCAN_N) rounds (5-9 instead of 10)
    for(int stride=1; stride<SCAN_N; stride<<=1) {
        int idx = (tid+1) * (stride<<1) - 1;
        if(idx < SCAN_N) s_scan[idx] += s_scan[idx - stride];
        __syncthreads();
    }
    if(tid == 0) s_scan[SCAN_N-1] = 0;
    __syncthreads();
    // Down-sweep
    for(int stride=SCAN_N>>1; stride>=1; stride>>=1) {
        int idx = (tid+1) * (stride<<1) - 1;
        if(idx < SCAN_N) {
            int tmp = s_scan[idx - stride];
            s_scan[idx - stride] = s_scan[idx];
            s_scan[idx] += tmp;
        }
        __syncthreads();
    }

    // Store offsets
    int my_offset = 0;
    if(tid < E) {
        my_offset = s_scan[tid];
        s_eo[tid] = my_offset;
        eo[tid] = my_offset;
    }
    int tp;
    if(tid == 0) {
        int last_pad = (E > 0) ? ((s_ec[E-1] + bm - 1) / bm) * bm : 0;
        tp = s_scan[min(E-1, SCAN_N-1)] + last_pad;
        eo[E] = tp;
        nv[0] = total; nv[1] = 0;
        s_eo[E] = tp;
    }
    __syncthreads();
    tp = s_eo[E];

    // Phase 3: Parallel sei fill
    if(tid < E) {
        int blk_start = my_offset / bm;
        int num_blocks = padded / bm;
        for(int b=0; b<num_blocks; b++) sei[blk_start + b] = tid;
    }

    // Phase 4+5: Fused sentinel fill + scatter (save 1 __syncthreads)
    // Write sentinels to ALL positions first
    for(int i=tid; i<tp; i+=NT) { si[i] = total; sw[i] = 0.0f; }
    __syncthreads();
    // Then overwrite with actual tokens (sentinel values at unused positions remain)
    for(int i=tid; i<total; i+=NT) {
        int eid = ti[i];
        int slot = atomicAdd(&ec[eid], 1);
        int pos = s_eo[eid] + slot;
        si[pos] = i; sw[pos] = tw[i]; t2s[i] = pos;
    }
}
__global__ void quant_a_kernel(
    const uint16_t* __restrict__ A, uint8_t* __restrict__ fp4,
    uint8_t* __restrict__ sorted_scale, const int* __restrict__ t2s,
    int M, int K, int scale_stride
) {
    int K32=K/32, total=M*K32;
    int idx=blockIdx.x*blockDim.x+threadIdx.x;
    if(idx>=total) return;
    int row=idx/K32, grp=idx%K32, base=row*K+grp*32;
    float vals[32]; float amax=0;
    // Vectorized 128-bit loads: 4 x uint4 = 4 x 8 bf16 = 32 bf16
    const uint4* Av = reinterpret_cast<const uint4*>(A + base);
    uint4 ld[4]; ld[0]=Av[0]; ld[1]=Av[1]; ld[2]=Av[2]; ld[3]=Av[3];
    const uint16_t* pp = reinterpret_cast<const uint16_t*>(ld);
    for(int i=0;i<32;i++){vals[i]=u2f((uint32_t)pp[i]<<16);amax=fmaxf(amax,fabsf(vals[i]));}
    uint32_t au=(f2u(amax)+0x200000u)&0xFF800000u;
    int eb=(int)((au>>23)&0xFFu);
    int si2=(eb==0)?-127:max(min(eb-129,127),-127);
    uint8_t e8m0=(uint8_t)(si2+127);
    int srow = t2s ? t2s[row] : row;
    sorted_scale[srow * scale_stride + grp] = e8m0;
    float qs=u2f((uint32_t)(si2+127)<<23);
    uint32_t pk[4]={0,0,0,0};
    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);
    }
    // Vectorized 128-bit store
    *reinterpret_cast<uint4*>(fp4 + row*(K/2) + grp*16) = *reinterpret_cast<const uint4*>(pk);
}
void launch_sort(torch::Tensor ti, torch::Tensor tw, torch::Tensor si, torch::Tensor sw,
    torch::Tensor sei, torch::Tensor nv, torch::Tensor ec, torch::Tensor eo,
    torch::Tensor t2s, int M, int topk, int E, int bm){
    moe_sort_kernel<<<1,1024>>>((const int*)ti.data_ptr(),(const float*)tw.data_ptr(),
        (int*)si.data_ptr(),(float*)sw.data_ptr(),(int*)sei.data_ptr(),(int*)nv.data_ptr(),
        (int*)ec.data_ptr(),(int*)eo.data_ptr(),(int*)t2s.data_ptr(),M,topk,E,bm);
}
void launch_quant(torch::Tensor A, torch::Tensor fp4, torch::Tensor sorted_scale,
    torch::Tensor t2s, int M, int K, int scale_stride){
    int total=M*(K/32); int thr=512; int blk=(total+thr-1)/thr;
    quant_a_kernel<<<blk,thr>>>((const uint16_t*)A.data_ptr(),
        (uint8_t*)fp4.data_ptr(),(uint8_t*)sorted_scale.data_ptr(),
        t2s.numel()>0?(const int*)t2s.data_ptr():nullptr, M,K,scale_stride);
}
"""
CPP_SRC = """
void launch_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_quant(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,int,int,int);
"""
try:
    _hip = load_inline(name='moe_opt_all_v8', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
        functions=['launch_sort','launch_quant'], verbose=True,
        extra_cuda_cflags=["--offload-arch=gfx950","-std=c++20","-O3"])
    _HAS = True
except:
    _HAS = False

import aiter, aiter.fused_moe as _fm
_sort_bufs = {}
_last_t2s = None

def _fast_sort(topk_ids, topk_weights, num_experts, model_dim,
               moebuf_dtype, block_size, expert_mask, num_local_tokens,
               dispatch_policy, use_opus):
    global _last_t2s
    if not _HAS:
        fwd_fn = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
        device = topk_ids.device; M, topk = topk_ids.shape; bm = int(block_size)
        mp = M*topk + num_experts*bm - topk; mb = (mp+bm-1)//bm
        key = (M, num_experts, model_dim, bm)
        if key not in _sort_bufs:
            _sort_bufs[key] = {
                'si': torch.empty(mp, dtype=torch.int32, device=device),
                'sw': torch.empty(mp, dtype=torch.float32, device=device),
                'sei': torch.empty(mb, dtype=torch.int32, device=device),
                'nv': torch.empty(2, dtype=torch.int32, device=device),
                'ec': torch.zeros(num_experts, dtype=torch.int32, device=device),
                'eo': torch.zeros(num_experts+1, dtype=torch.int32, device=device),
                'buf': torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
            }
        b = _sort_bufs[key]
        fwd_fn(topk_ids, topk_weights, b['si'], b['sw'], b['sei'], b['nv'],
               b['buf'], num_experts, bm, expert_mask, num_local_tokens, dispatch_policy)
        _last_t2s = None
        return b['si'], b['sw'], b['sei'], b['nv'], b['buf']
    device = topk_ids.device; M, topk = topk_ids.shape
    bm = int(block_size); E = num_experts
    mp = M*topk + E*bm - topk; mb = (mp+bm-1)//bm
    key = (M, E, model_dim, bm)
    if key not in _sort_bufs:
        _sort_bufs[key] = {
            'si': torch.empty(mp, dtype=torch.int32, device=device),
            'sw': torch.empty(mp, dtype=torch.float32, device=device),
            'sei': torch.empty(mb, dtype=torch.int32, device=device),
            'nv': torch.empty(2, dtype=torch.int32, device=device),
            'ec': torch.zeros(E, dtype=torch.int32, device=device),
            'eo': torch.zeros(E+1, dtype=torch.int32, device=device),
            't2s': torch.empty(M*topk, dtype=torch.int32, device=device),
            'buf': torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
        }
    b = _sort_bufs[key]
    _hip.launch_sort(topk_ids, topk_weights, b['si'], b['sw'], b['sei'], b['nv'],
                     b['ec'], b['eo'], b['t2s'], M, topk, E, bm)
    _last_t2s = b['t2s']
    return b['si'], b['sw'], b['sei'], b['nv'], b['buf']
_fm._moe_sorting_impl = _fast_sort

from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort as _orig_quant
from aiter.utility import fp4_utils as _fp4u
_qcache = {}

def _fast_quant(input, sorted_ids, num_valid_ids, token_num, topk, block_size):
    if not _HAS or input.dtype != torch.bfloat16:
        return _orig_quant(input, sorted_ids, num_valid_ids, token_num, topk, block_size)
    M, K = input.shape; K32 = K//32; dev = input.device
    max_pad = sorted_ids.shape[0]
    key = (M, K, max_pad)
    if key not in _qcache:
        _qcache[key] = {
            'fp4': torch.empty(M, K//2, dtype=torch.uint8, device=dev),
            'ss': torch.full((max_pad, K32), 127, dtype=torch.uint8, device=dev),
            'stmp': torch.empty(M, K32, dtype=torch.uint8, device=dev),
        }
    c = _qcache[key]
    if _last_t2s is not None:
        _hip.launch_quant(input, c['fp4'], c['ss'], _last_t2s, M, K, K32)
        fp4 = _get_fp4_view(c['fp4'], (M, K//2), dev)
        ss = _get_e8m0_view(c['ss'], (max_pad, K32), dev)
        return fp4, ss
    else:
        et = torch.empty(0, dtype=torch.int32, device=dev)
        _hip.launch_quant(input, c['fp4'], c['stmp'], et, M, K, K32)
        fp4 = _get_fp4_view(c['fp4'], (M, K//2), dev)
        scale = _get_e8m0_view(c['stmp'], (M, K32), dev)
        scale_sorted = _fp4u.moe_mxfp4_sort(
            scale.view(token_num, topk, -1) if topk > 1 else scale,
            sorted_ids=sorted_ids, num_valid_ids=num_valid_ids,
            token_num=token_num, block_size=block_size)
        return fp4, scale_sorted

import aiter.ops.triton.quant.fused_mxfp4_quant as _qmod
_qmod.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant
_fm.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant

from aiter.fused_moe import fused_moe
_ACT = ActivationType.Silu; _QT = QuantType.per_1x32
_last_env = [None, None]  # [KSPLIT, BYPASS] to avoid redundant env var ops

# Cache tensor view objects to avoid per-call Python object creation
_view_cache = {}

def _get_fp4_view(buf, shape, device):
    key = ('fp4', id(buf.untyped_storage()), shape)
    if key not in _view_cache:
        _view_cache[key] = torch.tensor([], dtype=torch.float4_e2m1fn_x2, device=device).set_(
            buf.untyped_storage(), buf.storage_offset(), shape, (shape[1], 1))
    return _view_cache[key]

def _get_e8m0_view(buf, shape, device):
    key = ('e8m0', id(buf.untyped_storage()), shape)
    if key not in _view_cache:
        _view_cache[key] = torch.tensor([], dtype=torch.float8_e8m0fnu, device=device).set_(
            buf.untyped_storage(), buf.storage_offset(), shape, (shape[1], 1))
    return _view_cache[key]

@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"]
    dhp = cfg["d_hidden_pad"]; dh = cfg["d_hidden"]
    dep = cfg["d_expert_pad"]; de = cfg["d_expert"]
    # Per-shape env var dispatch
    want_ks = '7' if M <= 16 and E > 64 else None
    want_bp = '1' if M <= 16 and E > 64 else None
    if _last_env[0] != want_ks:
        if want_ks: os.environ['AITER_KSPLIT'] = want_ks
        else: os.environ.pop('AITER_KSPLIT', None)
        _last_env[0] = want_ks
    if _last_env[1] != want_bp:
        if want_bp: os.environ['AITER_BYPASS_TUNE_CONFIG'] = want_bp
        else: os.environ.pop('AITER_BYPASS_TUNE_CONFIG', None)
        _last_env[1] = want_bp
    return fused_moe(hs, w1s, w2s, tw, ti,
        expert_mask=None, activation=_ACT, quant_type=_QT,
        doweight_stage1=False, w1_scale=w1ss, w2_scale=w2ss,
        a1_scale=None, a2_scale=None,
        hidden_pad=dhp-dh, intermediate_pad=dep-de)
scrolls · 297 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 743969.

⋯ 12 unchanged lines
#include <hip/hip_runtime.h>
__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;}
- // Optimized sort: parallel prefix sum for offsets, parallel sei fill, fused sentinel+scatter
- __shared__ int s_ec[1024]; // expert counts (padded)
- __shared__ int s_eo[1024]; // expert offsets
- __shared__ int s_blk_offset[1024]; // block offset for sei
+ // Sort v4: reduced scan size (512 not 1024) + fused sentinel/scatter
+ __shared__ int s_ec[512]; // expert counts (E<=257)
+ __shared__ int s_scan[512]; // prefix sum workspace
+ __shared__ int s_eo[258]; // expert offsets (E+1, max 258)
__global__ void moe_sort_kernel(
const int* __restrict__ ti, const float* __restrict__ tw,
⋯ 7 unchanged lines
const int tid = threadIdx.x;
const int NT = blockDim.x; // 1024
const int total = M * topk;
+ // Scan size: next power of 2 >= E, capped at 512
+ const int SCAN_N = (E <= 32) ? 32 : (E <= 64) ? 64 : (E <= 128) ? 128 : (E <= 256) ? 256 : 512;
- // Phase 1: Count experts using shared memory (avoid global atomics)
+ // Phase 1: Count experts in shared memory
for(int i=tid; i<E; i+=NT) s_ec[i] = 0;
__syncthreads();
for(int i=tid; i<total; i+=NT) atomicAdd(&s_ec[ti[i]], 1);
__syncthreads();
- // Phase 2: Compute padded counts and parallel prefix sum for offsets
- // Each thread handles one expert (E <= 1024 = NT)
+ // Phase 2: Parallel prefix sum (Blelloch on SCAN_N elements, not 1024)
int padded = 0;
if(tid < E) {
- int c = s_ec[tid];
- padded = ((c + bm - 1) / bm) * bm;
- ec[tid] = 0; // reset global ec for scatter phase
+ padded = ((s_ec[tid] + bm - 1) / bm) * bm;
+ ec[tid] = 0; // reset for scatter
}
- // Warp-level inclusive prefix sum within each warp
- // Then combine across warps
- // For E<=1024 and NT=1024, we use a simple shared-memory scan
- __shared__ int s_padded[1024];
- s_padded[tid] = (tid < E) ? padded : 0;
+ if(tid < SCAN_N) s_scan[tid] = (tid < E) ? padded : 0;
__syncthreads();
- // Blelloch-style parallel prefix sum (up-sweep + down-sweep)
- // Up-sweep
- for(int stride=1; stride<NT; stride<<=1) {
+ // Up-sweep: log2(SCAN_N) rounds (5-9 instead of 10)
+ for(int stride=1; stride<SCAN_N; stride<<=1) {
int idx = (tid+1) * (stride<<1) - 1;
- if(idx < NT) s_padded[idx] += s_padded[idx - stride];
+ if(idx < SCAN_N) s_scan[idx] += s_scan[idx - stride];
__syncthreads();
}
- // Set last to 0 for exclusive scan
- if(tid == 0) { s_padded[NT-1] = 0; }
+ if(tid == 0) s_scan[SCAN_N-1] = 0;
__syncthreads();
// Down-sweep
- for(int stride=NT>>1; stride>=1; stride>>=1) {
+ for(int stride=SCAN_N>>1; stride>=1; stride>>=1) {
int idx = (tid+1) * (stride<<1) - 1;
- if(idx < NT) {
- int tmp = s_padded[idx - stride];
- s_padded[idx - stride] = s_padded[idx];
- s_padded[idx] += tmp;
+ if(idx < SCAN_N) {
+ int tmp = s_scan[idx - stride];
+ s_scan[idx - stride] = s_scan[idx];
+ s_scan[idx] += tmp;
}
__syncthreads();
}
- // Now s_padded[tid] = exclusive prefix sum = offset for expert tid
+ // Store offsets
int my_offset = 0;
if(tid < E) {
- my_offset = s_padded[tid];
+ my_offset = s_scan[tid];
s_eo[tid] = my_offset;
eo[tid] = my_offset;
}
- // Total padded count
- int tp = 0;
+ int tp;
if(tid == 0) {
- // Last expert's offset + its padded count
int last_pad = (E > 0) ? ((s_ec[E-1] + bm - 1) / bm) * bm : 0;
- tp = s_padded[E-1] + last_pad;
+ tp = s_scan[min(E-1, SCAN_N-1)] + last_pad;
eo[E] = tp;
nv[0] = total; nv[1] = 0;
- s_eo[E] = tp; // store for later use
+ s_eo[E] = tp;
}
__syncthreads();
tp = s_eo[E];
- // Phase 3: Parallel sei fill — each thread handles its expert's blocks
+ // Phase 3: Parallel sei fill
if(tid < E) {
- int num_blocks = padded / bm;
- // Need to know the block offset for this expert
- // Compute: sum of (padded_count/bm) for experts 0..tid-1
- // We can compute this from the padded prefix sum
- // block_start = sum of ceil(ec[e]/bm) for e<tid
- // But we already have offsets: my_offset / bm = block_start (since all are multiples of bm)
int blk_start = my_offset / bm;
+ int num_blocks = padded / bm;
for(int b=0; b<num_blocks; b++) sei[blk_start + b] = tid;
}
- __syncthreads();
- // Phase 4: Sentinel fill (parallel)
+ // Phase 4+5: Fused sentinel fill + scatter (save 1 __syncthreads)
+ // Write sentinels to ALL positions first
for(int i=tid; i<tp; i+=NT) { si[i] = total; sw[i] = 0.0f; }
__syncthreads();
-
- // Phase 5: Scatter tokens to sorted positions
+ // Then overwrite with actual tokens (sentinel values at unused positions remain)
for(int i=tid; i<total; i+=NT) {
int eid = ti[i];
int slot = atomicAdd(&ec[eid], 1);
scrolls · 129 diff lines total

Best evidence level for this revision: reported

JSON