submission 748664
Jayluci4 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 120 lines, June 9 Researcher Reciprocity License v1.0.
submission_gemm_final2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-748664?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:5460e03ea583c0c7f5aa0271a5a38d77508fb7ca8cb98d4f491832a4f7f57a47
license declaredunknown
license concludedunknown
authorsJayluci4
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
fp4,ssh=_SC[key];sv=triton.cdiv(k,32);sp=triton.cdiv(sv,8)*8stages = 1
num_warps=nw,waves_per_eu=0,num_stages=1)Kernel source
submission_gemm_final2.py120 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
GEMM final2: hybrid dispatch.
Small M (4-32): gemm_a16wfp4_preshuffle (zero-copy, single fused launch)
Large M (64-256): fused quant+shuffle + ASM GEMM (v6 path, proven 14-15us)
"""
import gc, sys, os
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
gc.disable()
sys.setswitchinterval(1000.0)
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# Preshuffle path
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
# ASM path
import aiter
from aiter import dtypes
_F = dtypes.fp4x2; _E = dtypes.fp8_e8m0; _B = dtypes.bf16
_K32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_K64 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
# Fused quant+shuffle kernel (from v6, proven bit-exact)
@triton.jit
def _mxfp4_qop(x, BSN, BSM, QBS):
nqb: tl.constexpr = BSN // QBS
x = x.reshape(BSM, nqb, QBS)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
su = tl.log2(amax).floor() - 2; su = tl.clamp(su, min=-127, max=127)
bs = su.to(tl.uint8) + 127; qs = tl.exp2(-su)
qx = (x * qs).to(tl.uint32, bitcast=True); s = qx & 0x80000000; qx = qx ^ s
qf = qx.to(tl.float32, bitcast=True)
sat = qf >= 6; den = (~sat) & (qf < 1); nor = ~(sat | den)
DE: tl.constexpr = 149 << 23; DF: tl.constexpr = tl.cast(DE, tl.float32, bitcast=True)
dx = qf + DF; dx = dx.to(tl.uint32, bitcast=True) - DE; dx = dx.to(tl.uint8)
nx = qx.to(tl.int32); mo = (nx >> 22) & 1
nx += ((1-127) << 23) + (1 << 21) - 1 + mo; nx = (nx >> 22).to(tl.uint8)
e = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e = tl.where(nor, nx, e); e = tl.where(den, dx, e)
e = e | (s >> 28).to(tl.uint8)
e = tl.reshape(e, [BSM, nqb, QBS // 2, 2])
ev, od = tl.split(e)
return (ev | (od << 4)).reshape(BSM, BSN // 2), bs.reshape(BSM, nqb)
@triton.heuristics({"EV": lambda a: a["M"]%a["BM"]==0 and a["N"]%(a["BN"]*a["NI"])==0})
@triton.jit
def _fqs(xp,fp,bp,sxm,sxn,sfm,sfn,M,N,SV,SP,BM:tl.constexpr,BN:tl.constexpr,
NI:tl.constexpr,NS:tl.constexpr,QBS:tl.constexpr,EV:tl.constexpr):
pm=tl.program_id(0);sn=tl.program_id(1)*NI
nq:tl.constexpr=BN//QBS
for pn in tl.range(sn,min(sn+NI,N),num_stages=NS):
om=pm*BM+tl.arange(0,BM);on=pn*BN+tl.arange(0,BN)
offs=om[:,None]*tl.cast(sxm,tl.int64)+on[None,:]*tl.cast(sxn,tl.int64)
if EV: x=tl.load(xp+offs,cache_modifier=".cg").to(tl.float32)
else: x=tl.load(xp+offs,mask=(om<M)[:,None]&(on<N)[None,:],cache_modifier=".cg").to(tl.float32)
ot,bse=_mxfp4_qop(x,BN,BM,QBS)
oom=pm*BM+tl.arange(0,BM);oon=pn*BN//2+tl.arange(0,BN//2)
ooffs=oom[:,None]*tl.cast(sfm,tl.int64)+oon[None,:]*tl.cast(sfn,tl.int64)
if EV: tl.store(fp+ooffs,ot)
else: tl.store(fp+ooffs,ot,mask=(oom<M)[:,None]&(oon<N//2)[None,:])
br=pm*BM+tl.arange(0,BM);bc=pn*nq+tl.arange(0,nq)
r32=br[:,None]%32
bfo=((br[:,None]//32)*(32*SP)+(bc[None,:]//8)*256+(bc[None,:]%4)*64+(r32%16)*4+((bc[None,:]%8)//4)*2+r32//16)
tl.store(bp+bfo,bse,mask=(br<M)[:,None]&(bc<SV)[None,:])
_SC={}; _OA={}
def _qa(x):
m,k=x.shape;key=(m,k)
if key not in _SC:
sv=triton.cdiv(k,32);sp=triton.cdiv(sv,8)*8;sm=triton.cdiv(m,256)*256
_SC[key]=(torch.empty((m,k//2),dtype=torch.uint8,device=x.device),
torch.empty((sm,sp),dtype=torch.uint8,device=x.device))
fp4,ssh=_SC[key];sv=triton.cdiv(k,32);sp=triton.cdiv(sv,8)*8
if m<=32: ni,bm,bn,nw,ks=1,triton.next_power_of_2(m),32,1,1
else: ni,bm,bn,nw,ks=4,64,64,4,2;(bm,bn)=(32,128) if k<=16384 else (bm,bn)
if k<=1024: ni,ks,nw=1,1,4;bn=max(32,min(256,triton.next_power_of_2(k)));bm=min(8,triton.next_power_of_2(m))
_fqs[(triton.cdiv(m,bm),triton.cdiv(k,bn*ni))](
x,fp4,ssh,*x.stride(),*fp4.stride(),M=m,N=k,SV=sv,SP=sp,QBS=32,NI=ni,BM=bm,BN=bn,NS=ks,
num_warps=nw,waves_per_eu=0,num_stages=1)
return fp4,ssh
# Dispatch: preshuffle for small M, ASM for large M
_USE_PRESHUFFLE = {4, 8, 16, 32}
_OUT_PS = {}
_OUT_ASM = {}
def custom_kernel(data: input_t) -> output_t:
a = data[0]
m, k = a.shape
n = data[3].shape[0]
if m in _USE_PRESHUFFLE:
# Zero-copy preshuffle path (single fused launch)
key = (m, n)
if key not in _OUT_PS:
_OUT_PS[key] = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
w = data[3].view(torch.uint8).reshape(n // 16, (k // 2) * 16)
sc = data[4].view(torch.uint8)
n_pad = sc.shape[0]
w_sc = sc.reshape(n_pad // 32, -1)[:n // 32, :k]
return gemm_a16wfp4_preshuffle(a, w, w_sc, prequant=True, y=_OUT_PS[key])
else:
# ASM path (fused quant+shuffle + ASM GEMM)
if (m, n) not in _OUT_ASM:
_OUT_ASM[(m, n)] = torch.empty((((m+31)//32)*32, n), dtype=torch.bfloat16, device=a.device)
o = _OUT_ASM[(m, n)]
aq, ash = _qa(a)
aiter.gemm_a4w4_asm(aq.view(_F).view(m, k//2), data[3], ash.view(_E), data[4],
o, _K64 if m <= 8 else _K32, bpreshuffle=True, log2_k_split=0)
return o[:m]
scrolls · 120 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