Skip to content
KernelIndex
Search⌘K

submission 725123

anAirdrop · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_s45.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-725123?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
19.2µs
#710 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0cbe2a7f0eaf8221635654b886ddb9a754486a71e40684c18cd7fcd20dbab548
license declaredunknown
license concludedunknown
authorsanAirdrop
imported2026-08-26

Kernel source

submission_s45.py135 lines
"""Hybrid: fused HIP MFMA for K≤512, AITER for K>512. Works in leaderboard mode!
Fused kernel: single launch, no quant caching needed (~14µs for K=512)
AITER path: quant + shuffle + GEMM (3 launches, ~20-34µs for large K)
"""
from task import input_t, output_t
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
import torch
import sys

_hip_module = None
def _get_module():
    global _hip_module
    if _hip_module is not None:
        return _hip_module
    from torch.utils.cpp_extension import load_inline
    hip_src = r"""
#include <hip/hip_runtime.h>
typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef int __attribute__((ext_vector_type(8))) int8_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;

__device__ __forceinline__ unsigned char load_b_scale(
    const unsigned char* bs, int row, int col, int sn) {
    int d0=row/32, d1=(row%32)/16, d2=row%16, d3=col/8, d4=(col%8)/4, d5=col%4;
    return bs[d0*(sn/8*256)+d3*256+d5*64+d2*4+d4*2+d1];
}

__device__ __forceinline__ unsigned char fp4_code(float ax) {
    return (unsigned char)(
        (ax > 0.25f) + (ax >= 0.75f) + (ax > 1.25f) + (ax >= 1.75f) +
        (ax > 2.5f) + (ax >= 3.5f) + (ax > 5.0f));
}

extern "C" __global__ void fused_mfma_gemm(
    const unsigned short* __restrict__ A_bf16,
    const unsigned char* __restrict__ B_q,
    const unsigned char* __restrict__ B_scale_sh,
    unsigned short* __restrict__ C,
    int M, int N, int K,
    int stride_a, int stride_bq, int pad_Ks, int stride_c) {
    int m_tile=blockIdx.x*16, n_tile=blockIdx.y*16;
    int lane=threadIdx.x, m_lane=lane&15, k_lane=lane>>4;
    int a_row=m_tile+m_lane, b_row=n_tile+m_lane;
    float4_vec acc={0,0,0,0};
    int Kg=K/32;
    for(int k128=0; k128<K; k128+=128){
        int a_k_start=k128+k_lane*32;
        float max_abs=0.0f; float a_f[32];
        if(a_row<M){
            const unsigned short* ap=A_bf16+a_row*stride_a+a_k_start;
            for(int i=0;i<32;i++){
                float v=__uint_as_float((unsigned int)ap[i]<<16);
                a_f[i]=v; max_abs=fmaxf(max_abs,fabsf(v));
            }
        } else { for(int i=0;i<32;i++) a_f[i]=0.0f; }
        max_abs=fmaxf(max_abs,1.175494e-38f);
        unsigned int mi=__float_as_uint(max_abs);
        mi=(mi+0x200000u)&0x7F800000u;
        int su=(int)(mi>>23)-127-2;
        su=min(max(su,-127),127);
        float qs=exp2f((float)(-su));
        unsigned char asc=(unsigned char)(su+127);
        unsigned char pk[16];
        for(int i=0;i<16;i++){
            float v0=a_f[2*i]*qs, v1=a_f[2*i+1]*qs;
            unsigned char s0=(v0<0.0f)?8:0, s1=(v1<0.0f)?8:0;
            pk[i]=(s0|fp4_code(fabsf(v0)))|((s1|fp4_code(fabsf(v1)))<<4);
        }
        int4_vec a_reg; __builtin_memcpy(&a_reg,pk,16);
        int8_vec a_pad={a_reg[0],a_reg[1],a_reg[2],a_reg[3],0,0,0,0};
        int4_vec b_reg;
        if(b_row<N) b_reg=*reinterpret_cast<const int4_vec*>(B_q+b_row*stride_bq+(k128/2)+k_lane*16);
        else b_reg={0,0,0,0};
        int8_vec b_pad={b_reg[0],b_reg[1],b_reg[2],b_reg[3],0,0,0,0};
        int akg=k128/32+k_lane;
        unsigned char bsc=0x7F;
        if(b_row<N&&akg<Kg) bsc=load_b_scale(B_scale_sh,b_row,akg,pad_Ks);
        acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_pad,b_pad,acc,4,4,0,(int)asc,0,(int)bsc);
    }
    int cm=m_tile+k_lane*4, cn=n_tile+m_lane;
    for(int i=0;i<4;i++){
        if(cm+i<M&&cn<N) C[(cm+i)*stride_c+cn]=(unsigned short)(__float_as_uint(acc[i])>>16);
    }
}

torch::Tensor launch_fused(torch::Tensor A, torch::Tensor B_q, torch::Tensor B_scale_sh,
                           int M, int N, int K, int pad_Ks) {
    auto C=torch::empty({M,N},torch::dtype(torch::kBFloat16).device(A.device()));
    fused_mfma_gemm<<<dim3((M+15)/16,(N+15)/16),dim3(64)>>>(
        reinterpret_cast<const unsigned short*>(A.data_ptr<at::BFloat16>()),
        B_q.data_ptr<unsigned char>(),B_scale_sh.data_ptr<unsigned char>(),
        reinterpret_cast<unsigned short*>(C.data_ptr<at::BFloat16>()),
        M,N,K,(int)A.stride(0),(int)B_q.stride(0),pad_Ks,(int)C.stride(0));
    return C;
}
"""
    cpp_src = "torch::Tensor launch_fused(torch::Tensor,torch::Tensor,torch::Tensor,int,int,int,int);"
    try:
        _hip_module = load_inline(name='fused_mfma_rtne', cpp_sources=[cpp_src],
            cuda_sources=[hip_src], functions=['launch_fused'],
            extra_cuda_cflags=['-O3','--offload-arch=gfx950'], verbose=False)
    except Exception as e:
        print(f"[s45] compile FAILED: {e}", file=sys.stderr)
    return _hip_module

import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n = B.shape[0]
    mod = _get_module()

    if k <= 512 and mod is not None:
        # Fused: single launch, ~14µs for K=512 (vs ~20µs AITER 3-launch)
        pad_Ks = B_scale_sh.view(torch.uint8).shape[1]
        return mod.launch_fused(A, B_q.view(torch.uint8), B_scale_sh.view(torch.uint8), m, n, k, pad_Ks)
    else:
        # AITER: 3 launches but fast ASM GEMM for large K
        A_q_fp4, A_scale = dynamic_mxfp4_quant(A)
        A_q = A_q_fp4.view(dtypes.fp4x2)
        scale = A_scale; sc_m, sc_n = scale.shape
        pad_m = ((sc_m+255)//256)*256; pad_n = ((sc_n+7)//8)*8
        if pad_m != sc_m or pad_n != sc_n:
            sp = torch.empty(pad_m, pad_n, dtype=scale.dtype, device=scale.device)
            sp[:sc_m,:sc_n] = scale; scale = sp
        sm, sn = scale.shape
        A_scale_sh = scale.view(sm//32,2,16,sn//8,2,4).permute(0,3,5,2,4,1).contiguous().view(sm,sn).view(dtypes.fp8_e8m0)
        return aiter.gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 135 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