Skip to content
KernelIndex
Search⌘K

submission 683418

AlexNebula · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-683418?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
14.3µs
#495 of 1143
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:13a0a7a189bf4b25ade2e7a4c9a908550be57408b37ca904f8d9d0a0bfc31c8f
license declaredunknown
license concludedunknown
authorsAlexNebula
imported2026-08-26

Techniques

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

fp4MXFP4-MM v56: CVT(val, scale) 直接用 — 去掉手动乘+clamp。
vector-width = uint4const uint4* p128=(const uint4*)(A+row*k+col*32+quarter*8);

Kernel source

submission.py160 lines
"""
MXFP4-MM v56: CVT(val, scale) 直接用 — 去掉手动乘+clamp。
CVT scale = 除法,所以传 2^(scale_unbiased) 让 CVT 做 val / 2^su = val * 2^(-su)。
之前 v14 失败因为 scale 传反了。现在传正确的 divisor。
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ["CXX"] = "clang++"
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

CPP_SRC = """
#include <torch/extension.h>
torch::Tensor fused_single(torch::Tensor A, torch::Tensor B, torch::Tensor Bq,
    torch::Tensor C, int m, int n, int k);
void hip_quant_a(torch::Tensor A, torch::Tensor Aq, torch::Tensor As,
    int m, int k, int sm, int sn);
"""

HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
typedef uint8_t fp4x2_t;
typedef fp4x2_t fp4_vec32 __attribute__((ext_vector_type(32)));
typedef float f32_vec4 __attribute__((ext_vector_type(4)));
#define FP4_TYPE 4

// Single kernel for small K: CVT(val, divisor) — no manual multiply or clamp
__global__ __launch_bounds__(64)
void single_kernel(const __hip_bfloat16* A, const __hip_bfloat16* B,
    const uint8_t* Bq, __hip_bfloat16* C, int m, int n, int k, int bqs) {
    int bm=blockIdx.x, bn=blockIdx.y, lane=threadIdx.x, rg=lane/16, rit=lane%16;
    int a_row=bm*16+rit, b_col=bn*16+rit;
    f32_vec4 c={};
    for(int kt=0; kt<k/128; kt++){
        fp4_vec32 ar={}; uint8_t ae=0;
        if(a_row<m){
            const uint32_t* ap32=(const uint32_t*)(A+a_row*k+kt*128+rg*32);
            uint32_t mx=0; float v[32];
            #pragma unroll
            for(int i=0;i<16;i++){
                uint32_t pair=ap32[i]; mx=max(mx,max(pair&0x7FFFu,(pair>>16)&0x7FFFu));
                uint16_t lo=(uint16_t)(pair&0xFFFF), hi=(uint16_t)(pair>>16);
                v[i*2]=__bfloat162float(*(__hip_bfloat16*)&lo);
                v[i*2+1]=__bfloat162float(*(__hip_bfloat16*)&hi);
            }
            if(mx>0){
                uint16_t mb=(uint16_t)mx; float ax=__bfloat162float(*(__hip_bfloat16*)&mb);
                uint32_t ab=__float_as_uint(ax); uint32_t rb=(ab+0x200000u)&0xFF800000u;
                int su=max(-127,min(127,(int)((rb>>23)&0xFF)-129)); ae=(uint8_t)(su+127);
                // CVT divisor = 2^(su) — CVT does val / divisor = val * 2^(-su) = val * quant_scale
                float divisor=__uint_as_float((uint32_t)(su+127)<<23);
                uint32_t pk[4];
                #pragma unroll
                for(int g=0;g<4;g++){uint32_t p=0;
                    // CVT handles clamp internally when using proper scale
                    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+0],v[g*8+1],divisor,0);
                    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+2],v[g*8+3],divisor,1);
                    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+4],v[g*8+5],divisor,2);
                    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+6],v[g*8+7],divisor,3);pk[g]=p;}
                *(uint4*)(&ar)=*(uint4*)(pk);
            }
        }
        fp4_vec32 br={}; uint8_t be=127;
        if(b_col<n){
            *(uint4*)(&br)=*(const uint4*)(Bq+b_col*bqs+kt*64+rg*16);
            const uint32_t* bp=(const uint32_t*)(B+b_col*k+kt*128+rg*32);
            uint32_t bmx=0;
            #pragma unroll
            for(int i=0;i<16;i++){uint32_t bv=bp[i];bmx=max(bmx,max(bv&0x7FFFu,(bv>>16)&0x7FFFu));}
            if(bmx>0){uint16_t mb=(uint16_t)bmx;float bax=__bfloat162float(*(__hip_bfloat16*)&mb);
                uint32_t ab=__float_as_uint(bax);uint32_t rb=(ab+0x200000u)&0xFF800000u;
                be=(uint8_t)(max(-127,min(127,(int)((rb>>23)&0xFF)-129))+127);}
        }
        c=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(ar,br,c,FP4_TYPE,FP4_TYPE,0,(int)ae,0,(int)be);
    }
    for(int i=0;i<4;i++){int gr=bm*16+rg*4+i,gc=bn*16+rit;
        if(gr<m&&gc<n) C[gr*n+gc]=(__hip_bfloat16)c[i];}
}
torch::Tensor fused_single(torch::Tensor A,torch::Tensor B,torch::Tensor Bq,
    torch::Tensor C,int m,int n,int k){
    single_kernel<<<dim3((m+15)/16,(n+15)/16),64>>>(
        (__hip_bfloat16*)A.data_ptr(),(__hip_bfloat16*)B.data_ptr(),
        (uint8_t*)Bq.data_ptr(),(__hip_bfloat16*)C.data_ptr(),m,n,k,(int)Bq.stride(0));
    return C;}

// Large K: HIP A quant with CVT(val, divisor) — no manual multiply or clamp
__global__ void quant_a_kernel(const __hip_bfloat16* A, uint8_t* Aq, uint8_t* As,
    int m, int k, int aq_stride, int sm, int sn) {
    int idx=blockIdx.x*blockDim.x+threadIdx.x;
    int total=m*(k/32)*4; if(idx>=total) return;
    int block_id=idx/4, quarter=idx%4;
    int row=block_id/(k/32), col=block_id%(k/32);
    const uint4* p128=(const uint4*)(A+row*k+col*32+quarter*8);
    uint4 data=*p128; uint32_t vals[4]={data.x,data.y,data.z,data.w};
    uint32_t partial_max=0; float v[8];
    #pragma unroll
    for(int i=0;i<4;i++){uint32_t pair=vals[i];
        partial_max=max(partial_max,max(pair&0x7FFFu,(pair>>16)&0x7FFFu));
        uint16_t lo=(uint16_t)(pair&0xFFFF),hi=(uint16_t)(pair>>16);
        v[i*2]=__bfloat162float(*(__hip_bfloat16*)&lo);v[i*2+1]=__bfloat162float(*(__hip_bfloat16*)&hi);}
    uint32_t mx=partial_max; mx=max(mx,__shfl_xor(mx,1)); mx=max(mx,__shfl_xor(mx,2));
    uint8_t e=0; float divisor=1.0f;
    if(mx>0){uint16_t mb=(uint16_t)mx;float ax=__bfloat162float(*(__hip_bfloat16*)&mb);
        uint32_t ab=__float_as_uint(ax);uint32_t rb=(ab+0x200000u)&0xFF800000u;
        int su=max(-127,min(127,(int)((rb>>23)&0xFF)-129));e=(uint8_t)(su+127);
        divisor=__uint_as_float((uint32_t)(su+127)<<23);}
    if(quarter==0){int d0=row/32,d1=(row/16)%2,d2=row%16,d3=col/8,d4=(col/4)%2,d5=col%4;
        int si=d0*(sn/8*256)+d3*256+d5*64+d2*4+d4*2+d1;
        if(si<sm*sn) As[si]=e;}
    // CVT(val, divisor) — no manual multiply or clamp needed
    uint32_t pk=0;
    pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[0],v[1],divisor,0);
    pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[2],v[3],divisor,1);
    pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[4],v[5],divisor,2);
    pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[6],v[7],divisor,3);
    *(uint32_t*)(Aq+row*aq_stride+col*16+quarter*4)=pk;
}
void hip_quant_a(torch::Tensor A,torch::Tensor Aq,torch::Tensor As,int m,int k,int sm,int sn){
    int total=m*(k/32)*4;
    quant_a_kernel<<<(total+255)/256,256>>>((__hip_bfloat16*)A.data_ptr(),(uint8_t*)Aq.data_ptr(),
        (uint8_t*)As.data_ptr(),m,k,(int)Aq.stride(0),sm,sn);}
"""

