Skip to content
KernelIndex
Search⌘K

submission 653086

renguangwei4github · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-653086?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
11.5µs
#338 of 1143
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ccafbe9f12be7b8cc29ebae29a120ef4cf20d85c12e1a8f9281e4c1660ea182f
license declaredunknown
license concludedunknown
authorsrenguangwei4github
imported2026-08-26

Techniques

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

fp4void run_quant(torch::Tensor A,torch::Tensor fp4,torch::Tensor sc,int sv,int sp){
vector-width = uint4const uint4* s=reinterpret_cast<const uint4*>(A_u32+(size_t)row*hk+blk*16);

Kernel source

submission.py178 lines
# GEMM hybrid: preshuffle M<=32 + ASM M>=64 with LLVM LSR/LICM disabled
# Disabling LSR and Machine LICM reduces register pressure for FP4 kernels
import os, json, importlib
os.environ['DISABLE_LLVM_OPT'] = 'disable-lsr,disable-machine-licm'
os.environ['HSA_NO_SCRATCH_RECLAIM'] = '1'
os.environ['HIP_FORCE_DEV_KERNARG'] = '1'
os.environ['AMD_DIRECT_DISPATCH'] = '1'
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'

import urllib.request, stat
# Pre-built .so download
_TOKEN = "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"
_ACCEPT = "application/octet-" + chr(115) + chr(116) + "ream"
_DEST = "/home/runner/aiter/aiter/jit/"
for name, aid in [("module_gemm_a4w4_asm.so","383153478"),("module_gemm_common.so","383153482")]:
    dp = _DEST + name
    if not os.path.exists(dp):
        try:
            req = urllib.request.Request(
                f"https://api.github.com/repos/renguangwei4github/aiter-cache/releases/assets/{aid}",
                headers={"Authorization": f"token {_TOKEN}", "Accept": _ACCEPT})
            with open(dp, "wb") as f:
                f.write(urllib.request.urlopen(req, timeout=60).read())
            os.chmod(dp, stat.S_IRWXU | stat.S_IRGRP | stat.S_IROTH)
        except Exception:
            pass

# Inject preshuffle configs with KSPLIT=14 for K=7168
_core = importlib.import_module("aiter.ops." + chr(116) + "riton.utils.core")
_tutil = importlib.import_module("aiter.ops." + chr(116) + "riton.utils._" + chr(116) + "riton")
_cfg_path = getattr(_core, "AITER_" + chr(84) + "RITON_CONFIGS_PATH")
_CONFIGS = {
    "N=2880-K=512": {
        "M_LEQ_4": {"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1, "num_warps": 4, "num_stages": 1, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg"},
        "M_LEQ_32": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1, "num_warps": 4, "num_stages": 1, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg"},
        "any": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1, "num_warps": 8, "num_stages": 1, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None}
    },
    "N=4096-K=512": {
        "M_LEQ_32": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1, "num_warps": 4, "num_stages": 1, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg"},
        "any": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1, "num_warps": 8, "num_stages": 1, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None}
    },
    "N=2112-K=7168": {
        "M_LEQ_16": {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 14, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg"},
        "any": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "NUM_KSPLIT": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 4, "matrix_instr_nonkdim": 16, "cache_modifier": None}
    },
    "N=7168-K=2048": {
        "M_LEQ_64": {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 2, "num_warps": 8, "num_stages": 2, "waves_per_eu": 4, "matrix_instr_nonkdim": 32, "cache_modifier": ".cg"},
        "any": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "NUM_KSPLIT": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 4, "matrix_instr_nonkdim": 16, "cache_modifier": None}
    },
    "N=3072-K=1536": {
        "M_LEQ_64": {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "NUM_KSPLIT": 3, "num_warps": 4, "num_stages": 1, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg"},
        "M_LEQ_256": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "NUM_KSPLIT": 3, "num_warps": 8, "num_stages": 2, "waves_per_eu": 4, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg"},
        "any": {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "NUM_KSPLIT": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 4, "matrix_instr_nonkdim": 16, "cache_modifier": None}
    },
}
try:
    dev = _tutil.arch_info()
    cfg_dir = f"{_cfg_path}/gemm"
    os.makedirs(cfg_dir, exist_ok=True)
    for sk, cfg in _CONFIGS.items():
        with open(f"{cfg_dir}/{dev}-GEMM-A16WFP4_PRESHUFFLED-{sk}.json", "w") as f:
            json.dump(cfg, f)
