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
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-kernel
2. Persistent tile-loop (split-K shapes): atomicAdd + reduceshared-memory
__shared__ u8 ld[LDS_TOT];split-k
2. Persistent tile-loop (split-K shapes): atomicAdd + reducetile-k = 128
constexpr int BM=16,BN=128,BK=128;tile-m = 16
constexpr int BM=16,BN=128,BK=128;tile-n = 128
constexpr 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