module = load_inline(name='mxfp4_mm', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
    functions=['fused_single','hip_quant_a'], verbose=True,
    extra_cuda_cflags=["--offload-arch=gfx950","-std=c++20","-O3","-mllvm","-amdgpu-early-inline-all=true"])

_bufs = {}

def custom_kernel(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape; n = B.shape[0]
    bk = (m, n, k)
    if bk not in _bufs:
        _bufs.clear()
        sm = ((m+255)//256)*256; sn = ((k//32+7)//8)*8
        _bufs[bk] = {
            'C': torch.empty((m,n), dtype=torch.bfloat16, device=A.device),
            'Aq': torch.empty((m,k//2), dtype=torch.uint8, device=A.device),
            'As': torch.zeros((sm,sn), dtype=torch.uint8, device=A.device),
            'sm': sm, 'sn': sn,
        }
    C = _bufs[bk]['C']
    if k <= 1024:
        return module.fused_single(A, B, B_q.view(torch.uint8), C, m, n, k)
    else:
        sm, sn = _bufs[bk]['sm'], _bufs[bk]['sn']
        if sm > m or sn > k//32:
            _bufs[bk]['As'].zero_()
        module.hip_quant_a(A, _bufs[bk]['Aq'], _bufs[bk]['As'], m, k, sm, sn)
        return aiter.gemm_a4w4(
            _bufs[bk]['Aq'].view(dtypes.fp4x2), B_shuffle,
            _bufs[bk]['As'].view(dtypes.fp8_e8m0), B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 160 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