Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
155.9µs
#250 of 782
2026-04-01

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.

fp4const 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_inline
from task import input_t, output_t
from 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 lines
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:
- # 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