Skip to content
KernelIndex
Search⌘K

submission 688635

lgc0338 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2002cb88b32de30e76cd4a50f97d9179c6d962e637aa320f877e7a5aa3fa2f3e
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,

Kernel source

submission.py273 lines
#!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
"""
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 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 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;
    }
}

// ============ 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
) {
    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;
    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 unsorted
    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);
    }
    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];
}

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 threads = 256;
    int blocks = (total + threads - 1) / threads;
    quant_a_kernel<<<blocks, threads>>>((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_combo_v6', 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

# ============ Monkey-patch sorting ============
import aiter
import aiter.fused_moe as _fm

_sort_bufs = {}

def _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)
    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)
        max_pad = M*topk + num_experts*bm - topk
        max_blk = (max_pad + bm - 1) // bm
        key = (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),
                '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)
        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
    max_pad = M*topk + E*bm - topk
    max_blk = (max_pad + bm - 1) // bm

    key = (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),
            '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]
    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_quant
from aiter.utility import fp4_utils as _fp4u

_qcache = {}
_last_t2s = None  # set by _fast_sort

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, topk)
    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),
        }
    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)

        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
    else:
        # Intermediate quant (topk>1): HIP quant + HIP scale gather (replaces 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)

        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))

        scale = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(
            scale_buf.untyped_storage(), scale_buf.storage_offset(), (M, K32), (K32, 1))

        scale_sorted = _fp4u.moe_mxfp4_sort(
            scale.view(token_num, topk, -1),
            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

# ============ Main ============
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 · 273 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 686494.

⋯ 26 unchanged lines
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 M, int topk, int E, int bm
) {
if(blockIdx.x!=0) return;
⋯ 16 unchanged lines
__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];
+ int pos=eo[eid]+slot;si[pos]=i;sw[pos]=tw[i];t2s[i]=pos;
}
}
- // ============ A Quant Kernel (2D grid for any K) ============
- // Each thread handles one 32-element group. Grid: <<<ceil(M*K32/256), 256>>>
+ // ============ 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__ scale, int M, int K
+ 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
) {
int K32 = K/32;
int total = M * K32;
⋯ 6 unchanged lines
uint32_t au=(f2u(amax)+0x200000u)&0xFF800000u;
int eb=(int)((au>>23)&0xFFu);
int si2=(eb==0)?-127:max(min(eb-129,127),-127);
- scale[row*K32+grp]=(uint8_t)(si2+127);
+ uint8_t e8m0=(uint8_t)(si2+127);
+ // Write scale to sorted position (fused sort) or unsorted
+ 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++){
⋯ 9 unchanged lines
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,
- int M, int topk, int E, int bm){
+ 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(),M,topk,E,bm);
+ (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 scale, int M, int K){
+ 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(),
- (uint8_t*)fp4.data_ptr(),(uint8_t*)scale.data_ptr(),M,K);
+ (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,int,int,int,int);
- void launch_quant(torch::Tensor,torch::Tensor,torch::Tensor,int,int);
+ 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_v3', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
+ _hip = load_inline(name='moe_combo_v6', 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
⋯ 49 unchanged lines
'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]
+ global _last_t2s
_hip.launch_sort(topk_ids, topk_weights, b['si'], b['sw'], b['sei'], b['nv'],
- b['ec'], b['eo'], M, topk, E, bm)
+ 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
⋯ 3 unchanged lines
from aiter.utility import fp4_utils as _fp4u
_qcache = {}
+ _last_t2s = None # set by _fast_sort
def _fast_quant(input, sorted_ids, num_valid_ids, token_num, topk, block_size):
- # HIP for ALL quant: 2D grid ensures full wavefronts even for small K
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
- key = (M, K)
+ max_pad = sorted_ids.shape[0]
+
+ key = (M, K, max_pad, topk)
if key not in _qcache:
- _qcache[key] = (
- torch.empty(M, K//2, dtype=torch.uint8, device=dev),
- torch.empty(M, K//32, dtype=torch.uint8, device=dev),
- )
- fp4_buf, scale_buf = _qcache[key]
- _hip.launch_quant(input, fp4_buf, scale_buf, M, K)
+ _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),
+ }
+ c = _qcache[key]
+ fp4_buf = c['fp4']
- 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))
- scale = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(
- scale_buf.untyped_storage(), scale_buf.storage_offset(), (M, K//32), (K//32, 1))
+ 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)
- scale_sorted = _fp4u.moe_mxfp4_sort(
- scale, sorted_ids=sorted_ids, num_valid_ids=num_valid_ids,
- token_num=token_num, block_size=block_size)
- return fp4, scale_sorted
+ 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
+ else:
+ # Intermediate quant (topk>1): HIP quant + HIP scale gather (replaces 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)
+ 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))
+
+ scale = torch.tensor([], dtype=torch.float8_e8m0fnu, device=dev).set_(
+ scale_buf.untyped_storage(), scale_buf.storage_offset(), (M, K32), (K32, 1))
+
+ scale_sorted = _fp4u.moe_mxfp4_sort(
+ scale.view(token_num, topk, -1),
+ 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
scrolls · 169 diff lines total

Best evidence level for this revision: reported

JSON