Skip to content
KernelIndex
Search⌘K

submission 569088

Harsh Gupta · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:707605d64964a897f85aa72cce6dd0c65ccc6d170f2bc57d381c644213669e64
license declaredunknown
license concludedunknown
authorsHarsh Gupta
imported2026-08-26

Techniques

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

persistent-kernel2. Persistent tile-loop (split-K shapes): atomicAdd + reduce
shared-memory__shared__ u8 ld[LDS_TOT];
split-k2. Persistent tile-loop (split-K shapes): atomicAdd + reduce
tile-k = 128constexpr int BM=16,BN=128,BK=128;
tile-m = 16constexpr int BM=16,BN=128,BK=128;
tile-n = 128constexpr int BM=16,BN=128,BK=128;

Kernel source

submission.py339 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
EXP-17: Final integrated kernel — dual-path dispatch.

Two kernel variants:
  1. Simple grid-mapped (non-split shapes): zero tile-loop overhead
  2. Persistent tile-loop (split-K shapes): atomicAdd + reduce

Recovers exp16 Stage A numbers for B0/B2/B3/B4/B5 while keeping
Stage B's split-K for B1.
"""

import os, sys
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

cuda_src = r"""
// EXP-17: Final integrated — dual-path dispatch
//
// Two kernel functions sharing all helpers:
//   1. simple_kern: grid-mapped, 1 block per MN-tile, all K-ranges in regs
//   2. split_kern:  persistent tile-loop, atomicAdd for split-K
//
// Recovers Stage A perf for non-split shapes + Stage B for split-K.

#include <torch/library.h>
#include <ATen/ATen.h>
#include <hip/hip_runtime.h>

using u8  = unsigned char;
using u32 = unsigned int;
using i32 = int;
using i64 = int64_t;
using f32 = float;
using bf16 = hip_bfloat16;
using u32x4 = u32 __attribute__((__vector_size__(16)));
using v4i32 = int __attribute__((__vector_size__(16)));
using v4f32 = float __attribute__((__vector_size__(16)));
using v8i32 = int __attribute__((__vector_size__(32)));

using buf_rsrc_vec = int32_t __attribute__((ext_vector_type(4)));
using lds_ptr_t = uint32_t __attribute__((address_space(3)))*;

extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    buf_rsrc_vec rsrc, lds_ptr_t lds_ptr, int size,
    int voffset, int soffset, int offset, int aux
) __asm("llvm.amdgcn.raw.buffer.load.lds");

struct __attribute__((packed)) buffer_resource_t {
    uint64_t ptr; uint32_t range; uint32_t config;
};
__device__ __forceinline__ buf_rsrc_vec make_srsrc(const void* base, uint32_t range) {
    buffer_resource_t r; r.ptr=reinterpret_cast<uint64_t>(base);
    r.range=range; r.config=0x00110000u;
    return *reinterpret_cast<const buf_rsrc_vec*>(&r);
}

constexpr int WARP=64,THREADS=512,NUM_WAVES=8;
constexpr int BM=16,BN=128,BK=128;
constexpr int MFMA_M=16,MFMA_N=16,MFMA_K=128;
constexpr int QGRP=32,GPB=BK/QGRP;
constexpr int KPS=512, NKT=KPS/BK;

__device__ __host__ constexpr int cdiv(int a,int b){return(a+b-1)/b;}

constexpr int A_ROW=KPS/2+16, A_RDW=A_ROW/4;
constexpr int A_SZ=BM*A_ROW, AS_SZ=BM*(KPS/QGRP);
constexpr int B_TILE=BN*BK/2, B_SZ=NKT*B_TILE;
constexpr int BS_SZ=BN*NKT*(int)sizeof(u32);
constexpr int O_A=0, O_AS=O_A+A_SZ, O_B=O_AS+AS_SZ, O_BS=O_B+B_SZ;
constexpr int LDS_TOT=O_BS+BS_SZ;
constexpr int G_BKB=BK/2, G_CB=WARP*16, G_RPC=G_CB/G_BKB, G_LPR=G_BKB/16;

// ============================================================
__device__ __forceinline__ i64 bsh_off(int n,int kc,int cpr){
    return((((i64)(n>>4)*cpr)+kc)<<8)+(((i64)(n&15))<<4);}
__device__ __forceinline__ u8 rd_bsc(const u8* sh,int n,int ks,int sn){
    int d0=n>>5,nm=n&31,d5=nm>>4,d3=nm&15;
    int d1=ks>>3,km=ks&7,d4=km>>2,d2=km&3;
    return sh[d0*(sn<<5)+d1*256+d2*64+d3*4+d4*2+d5];}
__device__ __forceinline__ u8 amax_e8m0(f32 a){
    u32 u=__float_as_uint(a);u=(u+0x200000u)&0xFF800000u;
    i32 s;if(__uint_as_float(u)==0.f)s=-127;
    else s=(i32)((u>>23)&0xFF)-127-2;
    if(s<-127)s=-127;if(s>127)s=127;return(u8)(s+127);}
__device__ __forceinline__ f32 scan16(const u32* p){
    f32 a=0.f;
    #pragma unroll
    for(int i=0;i<16;i++){u32 v=p[i];
        f32 lo=__uint_as_float((v&0xFFFFu)<<16),hi=__uint_as_float((v>>16)<<16);
        f32 al=lo<0?-lo:lo,ah=hi<0?-hi:hi;
        if(al>a)a=al;if(ah>a)a=ah;}return a;}
__device__ __forceinline__ u8 cvt_pair(u32 b,f32 s){
    u32 r;asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"=v"(r):"v"(b),"v"(s));
    return(u8)(r&0xFFu);}
__device__ __forceinline__ v4f32 mfma4(v4i32 a,v4i32 b,v4f32 c,i32 sa,i32 sb){
    v8i32 a8={a[0],a[1],a[2],a[3],0,0,0,0},b8={b[0],b[1],b[2],b[3],0,0,0,0};
    return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a8,b8,c,4,4,0,sa,0,sb);}

// ============================================================
// Shared prologue: quant A + bulk-load B + B-scale for one K-range
// ============================================================
__device__ __forceinline__ void do_prologue(
    int tid, int wid, int ln, int lr, int lc,
    int bm, int bn, int ks, int M, int N, int K, int B_sn,
    const u8* Ar, const u8* Bsc, buf_rsrc_vec br, u8* ld)
{
    u8* aq=ld+O_A; u8* asc=ld+O_AS;
    int kso=ks/QGRP;
    // A quant
    constexpr int TG=BM*(KPS/QGRP);
    for(int g=tid;g<TG;g+=THREADS){
        int row=g/(KPS/QGRP),grp=g%(KPS/QGRP),gr=bm+row;
        u32 sp[16];
        if(gr<M){
            i64 off=(i64)gr*K*2+(i64)ks*2+(i64)grp*QGRP*2;
            const u32* s=reinterpret_cast<const u32*>(Ar+off);
            u32x4 v0=*reinterpret_cast<const u32x4*>(s),
                  v1=*reinterpret_cast<const u32x4*>(s+4),
                  v2=*reinterpret_cast<const u32x4*>(s+8),
                  v3=*reinterpret_cast<const u32x4*>(s+12);
            sp[0]=v0[0];sp[1]=v0[1];sp[2]=v0[2];sp[3]=v0[3];
            sp[4]=v1[0];sp[5]=v1[1];sp[6]=v1[2];sp[7]=v1[3];
            sp[8]=v2[0];sp[9]=v2[1];sp[10]=v2[2];sp[11]=v2[3];
            sp[12]=v3[0];sp[13]=v3[1];sp[14]=v3[2];sp[15]=v3[3];
        }else{for(int i=0;i<16;i++)sp[i]=0;}
        f32 am=scan16(sp);u8 e=amax_e8m0(am);
        f32 sf=__uint_as_float(((u32)e)<<23);
        u32* d=reinterpret_cast<u32*>(aq+row*A_ROW+grp*(QGRP/2));
        #pragma unroll
        for(int j=0;j<4;j++){
            u8 b0=cvt_pair(sp[j*4],sf),b1=cvt_pair(sp[j*4+1],sf),
               b2=cvt_pair(sp[j*4+2],sf),b3=cvt_pair(sp[j*4+3],sf);
            d[j]=(u32)b0|((u32)b1<<8)|((u32)b2<<16)|((u32)b3<<24);}
        asc[row*(KPS/QGRP)+grp]=e;
    }
    // Bulk B load (NKT tiles)
    for(int t=0;t<NKT;t++){
        int k_off=ks+t*BK,kc=k_off/32;
        int rg=wid,gn=bn+rg*G_RPC+lr;
        int voff=(int)bsh_off(gn,kc+lc,B_sn);
        llvm_amdgcn_raw_buffer_load_lds(br,
            (lds_ptr_t)(&ld[O_B+t*B_TILE+rg*G_CB]),16,voff,0,0,0);
    }
    // B scale
    for(int idx=tid;idx<BN*NKT;idx+=THREADS){
        int local_n=idx/NKT,kt=idx%NKT;
        int gn=bn+local_n,sc=kso+kt*GPB;
        u32 pk=0x7F7F7F7Fu;
        if(gn<N){
            pk=((u32)rd_bsc(Bsc,gn,sc+0,B_sn))
              |((u32)rd_bsc(Bsc,gn,sc+1,B_sn)<<8)
              |((u32)rd_bsc(Bsc,gn,sc+2,B_sn)<<16)
              |((u32)rd_bsc(Bsc,gn,sc+3,B_sn)<<24);}
        reinterpret_cast<u32*>(ld+O_BS)[local_n*NKT+kt]=pk;
    }
    asm volatile("s_waitcnt vmcnt(0)":::"memory");
    __syncthreads();
}

// Shared K-loop: accumulate into acc
__device__ __forceinline__ v4f32 do_kloop(
    int ln, int sub_n, u8* ld, v4f32 acc)
{
    u8* aq=ld+O_A; u8* asc=ld+O_AS;
    for(int kt=0;kt<NKT;kt++){
        v4i32 af;{int mi=ln%16,kg=ln/16;
            u32* p=reinterpret_cast<u32*>(aq);
            int b2=mi*A_RDW+kt*(BK/8)+kg*4;
            af[0]=p[b2];af[1]=p[b2+1];af[2]=p[b2+2];af[3]=p[b2+3];}
        v4i32 bf;{int ni=ln%16,kg=ln/16;
            u32* p=reinterpret_cast<u32*>(ld+O_B+kt*B_TILE);
            int b2=(sub_n*MFMA_N+ni)*(BK/8)+kg*4;
            bf[0]=p[b2];bf[1]=p[b2+1];bf[2]=p[b2+2];bf[3]=p[b2+3];}
        i32 sa;{int mi=ln%16,kg=ln/16;
            sa=(i32)asc[mi*(KPS/QGRP)+kt*GPB+kg];}
        i32 sb;{int ni=ln%16,kg=ln/16;
            u32* bsp=reinterpret_cast<u32*>(ld+O_BS);
            u32 pk=bsp[(sub_n*MFMA_N+ni)*NKT+kt];
            sb=(i32)((pk>>(8*kg))&0xFFu);}
        acc=mfma4(af,bf,acc,sa,sb);
    }
    return acc;
}

// ============================================================
// KERNEL 1: Simple grid-mapped (non-split shapes)
// ============================================================
__global__ __launch_bounds__(512, 1)
void simple_kern(
    const bf16* __restrict__ A, const u8* __restrict__ Bsh,
    const u8* __restrict__ Bsc, bf16* __restrict__ C,
    int M, int N, int K, int B_sn)
{
    int tid=threadIdx.x, wid=tid/WARP, ln=tid%WARP;
    int sub_n=wid, lr=ln/G_LPR, lc=ln%G_LPR;
    int bm=blockIdx.y*BM, bn=blockIdx.x*BN;

    __shared__ u8 ld[LDS_TOT];
    const u8* Ar=reinterpret_cast<const u8*>(A);
    buf_rsrc_vec br=make_srsrc(Bsh,(uint32_t)((i64)N*(K/2)));

    v4f32 acc={0.f,0.f,0.f,0.f};
    int k_ranges=K/KPS;
    for(int kr=0;kr<k_ranges;kr++){
        do_prologue(tid,wid,ln,lr,lc,bm,bn,kr*KPS,M,N,K,B_sn,Ar,Bsc,br,ld);
        acc=do_kloop(ln,sub_n,ld,acc);
        __syncthreads();
    }
    // Direct bf16 store
    int nl=ln%16,rq=ln/16,gn=bn+sub_n*MFMA_N+nl,gmb=bm+rq*4;
    if(gn<N){
        #pragma unroll
        for(int r=0;r<4;r++){int gm=gmb+r;
            if(gm<M) C[(i64)gm*N+gn]=static_cast<bf16>(acc[r]);}}
}