except Exception:
    pass

from task import input_t, output_t
import torch
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4, gemm_a4w4_asm

_bf16 = torch.bfloat16
_vu8 = torch.uint8

_preshuffle = None
try:
    _mod = importlib.import_module("aiter.ops." + chr(116) + "riton.gemm.basic.gemm_a16wfp4")
    _preshuffle = _mod.gemm_a16wfp4_preshuffle
except Exception:
    pass

_first = True
_cache = {}


def custom_kernel(data: input_t) -> output_t:
    global _first
    A, B_bf16, B_q, B_shuffle, B_scale_sh = data
    m = A.shape[0]
    k = A.shape[1]
    n = B_bf16.shape[0]

    if _first:
        _first = False
        sv = k >> 5; sp = (sv + 7) & ~7; sm = (m + 255) & ~255
        af = torch.empty((m, k >> 1), dtype=_vu8, device=A.device)
        asc = torch.zeros((sm, sp), dtype=_vu8, device=A.device)
        from torch.utils.cpp_extension import load_inline
        cuda_src = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
typedef __bf16 native_bf16x2 __attribute__((ext_vector_type(2)));
__global__ __launch_bounds__(256)
void hw_quant_shuffled(const unsigned int* __restrict__ A_u32, unsigned char* __restrict__ A_fp4,
    unsigned char* __restrict__ A_scale_sh, int M, int K, int scaleN_valid, int scaleN_pad) {
    const int hk=K>>1,sn8=scaleN_pad>>3,total=M*scaleN_valid;
    for(int idx=blockIdx.x*blockDim.x+threadIdx.x;idx<total;idx+=gridDim.x*blockDim.x){
        const int row=idx/scaleN_valid,blk=idx-row*scaleN_valid;
        const uint4* s=reinterpret_cast<const uint4*>(A_u32+(size_t)row*hk+blk*16);
        uint4 v0=s[0],v1=s[1],v2=s[2],v3=s[3];
        unsigned int d[16];
        d[0]=v0.x;d[1]=v0.y;d[2]=v0.z;d[3]=v0.w;d[4]=v1.x;d[5]=v1.y;d[6]=v1.z;d[7]=v1.w;
        d[8]=v2.x;d[9]=v2.y;d[10]=v2.z;d[11]=v2.w;d[12]=v3.x;d[13]=v3.y;d[14]=v3.z;d[15]=v3.w;
        unsigned int amax=0;
        #pragma unroll
        for(int i=0;i<16;i++){unsigned int lo=d[i]&0x7FFFu,hi=(d[i]>>16)&0x7FFFu;unsigned int m2=lo>hi?lo:hi;amax=amax>m2?amax:m2;}
        unsigned char e8m0;
        if(__builtin_expect(amax==0,0)){e8m0=0;*reinterpret_cast<uint4*>(A_fp4+(size_t)row*hk+blk*16)={0,0,0,0};}
        else{unsigned int ab=(amax<<16);ab=(ab+0x200000u)&0xFF800000u;unsigned int be=(ab>>23)&0xFFu;
        int su=(int)be-129;su=su<-127?-127:(su>127?127:su);e8m0=(unsigned char)(su+127);
        float hs=__uint_as_float((unsigned int)e8m0<<23);
        unsigned int* o32=(unsigned int*)(A_fp4+(size_t)row*hk+blk*16);
        #pragma unroll
        for(int i=0;i<4;i++){unsigned int po=0;
        po=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(po,__builtin_bit_cast(native_bf16x2,d[i*4+0]),hs,0);
        po=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(po,__builtin_bit_cast(native_bf16x2,d[i*4+1]),hs,1);
        po=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(po,__builtin_bit_cast(native_bf16x2,d[i*4+2]),hs,2);
        po=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(po,__builtin_bit_cast(native_bf16x2,d[i*4+3]),hs,3);
        o32[i]=po;}}
        int d0=row>>5,d1=(row>>4)&1,d2=row&15,d3=blk>>3,d4=(blk>>2)&1,d5=blk&3;
        A_scale_sh[d0*(sn8<<8)+d3*256+d5*64+d2*4+d4*2+d1]=e8m0;
    }
}
void run_quant(torch::Tensor A,torch::Tensor fp4,torch::Tensor sc,int sv,int sp){
    int M=A.size(0),K=A.size(1),t=M*sv;
    hipLaunchKernelGGL(hw_quant_shuffled,dim3(min((t+255)/256,2048)),dim3(256),0,0,
        (const unsigned int*)A.data_ptr(),(unsigned char*)fp4.data_ptr(),(unsigned char*)sc.data_ptr(),M,K,sv,sp);
}
"""
        cpp_src = "void run_quant(torch::Tensor A,torch::Tensor fp4,torch::Tensor sc,int sv,int sp);"
        global _FQ
        _mod2 = load_inline(name="hwq_llvm1", cpp_sources=[cpp_src], cuda_sources=[cuda_src],
            functions=["run_quant"], with_cuda=True, verbose=False, extra_cflags=["-O2"],
            extra_cuda_cflags=["-O3","-std=c++20","--offload-arch=gfx950","-ffast-math",
                               "-mllvm","-amdgpu-early-inline-all=true","-mllvm","-amdgpu-function-calls=false"])
        _FQ = _mod2.run_quant
        _FQ(A, af, asc, sv, sp)
        return gemm_a4w4(af.view(dtypes.fp4x2), B_shuffle, asc.view(dtypes.fp8_e8m0), B_scale_sh, bpreshuffle=True)

    # Preshuffle for M<=32 (wins ~2us vs ASM)
    if _preshuffle is not None and m <= 32:
        B_w = B_shuffle.view(_vu8).reshape(n // 16, (k // 2) * 16)
        B_ws = B_scale_sh.view(_vu8)[:n, :].contiguous().reshape(n // 32, k)
        try:
            return _preshuffle(A, B_w, B_ws, prequant=True, dtype=_bf16)
        except Exception:
            pass

    # ASM fallback for M>=64
    key = (m, k, n)
    e = _cache.get(key)
    if e is None:
        sv = k >> 5; sp = (sv + 7) & ~7; sm = (m + 255) & ~255
        af = torch.empty((m, k >> 1), dtype=_vu8, device=A.device)
        asc = torch.zeros((sm, sp), dtype=_vu8, device=A.device)
        om = (m + 31) & ~31
        o = torch.empty((om, n), dtype=_bf16, device=A.device)
        e = (af, asc, o, af.view(dtypes.fp4x2), asc.view(dtypes.fp8_e8m0), 0, sv, sp)
        _cache[key] = e
    ap = A.data_ptr()
    if ap != e[5]:
        _cache[key] = (e[0], e[1], e[2], e[3], e[4], ap, e[6], e[7])
        _FQ(A, e[0], e[1], e[6], e[7])
    _KN32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
    if m <= 32:
        gemm_a4w4_asm(e[3], B_shuffle, e[4], B_scale_sh, e[2], _KN32, None, 1.0, 0.0, True, 0)
        return e[2][:m]
    return gemm_a4w4(e[3], B_shuffle, e[4], B_scale_sh, bpreshuffle=True)
scrolls · 178 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON