Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
14.3µs
#497 of 1143
2026-04-05

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.

fp4M=16, K=7168, N=2112 │ fused_asm │ 20.9µs │ 21.9µs│ FP4 MFMA for large K
num-warps = 1num_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 = 1num_warps=1, waves_per_eu=0, num_stages=1)
tile-m = 4Tuned configs: (4,512,2880)→BM=4,BN=64 (32,512,4096)→BM=8,BN=128
tile-n = 64Tuned configs: (4,512,2880)→BM=4,BN=64 (32,512,4096)→BM=8,BN=128

Kernel 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