submission 685293
lgc0338 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 142 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-685293?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:1f51e00f4536c071b0b9d96e10e62f14859d63c4f1930d5d6a2fa100b8b14d62
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.py142 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
"""
Monkey-patch fused_dynamic_mxfp4_quant_moe_sort with our HIP hardware quant.
fused_moe still receives bf16 input (no KeyError), but A quant uses hardware instruction.
Only patches for M=512 (2-stage path). M≤128 uses CKTile (no quant needed).
"""
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 A quantization kernel
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 quant_a_kernel(
const uint16_t* __restrict__ A, uint8_t* __restrict__ fp4,
uint8_t* __restrict__ scale, int M, int K
) {
int row=blockIdx.x, grp=threadIdx.x;
if(row>=M||grp>=K/32) return;
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);
scale[row*(K/32)+grp]=(uint8_t)(si2+127);
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_quant_a(torch::Tensor A, torch::Tensor fp4, torch::Tensor scale, int M, int K){
quant_a_kernel<<<M, K/32>>>((const uint16_t*)A.data_ptr(),
(uint8_t*)fp4.data_ptr(),(uint8_t*)scale.data_ptr(),M,K);
}
"""
CPP_SRC = "void launch_quant_a(torch::Tensor,torch::Tensor,torch::Tensor,int,int);"
try:
_hip = load_inline(name='hw_quant_v3', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
functions=['launch_quant_a'], verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"])
_HAS_HW = True
except:
_HAS_HW = False
# Monkey-patch: replace Triton quant with HIP hw quant
import aiter.fused_moe as _fm
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 = {}
_patch_call_count = 0
def _fast_quant_moe_sort(input, sorted_ids, num_valid_ids, token_num, topk, block_size):
"""Replace Triton fused quant+sort with HIP hw quant + separate scale sort."""
global _patch_call_count
_patch_call_count += 1
if _patch_call_count <= 2:
print(f"[PATCH] Called #{_patch_call_count}: M={token_num} topk={topk} shape={input.shape} dtype={input.dtype}", file=sys.stderr)
# ONLY replace A quant (topk==1). Keep Triton for ALL intermediate quant.
if not _HAS_HW or input.dtype != torch.bfloat16 or topk != 1:
return _orig_quant(input, sorted_ids, num_valid_ids, token_num, topk, block_size)
M = input.shape[0]
K = input.shape[1]
dev = input.device
key = (M, K)
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]
# Fast HIP quantization (~5µs vs Triton ~72µs)
_hip.launch_quant_a(input, fp4_buf, scale_buf, M, K)
# Reinterpret as typed tensors
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))
# Sort only the scale (like M>1024 separate path)
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
# Apply the patch
import aiter.ops.triton.quant.fused_mxfp4_quant as _qmod
_qmod.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant_moe_sort
# Also patch the import in fused_moe module
_fm.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant_moe_sort
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 · 142 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 676583.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- Best per-shape strategy:- - Small batch (M≤128): KSPLIT=2 + BYPASS for E=257 → CKTile a16w4 (skip quant)- - Large batch (M>128): default CK kernels (no KSPLIT, no BYPASS)+ Monkey-patch fused_dynamic_mxfp4_quant_moe_sort with our HIP hardware quant.+ fused_moe still receives bf16 input (no KeyError), but A quant uses hardware instruction.+ Only patches for M=512 (2-stage path). M≤128 uses CKTile (no quant needed)."""- import os+ import os, sys+ 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_tfrom aiter import ActivationType, QuantType- from aiter.fused_moe import fused_moe+ # HIP A quantization kernel+ 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 quant_a_kernel(+ const uint16_t* __restrict__ A, uint8_t* __restrict__ fp4,+ uint8_t* __restrict__ scale, int M, int K+ ) {+ int row=blockIdx.x, grp=threadIdx.x;+ if(row>=M||grp>=K/32) return;+ 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);+ scale[row*(K/32)+grp]=(uint8_t)(si2+127);+ 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_quant_a(torch::Tensor A, torch::Tensor fp4, torch::Tensor scale, int M, int K){+ quant_a_kernel<<<M, K/32>>>((const uint16_t*)A.data_ptr(),+ (uint8_t*)fp4.data_ptr(),(uint8_t*)scale.data_ptr(),M,K);+ }+ """+ CPP_SRC = "void launch_quant_a(torch::Tensor,torch::Tensor,torch::Tensor,int,int);"++ try:+ _hip = load_inline(name='hw_quant_v3', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],+ functions=['launch_quant_a'], verbose=True,+ extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"])+ _HAS_HW = True+ except:+ _HAS_HW = False++ # Monkey-patch: replace Triton quant with HIP hw quant+ import aiter.fused_moe as _fm+ 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 = {}++ _patch_call_count = 0+ def _fast_quant_moe_sort(input, sorted_ids, num_valid_ids, token_num, topk, block_size):+ """Replace Triton fused quant+sort with HIP hw quant + separate scale sort."""+ global _patch_call_count+ _patch_call_count += 1+ if _patch_call_count <= 2:+ print(f"[PATCH] Called #{_patch_call_count}: M={token_num} topk={topk} shape={input.shape} dtype={input.dtype}", file=sys.stderr)+ # ONLY replace A quant (topk==1). Keep Triton for ALL intermediate quant.+ if not _HAS_HW or input.dtype != torch.bfloat16 or topk != 1:+ return _orig_quant(input, sorted_ids, num_valid_ids, token_num, topk, block_size)++ M = input.shape[0]+ K = input.shape[1]+ dev = input.device+ key = (M, K)+ 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]++ # Fast HIP quantization (~5µs vs Triton ~72µs)+ _hip.launch_quant_a(input, fp4_buf, scale_buf, M, K)++ # Reinterpret as typed tensors+ 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))++ # Sort only the scale (like M>1024 separate path)+ 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++ # Apply the patch+ import aiter.ops.triton.quant.fused_mxfp4_quant as _qmod+ _qmod.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant_moe_sort+ # Also patch the import in fused_moe module+ _fm.fused_dynamic_mxfp4_quant_moe_sort = _fast_quant_moe_sort++ from aiter.fused_moe import fused_moe_ACT = ActivationType.Silu_QT = QuantType.per_1x32⋯ 7 unchanged linesdep = 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:- # Small batch: CKTile a16w4 path (skip quant, block_m=16)os.environ['AITER_KSPLIT'] = '2'if E > 64:os.environ['AITER_BYPASS_TUNE_CONFIG'] = '1'- else:- os.environ.pop('AITER_BYPASS_TUNE_CONFIG', None)- else:- # Large batch: default CK kernels- os.environ.pop('AITER_KSPLIT', None)- os.environ.pop('AITER_BYPASS_TUNE_CONFIG', None)return fused_moe(hs, w1s, w2s, tw, ti,expert_mask=None, activation=_ACT, quant_type=_QT,
scrolls · 144 diff lines total
Best evidence level for this revision: reported
JSON