submission 733989
pawan2411 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 253 lines, June 9 Researcher Reciprocity License v1.0.
submission_v97_optimal.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-733989?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:4643e1264a26b742f7bf66ec128c1b2cc1f92fb1889b111275790c89109ab2d9
license declaredunknown
license concludedunknown
authorspawan2411
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
M=16, K=7168, N=2112 │ fused_asm │ 20.9µs │ 21.9µs│ FP4 MFMA for large Knum-warps = 1
num_warps=1, waves_per_eu=0, num_stages=1)split-k
'NUM_KSPLIT':1,'SPLITK_BLOCK_SIZE':2*Kp,'num_warps':nw,'num_stages':ns,stages = 1
num_warps=1, waves_per_eu=0, num_stages=1)tile-m = 4
Tuned configs: (4,512,2880)→BM=4,BN=64 (32,512,4096)→BM=8,BN=128tile-n = 64
Tuned configs: (4,512,2880)→BM=4,BN=64 (32,512,4096)→BM=8,BN=128Kernel source
submission_v97_optimal.py253 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
GEMM v93 — BEST SUBMISSION (Ranked geo mean: 14.5µs, benchmark: 13.6µs)
=== PROFILING (torch profiler on MI355X) ===
ALL GPU kernels complete in <1µs total.
The 13.6µs benchmark time is 93% CPU dispatch overhead.
=== THREE CODE PATHS (routed by M and K) ===
PATH 1: a16wfp4 (M<=32, K<=1024) — shapes M=4 and M=32 with K=512
┌─────────────────────────────────────────────────────────────────┐
│ Step │ GPU │ CPU dispatch │ Kernel │
│─────────────────────────│─────────│──────────────│──────────────│
│ 1. inverse_e8m0_shuffle │ 0.48µs │ ~2µs │ HIP custom │
│ 2. _gemm_a16wfp4_kernel │ 0.48µs │ ~3µs │ Triton MFMA │
│─────────────────────────│─────────│──────────────│──────────────│
│ TOTAL │ <1µs │ ~5µs │ 2 launches │
│ + Python overhead │ │ ~5µs │ │
│ BENCHMARK TIME │ │ ~10µs │ │
└─────────────────────────────────────────────────────────────────┘
Tuned configs: (4,512,2880)→BM=4,BN=64 (32,512,4096)→BM=8,BN=128
PATH 2: fused_asm (M<=32, K>1024) — shape M=16,K=7168
┌─────────────────────────────────────────────────────────────────┐
│ Step │ GPU │ CPU dispatch │ Kernel │
│─────────────────────────│─────────│──────────────│──────────────│
│ 1. _fqs (quant+shuffle) │ 0.48µs │ ~3µs │ Triton fused │
│ 2. gemm_a4w4_asm │ 0.48µs │ ~3µs │ CK ASM MFMA │
│─────────────────────────│─────────│──────────────│──────────────│
│ TOTAL │ <1µs │ ~6µs │ 2 launches │
│ + Python overhead │ │ ~15µs │ │
│ BENCHMARK TIME │ │ ~21µs │ │
└─────────────────────────────────────────────────────────────────┘
Uses custom fused Triton kernel that quantizes A + shuffles scales in one pass.
PATH 3: wrapper (M>32) — shapes M=64 and M=256
┌─────────────────────────────────────────────────────────────────┐
│ Step │ GPU │ CPU dispatch │ Kernel │
│─────────────────────────│─────────│──────────────│──────────────│
│ 1. _mxfp4_kernel (quant)│ 0.48µs │ ~3µs │ Triton quant │
│ 2. es_do (HIP shuffle) │ 0.48µs │ ~2µs │ HIP custom │
│ 3. gemm_a4w4 (wrapper) │ 0.48µs │ ~3µs │ CK ASM auto │
│─────────────────────────│─────────│──────────────│──────────────│
│ TOTAL │ ~1.5µs │ ~8µs │ 3 launches │
│ + Python overhead │ │ ~8µs │ │
│ BENCHMARK TIME │ │ ~17µs │ │
└─────────────────────────────────────────────────────────────────┘
Key: gemm_a4w4 WRAPPER auto-selects optimal tile (192x128 for M=256).
This beats v64's a16wfp4 by 4µs on M=256 (16.2 vs 20.6µs).
=== PER-SHAPE PERFORMANCE ===
Shape (M,K,N) │ Path │ Bench │ Ranked │ Why this path wins
───────────────────────│───────────│────────│────────│─────────────────────
M=4, K=512, N=2880 │ a16wfp4 │ 9.7µs │ 10.2µs│ No A quant needed
M=16, K=7168, N=2112 │ fused_asm │ 20.9µs │ 21.9µs│ FP4 MFMA for large K
M=32, K=512, N=4096 │ a16wfp4 │ 10.6µs │ 11.6µs│ No A quant needed
M=32, K=512, N=2880 │ a16wfp4 │ 11.3µs │ 12.4µs│ Marginal (ASM=11.0)
M=64, K=2048, N=7168 │ wrapper │ 17.5µs │ 18.2µs│ Auto-tuned tile size
M=256, K=1536, N=3072 │ wrapper │ 16.2µs │ 17.0µs│ Auto-tuned 192x128 tile
=== ROUTING LOGIC ===
if M <= 32 and K <= 1024: → PATH 1 (a16wfp4: bf16 A × fp4 B via mixed MFMA)
elif M <= 32: → PATH 2 (fused_asm: fused quant+shuffle + ASM)
else: → PATH 3 (wrapper: quant + shuffle + gemm_a4w4 auto)
=== VS THEORETICAL OPTIMAL ===
v93 geo mean: 13.6µs (benchmark), 14.5µs (ranked)
Optimal combo: 13.4µs (cherry-pick best per shape across ALL versions)
Gap: 0.2µs (within benchmark noise)
=== WHAT WE TRIED THAT DIDN'T WORK ===
Approach │ Result │ Why
──────────────────────│───────────│──────────────────────────────────
torch.mm(A, B.T) │ ❌ Failed │ bf16×bf16 ≠ fp4 reference
Custom HIP GEMM │ ❌ Failed │ bf16 mm ≠ a4w4 output
FlyDSL MLIR kernel │ ✅ <1µs GPU│ 60s compilation → timeout
FlyDSL HSACO extract │ ✅ Loaded │ Can't pack GPU kernel args
CUDA graphs │ ✅ Captured│ copy_() overhead ≥ dispatch savings
NUM_KSPLIT=2 │ ❌ Failed │ ATOMIC_ADD wrong results
gemm_afp4wfp4 │ ❌ Failed │ fp4x2 dtype not recognized
a16wfp4 for ALL shapes│ ✅ Correct │ 2x slower for K=7168 (39.9µs)
GROUP_SIZE_M=4 │ ✅ Correct │ Marginal help on K=7168 only
ctypes hipModuleLoad │ ✅ Loaded │ kernarg format mismatch
"""
import os
os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')
import torch
import triton
import triton.language as tl
import sys
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _gemm_a16wfp4_kernel, _get_config
_qm = sys.modules[dynamic_mxfp4_quant.__module__]
_mxfp4_kernel = _qm._dynamic_mxfp4_quant_kernel
# Fused quant+shuffle for ASM path (M<=32)
@triton.heuristics({"EM": lambda a: a["M"]%a["BSM"]==0 and a["N"]%(a["BSN"]*a["NI"])==0})
@triton.jit
def _fqs(xp,fp,sp,sxm,sxn,sfm,sfn,M,N,sn,BSM:tl.constexpr,BSN:tl.constexpr,NI:tl.constexpr,NS:tl.constexpr,QBS:tl.constexpr,EM:tl.constexpr,SM:tl.constexpr):
pm=tl.program_id(0);sn2=tl.program_id(1)*NI;xm2=tl.cast(sxm,tl.int64);xn2=tl.cast(sxn,tl.int64);fm2=tl.cast(sfm,tl.int64);fn2=tl.cast(sfn,tl.int64)
NQ:tl.constexpr=BSN//QBS;KB8=sn//8;s0=KB8*256
for pn in tl.range(sn2,min(sn2+NI,N),num_stages=NS):
xm=pm*BSM+tl.arange(0,BSM);xn=pn*BSN+tl.arange(0,BSN);xo=xm[:,None]*xm2+xn[None,:]*xn2
if EM:x=tl.load(xp+xo,cache_modifier=".cg").to(tl.float32)
else:x=tl.load(xp+xo,mask=(xm<M)[:,None]&(xn<N)[None,:],cache_modifier=".cg").to(tl.float32)
ot,bs=_mxfp4_quant_op(x,BSN,BSM,QBS)
om=pm*BSM+tl.arange(0,BSM);on=pn*BSN//2+tl.arange(0,BSN//2);oo=om[:,None]*fm2+on[None,:]*fn2
if EM:tl.store(fp+oo,ot)
else:tl.store(fp+oo,ot,mask=(om<M)[:,None]&(on<(N//2))[None,:])
ri=pm*BSM+tl.arange(0,BSM);ci=pn*NQ+tl.arange(0,NQ)
i0=ri//32;i1=(ri%32)//16;i2=ri%16;i3=ci//8;i4=(ci%8)//4;i5=ci%4
so=i0[:,None]*s0+i3[None,:]*256+i5[None,:]*64+i2[:,None]*4+i4[None,:]*2+i1[:,None]
if EM:tl.store(sp+so,bs)
else:tl.store(sp+so,bs,mask=(ri<M)[:,None]&(ci<((N+QBS-1)//QBS))[None,:])
# HIP kernels
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
__global__ void isk(const uint8_t*s,uint8_t*u,int N,int K,int sm,int sn){
int idx=blockIdx.x*blockDim.x+threadIdx.x;if(idx>=sm*sn)return;
int KB8=sn/8,s0=KB8*256,r=idx,i0=r/s0;r%=s0;int i3=r/256;r%=256;
int i5=r/64;r%=64;int i2=r/4;r%=4;int i4=r/2,i1=r%2;
if(i0*32+i1*16+i2<N&&i3*8+i4*4+i5<K)u[(i0*32+i1*16+i2)*K+i3*8+i4*4+i5]=s[idx];}
void is_into(torch::Tensor s,torch::Tensor u,int N,int K){
int sm=s.size(0),sn=s.size(1);
isk<<<((sm*sn)+255)/256,256>>>(s.data_ptr<uint8_t>(),u.data_ptr<uint8_t>(),N,K,sm,sn);}
__global__ void esk(const uint8_t*si,uint8_t*so,int M,int K,int sm,int sn,int sc){
int idx=blockIdx.x*blockDim.x+threadIdx.x;if(idx>=sm*sn)return;
int KB8=sn/8,s0=KB8*256,r=idx,o0=r/s0;r%=s0;int o1=r/256;r%=256;
int o2=r/64;r%=64;int o3=r/4;r%=4;int o4=r/2,o5=r%2;
int ir=o0*32+o5*16+o3,ic=o1*8+o4*4+o2;uint8_t v=0;
if(ir<M&&ic<K)v=si[ir+ic*sc];so[idx]=v;}
void es_do(torch::Tensor si,torch::Tensor so){
int M=si.size(0),K=si.size(1),sm=so.size(0),sn=so.size(1),sc=si.stride(1);
esk<<<((sm*sn)+255)/256,256>>>(si.data_ptr<uint8_t>(),so.data_ptr<uint8_t>(),M,K,sm,sn,sc);}
"""
_hip = load_inline(name='v93h',
cpp_sources=["void is_into(torch::Tensor,torch::Tensor,int,int);void es_do(torch::Tensor,torch::Tensor);"],
cuda_sources=[HIP_SRC], functions=['is_into','es_do'], verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950","-std=c++20","-O3"])
def _an(t):
fn=f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{t}";return f"_ZN5aiter{len(fn)}{fn}E"
_KN32=_an("32x128")
# v64's tuned configs for a16wfp4
_TC = {(4,512,2880):(4,64,4,1,2),(32,512,4096):(8,128,8,2,2),(32,512,2880):(8,128,4,2,2)}
_bufs = {}
def _get_bufs(M, K, N, device):
key = (M, K, N)
if key not in _bufs:
Kb = K // 32; Kp = K // 2
# Routing: a16wfp4 for M<=32 with K<=1024, ASM for everything else
use_a16 = (M <= 32 and K <= 1024)
if use_a16:
# v64 a16wfp4 path
if key in _TC: bm,bn,nw,ns,wpe = _TC[key]
else: bm,bn,nw,ns,wpe = 8,128,4,1,2
BK = max(64,min(512,triton.next_power_of_2(2*Kp)))
config = {'BLOCK_SIZE_M':bm,'BLOCK_SIZE_N':bn,'BLOCK_SIZE_K':BK,'GROUP_SIZE_M':1,
'NUM_KSPLIT':1,'SPLITK_BLOCK_SIZE':2*Kp,'num_warps':nw,'num_stages':ns,
'waves_per_eu':wpe,'matrix_instr_nonkdim':16,'cache_modifier':'.cg'}
grid = (triton.cdiv(M,bm)*triton.cdiv(N,bn),)
sn = ((Kb+7)//8)*8
_bufs[key] = {'path':'a16','config':config,'grid':grid,'Kp':Kp,'Kb':Kb,
'Bs':torch.empty(N,Kb,dtype=torch.uint8,device=device),'sn':sn,
'out':torch.empty(M,N,dtype=torch.bfloat16,device=device)}
elif M <= 32:
# Fused quant+shuffle + ASM (for M<=32, K>1024)
sm = ((M+255)//256)*256; sn = ((Kb+7)//8)*8
BSM = triton.next_power_of_2(M); BSN = 32
_bufs[key] = {'path':'fused_asm',
'xfp4':torch.empty(M,K//2,dtype=torch.uint8,device=device),
'shuf':torch.empty(sm,sn,dtype=torch.uint8,device=device),
'out':torch.empty(M,N,dtype=torch.bfloat16,device=device),
'sn':sn,'BSM':BSM,
'grid':(triton.cdiv(M,BSM),triton.cdiv(K,BSN))}
else:
# M>32: quant + HIP shuffle + gemm_a4w4 wrapper (637822 approach)
sm = ((M+255)//256)*256; sn = ((Kb+7)//8)*8
BSM = 32; BSN = 128; NI = 4; NW = 4; NSK = 2
_bufs[key] = {'path':'wrapper',
'xfp4':torch.empty(M,K//2,dtype=torch.uint8,device=device),
'bs':torch.empty((Kb,M),dtype=torch.uint8,device=device).T,
'shuf':torch.empty(sm,sn,dtype=torch.uint8,device=device),
'sn':sn,
'BLOCK_SIZE_M':BSM,'BLOCK_SIZE_N':BSN,'NUM_ITER':NI,'NUM_WARPS':NW,'NUM_STAGES':NSK,
'grid':(triton.cdiv(M,BSM),triton.cdiv(K,BSN*NI))}
return _bufs[key]
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape; N = B_q.shape[0]
b = _get_bufs(M, K, N, A.device)
if b['path'] == 'a16':
# a16wfp4: inverse shuffle + Triton GEMM
bs = B_scale_sh.view(torch.uint8)
if bs.dim() != 2: bs = bs.view(bs.numel()//b['sn'], b['sn'])
_hip.is_into(bs, b['Bs'], N, b['Kb'])
bq = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
_gemm_a16wfp4_kernel[b['grid']](A, bq, b['out'], b['Bs'],
M, N, b['Kp'], A.stride(0), A.stride(1), 1, b['Kp'],
0, b['out'].stride(0), b['out'].stride(1),
b['Bs'].stride(0), b['Bs'].stride(1),
ATOMIC_ADD=False, **b['config'])
return b['out']
elif b['path'] == 'fused_asm':
# Fused quant+shuffle + ASM
_fqs[b['grid']](A, b['xfp4'], b['shuf'],
A.stride(0), A.stride(1), b['xfp4'].stride(0), b['xfp4'].stride(1),
M=M, N=K, sn=b['sn'], BSM=b['BSM'], BSN=32, NI=1, NS=1, QBS=32, SM=0,
num_warps=1, waves_per_eu=0, num_stages=1)
aiter.gemm_a4w4_asm(b['xfp4'].view(dtypes.fp4x2), B_shuffle,
b['shuf'].view(dtypes.fp8_e8m0), B_scale_sh, b['out'],
_KN32, bpreshuffle=True)
return b['out']
else:
# M>32: quant + HIP shuffle + gemm_a4w4 wrapper
_mxfp4_kernel[b['grid']](A, b['xfp4'], b['bs'],
*A.stride(), *b['xfp4'].stride(), *b['bs'].stride(),
M=M, N=K, MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
NUM_ITER=b['NUM_ITER'], BLOCK_SIZE_M=b['BLOCK_SIZE_M'], BLOCK_SIZE_N=b['BLOCK_SIZE_N'],
NUM_STAGES=b['NUM_STAGES'], num_warps=b['NUM_WARPS'], waves_per_eu=0, num_stages=1)
A_q = b['xfp4'].view(dtypes.fp4x2)
_hip.es_do(b['bs'], b['shuf'])
A_scale_sh = b['shuf'].view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 253 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