submission 727750
lgc0338 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 214 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-727750?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:d5bfabeb8b0d36e7617c64f632527b009aa47d25f2e022dc7f5d008df1afc181
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
const uint16_t* __restrict__ A, uint8_t* __restrict__ fp4,Kernel source
submission.py214 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;}
__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;
for(int i=threadIdx.x;i<E;i+=blockDim.x) ec[i]=0;
__syncthreads();
for(int i=threadIdx.x;i<M*topk;i+=blockDim.x) atomicAdd(&ec[ti[i]],1);
__syncthreads();
if(threadIdx.x==0){
int off=0;
for(int e=0;e<E;e++){eo[e]=off;int p=((ec[e]+bm-1)/bm)*bm;off+=p;}
eo[E]=off; nv[0]=M*topk; nv[1]=0;
int blk=0;
for(int e=0;e<E;e++){int p=((ec[e]+bm-1)/bm)*bm;for(int b=0;b<p/bm;b++)sei[blk++]=e;}
}
__syncthreads();
int tp=eo[E];
for(int i=threadIdx.x;i<tp;i+=blockDim.x){si[i]=M*topk;sw[i]=0.0f;}
__syncthreads();
for(int i=threadIdx.x;i<E;i+=blockDim.x)ec[i]=0;
__syncthreads();
for(int i=threadIdx.x;i<M*topk;i+=blockDim.x){
int eid=ti[i];int slot=atomicAdd(&ec[eid],1);
int pos=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;
for(int i=0;i<32;i++){vals[i]=u2f((uint32_t)A[base+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=256; 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_v1', 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:
# Fused scale sort for ALL quant (A and intermediate)
_hip.launch_quant(input, c['fp4'], c['ss'], _last_t2s, M, K, K32)
fp4 = torch.tensor([], dtype=torch.float4_e2m1fn_x2, device=dev).set_(
c['fp4'].untyped_storage(), c['fp4'].storage_offset(), (M, K//2), (K//2, 1))
ss = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(
c['ss'].untyped_storage(), c['ss'].storage_offset(), (max_pad, K32), (K32, 1))
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 = torch.tensor([], dtype=torch.float4_e2m1fn_x2, device=dev).set_(
c['fp4'].untyped_storage(), c['fp4'].storage_offset(), (M, K//2), (K//2, 1))
scale = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(
c['stmp'].untyped_storage(), c['stmp'].storage_offset(), (M, K32), (K32, 1))
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
@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"]
os.environ.pop('AITER_KSPLIT', None)
os.environ.pop('AITER_BYPASS_TUNE_CONFIG', None)
if M <= 128:
os.environ['AITER_KSPLIT'] = '2'
if E > 64: os.environ['AITER_BYPASS_TUNE_CONFIG'] = '1'
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 · 214 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 709268.
#!POPCORN leaderboard amd-moe-mxfp4#!POPCORN gpu MI355X-- """- Combined optimization: HIP sort + HIP A quant, both monkey-patched into fused_moe.- - HIP sort: ~5µs vs AITER ~40µs → save ~35µs- - HIP A quant: ~5µs vs Triton ~30µs → save ~25µs (topk==1 only)- - Total savings: ~60µs for M=512 shapes- """+ """Opt2: Fused scale sort for ALL quant (A+intermediate) (topk==1) via t2s mapping."""import os, sysos.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')os.environ.setdefault('CXX', 'clang++')-import torchfrom torch.utils.cpp_extension import load_inlinefrom task import input_t, output_t⋯ 3 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;}-- // ============ Sort Kernel ============__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, // token_to_sorted: flat_idx → sorted_pos+ int* __restrict__ t2s,int M, int topk, int E, int bm) {if(blockIdx.x!=0) return;⋯ 19 unchanged linesint pos=eo[eid]+slot;si[pos]=i;sw[pos]=tw[i];t2s[i]=pos;}}-- // ============ A Quant Kernel (2D grid, optional fused scale sort) ============__global__ void quant_a_kernel(const uint16_t* __restrict__ A, uint8_t* __restrict__ fp4,- uint8_t* __restrict__ sorted_scale, // write scale directly to sorted positions- const int* __restrict__ t2s, // token_to_sorted mapping (NULL = no sort)- int M, int K, int scale_stride // scale_stride = K/32 for sorted output rows+ uint8_t* __restrict__ sorted_scale, const int* __restrict__ t2s,+ int M, int K, int scale_stride) {- int K32 = K/32;- int total = M * K32;- int idx = blockIdx.x * blockDim.x + threadIdx.x;- if(idx >= total) return;- int row = idx / K32, grp = idx % K32;- int base=row*K+grp*32;+ 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;for(int i=0;i<32;i++){vals[i]=u2f((uint32_t)A[base+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);- // Write scale to sorted position (fused sort) or unsortedint srow = t2s ? t2s[row] : row;sorted_scale[srow * scale_stride + grp] = e8m0;float qs=u2f((uint32_t)(si2+127)<<23);⋯ 4 unchanged linespk[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);}- int off=row*(K/2)+grp*16;- const uint8_t*p=(const uint8_t*)pk;- for(int i=0;i<16;i++) fp4[off+i]=p[i];+ // 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){⋯ 3 unchanged lines}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 threads = 256;- int blocks = (total + threads - 1) / threads;- quant_a_kernel<<<blocks, threads>>>((const uint16_t*)A.data_ptr(),+ int total=M*(K/32); int thr=256; 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);+ 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_combo_v6', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],+ _hip = load_inline(name='moe_opt_all_v1', 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 = Trueexcept:_HAS = False- # ============ Monkey-patch sorting ============- import aiter- import aiter.fused_moe as _fm-+ import aiter, aiter.fused_moe as _fm_sort_bufs = {}+ _last_t2s = Nonedef _fast_sort(topk_ids, topk_weights, num_experts, model_dim,moebuf_dtype, block_size, expert_mask, num_local_tokens,dispatch_policy, use_opus):- # Use HIP sort for ALL shapes (1024 threads handles large E efficiently)+ global _last_t2sif 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)- max_pad = M*topk + num_experts*bm - topk- max_blk = (max_pad + bm - 1) // bm+ device = topk_ids.device; M, topk = topk_ids.shape; bm = int(block_size)+ mp = M*topk + num_experts*bm - topk; mb = (mp+bm-1)//bmkey = (M, num_experts, model_dim, bm)if key not in _sort_bufs:_sort_bufs[key] = {- 'si': torch.empty(max_pad, dtype=torch.int32, device=device),- 'sw': torch.empty(max_pad, dtype=torch.float32, device=device),- 'sei': torch.empty(max_blk, dtype=torch.int32, device=device),+ '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),⋯ 2 unchanged linesb = _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 = Nonereturn 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- max_pad = M*topk + E*bm - topk- max_blk = (max_pad + bm - 1) // bm-+ 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)//bmkey = (M, E, model_dim, bm)if key not in _sort_bufs:_sort_bufs[key] = {- 'si': torch.empty(max_pad, dtype=torch.int32, device=device),- 'sw': torch.empty(max_pad, dtype=torch.float32, device=device),- 'sei': torch.empty(max_blk, dtype=torch.int32, device=device),+ '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),⋯ 1 unchanged lines'buf': torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),}b = _sort_bufs[key]- global _last_t2s_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- # ============ Monkey-patch A quant ============from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort as _orig_quantfrom aiter.utility import fp4_utils as _fp4u-_qcache = {}- _last_t2s = None # set by _fast_sortdef _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+ M, K = input.shape; K32 = K//32; dev = input.devicemax_pad = sorted_ids.shape[0]-- key = (M, K, max_pad, topk)+ key = (M, K, max_pad)if key not in _qcache:_qcache[key] = {'fp4': torch.empty(M, K//2, dtype=torch.uint8, device=dev),- 'sorted_scale': torch.full((max_pad, K32), 127, dtype=torch.uint8, device=dev),- 'scale_tmp': torch.empty(M, K32, 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]- fp4_buf = c['fp4']-- if topk == 1 and _last_t2s is not None:- # FUSED scale sort: quant kernel writes scale directly to sorted positions- sorted_scale_buf = c['sorted_scale'] # pre-filled with 127 at creation, padding stays valid- _hip.launch_quant(input, fp4_buf, sorted_scale_buf, _last_t2s, M, K, K32)-+ if _last_t2s is not None:+ # Fused scale sort for ALL quant (A and intermediate)+ _hip.launch_quant(input, c['fp4'], c['ss'], _last_t2s, M, K, K32)fp4 = torch.tensor([], dtype=torch.float4_e2m1fn_x2, device=dev).set_(- fp4_buf.untyped_storage(), fp4_buf.storage_offset(), (M, K//2), (K//2, 1))- sorted_scale = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(- sorted_scale_buf.untyped_storage(), sorted_scale_buf.storage_offset(),- (max_pad, K32), (K32, 1))- return fp4, sorted_scale+ c['fp4'].untyped_storage(), c['fp4'].storage_offset(), (M, K//2), (K//2, 1))+ ss = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(+ c['ss'].untyped_storage(), c['ss'].storage_offset(), (max_pad, K32), (K32, 1))+ return fp4, sselse:- # Intermediate quant (topk>1): use unsorted + moe_mxfp4_sort- empty_t2s = torch.empty(0, dtype=torch.int32, device=dev)- scale_buf = c['scale_tmp']- _hip.launch_quant(input, fp4_buf, scale_buf, empty_t2s, M, K, K32)-+ et = torch.empty(0, dtype=torch.int32, device=dev)+ _hip.launch_quant(input, c['fp4'], c['stmp'], et, M, K, K32)fp4 = torch.tensor([], dtype=torch.float4_e2m1fn_x2, device=dev).set_(- fp4_buf.untyped_storage(), fp4_buf.storage_offset(), (M, K//2), (K//2, 1))+ c['fp4'].untyped_storage(), c['fp4'].storage_offset(), (M, K//2), (K//2, 1))scale = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(- scale_buf.untyped_storage(), scale_buf.storage_offset(), (M, K32), (K32, 1))-+ c['stmp'].untyped_storage(), c['stmp'].storage_offset(), (M, K32), (K32, 1))scale_sorted = _fp4u.moe_mxfp4_sort(- scale.view(token_num, topk, -1),+ 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⋯ 2 unchanged lines_qmod.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant_fm.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant- # ============ Main ============from aiter.fused_moe import fused_moe- _ACT = ActivationType.Silu- _QT = QuantType.per_1x32+ _ACT = ActivationType.Silu; _QT = QuantType.per_1x32@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"]-+ 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"]os.environ.pop('AITER_KSPLIT', None)os.environ.pop('AITER_BYPASS_TUNE_CONFIG', None)-if M <= 128:os.environ['AITER_KSPLIT'] = '2'- if E > 64:- os.environ['AITER_BYPASS_TUNE_CONFIG'] = '1'-+ if E > 64: os.environ['AITER_BYPASS_TUNE_CONFIG'] = '1'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,
scrolls · 292 diff lines total
Best evidence level for this revision: reported
JSON