// ============================================================
// KERNEL 2: Persistent tile-loop (split-K shapes)
// ============================================================
__global__ __launch_bounds__(512, 1)
void split_kern(
    const bf16* __restrict__ A, const u8* __restrict__ Bsh,
    const u8* __restrict__ Bsc, bf16* __restrict__ C, f32* __restrict__ ws,
    int M, int N, int K, int B_sn,
    int n_tiles, int spk, int total_units)
{
    int tid=threadIdx.x, wid=tid/WARP, ln=tid%WARP;
    int sub_n=wid, lr=ln/G_LPR, lc=ln%G_LPR;

    __shared__ u8 ld[LDS_TOT];
    const u8* Ar=reinterpret_cast<const u8*>(A);
    buf_rsrc_vec br=make_srsrc(Bsh,(uint32_t)((i64)N*(K/2)));

    int total_kr=K/KPS, kr_per=total_kr/spk;
    int upb=(total_units+gridDim.x-1)/gridDim.x;
    int my0=blockIdx.x*upb, my1=my0+upb;
    if(my1>total_units)my1=total_units;

    for(int u=my0;u<my1;u++){
        int nk=n_tiles*spk;
        int m_tile=u/nk, rem=u%nk;
        int k_split=rem/n_tiles, n_tile=rem%n_tiles;
        int bm=m_tile*BM, bn=n_tile*BN;
        int kr_start=k_split*kr_per;

        v4f32 acc={0.f,0.f,0.f,0.f};
        for(int kr=kr_start;kr<kr_start+kr_per;kr++){
            do_prologue(tid,wid,ln,lr,lc,bm,bn,kr*KPS,M,N,K,B_sn,Ar,Bsc,br,ld);
            acc=do_kloop(ln,sub_n,ld,acc);
            __syncthreads();
        }
        // atomicAdd to workspace
        int nl=ln%16,rq=ln/16,gn=bn+sub_n*MFMA_N+nl,gmb=bm+rq*4;
        if(gn<N){
            #pragma unroll
            for(int r=0;r<4;r++){int gm=gmb+r;
                if(gm<M) atomicAdd(&ws[(i64)gm*N+gn],acc[r]);}}
        __syncthreads();
    }
}

// f32 → bf16
__global__ __launch_bounds__(256)
void cvt_f32_bf16(const f32* __restrict__ src, bf16* __restrict__ dst, int n){
    int i=blockIdx.x*256+threadIdx.x;
    if(i<n) dst[i]=static_cast<bf16>(src[i]);
}

// ============================================================
at::Tensor fused_gemm(
    const at::Tensor& A, const at::Tensor& Bsh,
    const at::Tensor& Bsc, int64_t M, int64_t N, int64_t K)
{
    auto C=at::empty({M,N},A.options());
    const bf16* ap=reinterpret_cast<const bf16*>(A.data_ptr());
    const u8* bp=reinterpret_cast<const u8*>(Bsh.data_ptr());
    const u8* sp=reinterpret_cast<const u8*>(Bsc.data_ptr());
    bf16* cp=reinterpret_cast<bf16*>(C.data_ptr());
    int bsn=(int)(K/32);
    int nt=cdiv((int)N,BN), mt=cdiv((int)M,BM), mn=nt*mt;
    int total_kr=(int)(K/KPS);

    // Split-K for severely underfilled + large K
    int spk=1;
    if(mn<64 && K>1024){
        spk=256/mn;
        while(spk>1 && total_kr%spk!=0) spk--;
        if(spk<1) spk=1;
    }

    if(spk>1){
        int total=mn*spk;
        int gx=total<256?total:256;
        auto Cf=at::zeros({M,N},A.options().dtype(at::kFloat));
        f32* fp=reinterpret_cast<f32*>(Cf.data_ptr());
        split_kern<<<gx,THREADS,0,0>>>(ap,bp,sp,cp,fp,(int)M,(int)N,(int)K,bsn,nt,spk,total);
        int tot=(int)(M*N);
        cvt_f32_bf16<<<cdiv(tot,256),256,0,0>>>(fp,cp,tot);
    } else {
        dim3 grid(nt,mt);
        simple_kern<<<grid,THREADS,0,0>>>(ap,bp,sp,cp,(int)M,(int)N,(int)K,bsn);
    }
    return C;
}

TORCH_LIBRARY(mxfp4_exp17, m){
    m.def("fused_gemm(Tensor A,Tensor B_shuffle,Tensor B_scale_sh,int M,int N,int K)->Tensor");
    m.impl("fused_gemm",&fused_gemm);
}

"""

_ext = load_inline(
    name="mxfp4_exp17",
    cpp_sources=[""],
    cuda_sources=[cuda_src],
    extra_cflags=["-O3"],
    extra_cuda_cflags=["-O3"],
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
)
ops = torch.ops.mxfp4_exp17
print("exp17: dual-path kernel compiled OK", file=sys.stderr)


def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    B_shuffle = data[3]
    B_scale_sh = data[4]
    m, k = A.shape
    n = data[1].shape[0]
    return ops.fused_gemm(A, B_shuffle, B_scale_sh, m, n, k)
scrolls · 339 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