submission 740264
vuxml · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 847 lines, June 9 Researcher Reciprocity License v1.0.
submission_v11_ku.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-740264?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:adb39b4cf7a4dd32c283038f394d74615f2a70abf42133bafcd1f7a0e3d115a4
license declaredunknown
license concludedunknown
authorsvuxml
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
if SK>1:_reduce_k[(rg,)](W,C,SK,m,n,m*n,n,n,BLK=256,SKC=16,num_warps=4)shared-memory
extern __shared__ float red[];split-k
v11: K-unrolled split-K + 1-launch fqmn + expanded Triton.tile-k = 128
ROOT CAUSE (v9/v10 analysis): HIP loops at BK=128 (16 K-iters @ k=2048,vector-width = int4
void hw_quant32(const int4* __restrict__ a4,i32x4& o,int& e8){Kernel source
submission_v11_ku.py847 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v11: K-unrolled split-K + 1-launch fqmn + expanded Triton.
ROOT CAUSE (v9/v10 analysis): HIP loops at BK=128 (16 K-iters @ k=2048,
`#pragma unroll 1`). Triton at BK=512 (4 iters). sched_barrier can't fix
trip-count. Need BK=512 in HIP too.
NEW KERNELS:
fgemm_ku<W,MR,KU>: split-K, KU-unrolled K-body. K-iters/(W*KU) outer loops.
KU=4 @ k=2048,W=4 -> 4 outer iters (= Triton). KU mfma back-to-back hide
latency without 2-stage bufs (keep VGPR low). LDS reduce (proven cheap
at MR<=4).
fgemm_fqmn<W,MR,NR>: 1-launch fused-quant, MR m-tiles x NR n-tiles.
For m=64 (MR=4, NR=2): quant once, 8 mfma. 1-LAUNCH FLOOR = 5.9us.
prequant_z: prequant + zero Cf in same grid (for ak path, 3->launch).
TRITON: +SK*PQ combos at m<=32; +BK=1024,ns=3 at m>=64.
KEPT: fqn<4,2> (small-M), fq (safety), full Triton fallback.
"""
import os, sys, time
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
import warnings; warnings.filterwarnings("ignore")
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
_L = lambda *a: print(*a, file=sys.stderr, flush=True)
_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
typedef int i32x4 __attribute__((ext_vector_type(4)));
typedef int i32x8 __attribute__((ext_vector_type(8)));
typedef float f32x4 __attribute__((ext_vector_type(4)));
typedef __hip_bfloat16 bf16;
__device__ __forceinline__ uint32_t f2u(float x){
union{float f;uint32_t u;}c;c.f=x;return c.u;}
__device__ __forceinline__ float e8f(uint8_t e){
union{uint32_t u;float f;}c;c.u=(uint32_t)e<<23;return c.f;}
#define QCV(o,a,b,s,bs) __builtin_amdgcn_cvt_scalef32_pk_fp4_f32((o),(a),(b),(s),(bs))
__device__ __forceinline__ long bsc_idx(long r,long c,long sn8){
return (r>>5)*(sn8*256)+(r&15)*4+((r>>4)&1)
+(c>>3)*256+(c&3)*64+((c>>2)&1)*2;
}
__device__ __forceinline__ i32x8 w8(i32x4 x){
i32x8 r={0,0,0,0,0,0,0,0};r[0]=x[0];r[1]=x[1];r[2]=x[2];r[3]=x[3];return r;}
__device__ __forceinline__
void hw_quant32(const int4* __restrict__ a4,i32x4& o,int& e8){
bf16 ab[32] __attribute__((aligned(16)));
*reinterpret_cast<int4*>(&ab[ 0])=a4[0];
*reinterpret_cast<int4*>(&ab[ 8])=a4[1];
*reinterpret_cast<int4*>(&ab[16])=a4[2];
*reinterpret_cast<int4*>(&ab[24])=a4[3];
float v[32];float amax=0.f;
#pragma unroll
for(int i=0;i<32;++i){v[i]=(float)ab[i];
float t=__builtin_fabsf(v[i]);amax=t>amax?t:amax;}
uint32_t au=(f2u(amax)+0x200000u)&0xFF800000u;
int su=au?(int)((au>>23)&0xFFu)-129:-127;
su=su<-127?-127:(su>127?127:su);e8=su+127;
float bsc=e8f((uint8_t)e8);
int w0=0,w1=0,w2=0,w3=0;
w0=QCV(w0,v[ 0],v[ 1],bsc,0);w0=QCV(w0,v[ 2],v[ 3],bsc,1);
w0=QCV(w0,v[ 4],v[ 5],bsc,2);w0=QCV(w0,v[ 6],v[ 7],bsc,3);
w1=QCV(w1,v[ 8],v[ 9],bsc,0);w1=QCV(w1,v[10],v[11],bsc,1);
w1=QCV(w1,v[12],v[13],bsc,2);w1=QCV(w1,v[14],v[15],bsc,3);
w2=QCV(w2,v[16],v[17],bsc,0);w2=QCV(w2,v[18],v[19],bsc,1);
w2=QCV(w2,v[20],v[21],bsc,2);w2=QCV(w2,v[22],v[23],bsc,3);
w3=QCV(w3,v[24],v[25],bsc,0);w3=QCV(w3,v[26],v[27],bsc,1);
w3=QCV(w3,v[28],v[29],bsc,2);w3=QCV(w3,v[30],v[31],bsc,3);
o=(i32x4){w0,w1,w2,w3};
}
__global__ __launch_bounds__(256)
void prequant(const bf16* __restrict__ A,uint8_t* __restrict__ Af,
uint8_t* __restrict__ As,int M,int K){
int g=blockIdx.x*256+threadIdx.x;
int K32=K>>5;
if(g>=M*K32)return;
int m=g/K32, kb=g%K32;
const bf16* Ap=A+(long)m*K+(long)kb*32;
int4 ai[4];
ai[0]=*reinterpret_cast<const int4*>(Ap);
ai[1]=*reinterpret_cast<const int4*>(Ap+8);
ai[2]=*reinterpret_cast<const int4*>(Ap+16);
ai[3]=*reinterpret_cast<const int4*>(Ap+24);
i32x4 o; int e8;
hw_quant32(ai,o,e8);
*reinterpret_cast<i32x4*>(Af+(long)m*(K>>1)+(long)kb*16)=o;
As[(long)m*K32+kb]=(uint8_t)e8;
}
// ───── fgemm_ku: split-K, KU-unrolled body (BK_eff = KU*128) ─────
// WAVES split K at KU*128 granularity. Each wave: outer loop × KU mfma.
// Compiler sees KU independent load->mfma chains per body, can batch-issue.
template<int WAVES,int M_REP,int KU>
__global__ __launch_bounds__(WAVES*64)
void fgemm_ku(
const uint8_t* __restrict__ Af,const uint8_t* __restrict__ As,
const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,
bf16* __restrict__ C,int M,int N,int K,long sn8,int NT)
{
const int tid=threadIdx.x,L=tid&63,w=tid>>6;
const int m16=L&15,kg=L>>4;
const int bid=blockIdx.x;
const int m_tile=bid/NT, n_tile=bid%NT;
const long K128=K>>7;
const long ksz=((K128+(long)WAVES*KU-1)/((long)WAVES*KU))*KU;
const long i_lo=(long)w*ksz, i_hi=min(i_lo+ksz,K128);
const long Kh=K>>1,K32=K>>5;
const long n_col=(long)n_tile*16+m16;
const uint8_t* Bsh_t=Bsh+(long)n_tile*(long)K*8;
f32x4 acc[M_REP];
#pragma unroll
for(int r=0;r<M_REP;++r)acc[r]=(f32x4){0,0,0,0};
long mrow[M_REP]; int mmsk[M_REP];
#pragma unroll
for(int r=0;r<M_REP;++r){
int mr=(m_tile*M_REP+r)*16+m16;
mmsk[r]=(mr<M)?0xFF:0;mrow[r]=(mr<M)?mr:0;
}
// Outer loop steps KU*128; body fully unrolls KU k-substeps.
// Hoist all KU*(1+MR) global loads BEFORE all KU*MR mfma so compiler
// can issue them together (s_waitcnt before each batch decreases).
for(long ib=i_lo; ib<i_hi; ib+=KU){
i32x4 bb[KU]; int bsv[KU];
i32x4 ab[KU][M_REP]; int asv[KU][M_REP];
#pragma unroll
for(int u=0;u<KU;++u){
long k=(ib+u)*128;
bb[u]=*reinterpret_cast<const i32x4*>(Bsh_t+(k>>5)*256+L*16);
bsv[u]=(int)Bsc[bsc_idx(n_col,(k>>5)+kg,sn8)];
long kf=(k>>1)+(long)kg*16, ks=(k>>5)+kg;
#pragma unroll
for(int r=0;r<M_REP;++r){
ab[u][r]=*reinterpret_cast<const i32x4*>(Af+mrow[r]*Kh+kf);
asv[u][r]=(int)As[mrow[r]*K32+ks]&mmsk[r];
}
}
__builtin_amdgcn_sched_barrier(0);
#pragma unroll
for(int u=0;u<KU;++u){
i32x8 b8=w8(bb[u]);
#pragma unroll
for(int r=0;r<M_REP;++r)
acc[r]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
w8(ab[u][r]),b8,acc[r],4,4,0,asv[u][r],0,bsv[u]);
}
}
extern __shared__ float red[];
#pragma unroll
for(int r=0;r<M_REP;++r)
#pragma unroll
for(int i=0;i<4;++i)red[((long)w*M_REP+r)*256+L*4+i]=acc[r][i];
__syncthreads();
if(w!=0)return;
#pragma unroll
for(int r=0;r<M_REP;++r)
#pragma unroll
for(int i=0;i<4;++i){
float s=0;
#pragma unroll
for(int ww=0;ww<WAVES;++ww)s+=red[((long)ww*M_REP+r)*256+L*4+i];
acc[r][i]=s;
}
#pragma unroll
for(int r=0;r<M_REP;++r)
#pragma unroll
for(int i=0;i<4;++i){
int mo=(m_tile*M_REP+r)*16+kg*4+i;
if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[r][i];
}
}
// ───── fgemm_fqmn: 1-launch fused-quant, MR x NR ─────
// Per K-step: NR× B-load + MR× (A-load+quant) -> MR*NR mfma.
// Quant amortized NR× (quant once per m-tile, use for all NR n-tiles).
template<int WAVES,int M_REP,int N_REP>
__global__ __launch_bounds__(WAVES*64)
void fgemm_fqmn(
const bf16* __restrict__ A,
const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,
bf16* __restrict__ C,int M,int N,int K,long sn8,int NT)
{
const int tid=threadIdx.x,L=tid&63,w=tid>>6;
const int m16=L&15,kg=L>>4;
const int bid=blockIdx.x;
const int NTG=(NT+N_REP-1)/N_REP;
const int m_tile=bid/NTG, ntg=bid%NTG;
long ksz=((K/128+WAVES-1)/WAVES)*128;
long k_lo=(long)w*ksz,k_hi=min(k_lo+ksz,(long)K);
long n_col[N_REP]; int vnm[N_REP]; const uint8_t* Bsh_t[N_REP];
#pragma unroll
for(int nr=0;nr<N_REP;++nr){
int nt=ntg*N_REP+nr; int v=(nt<NT);
vnm[nr]=v?0xFF:0;
long ntr=v?nt:0;
n_col[nr]=ntr*16+m16;
Bsh_t[nr]=Bsh+ntr*(long)K*8;
}
long mrow[M_REP]; int vmm[M_REP];
#pragma unroll
for(int mr=0;mr<M_REP;++mr){
int m_r=(m_tile*M_REP+mr)*16+m16;
vmm[mr]=(m_r<M)?1:0;
mrow[mr]=(m_r<M)?m_r:0;
}
f32x4 acc[M_REP][N_REP];
#pragma unroll
for(int mr=0;mr<M_REP;++mr)
#pragma unroll
for(int nr=0;nr<N_REP;++nr) acc[mr][nr]=(f32x4){0,0,0,0};
for(long k=k_lo;k<k_hi;k+=128){
long kb_=(k>>5)*256+L*16, ks=(k>>5)+kg;
i32x4 bb[N_REP]; int bsv[N_REP];
#pragma unroll
for(int nr=0;nr<N_REP;++nr){
bb[nr]=*reinterpret_cast<const i32x4*>(Bsh_t[nr]+kb_);
bsv[nr]=(int)Bsc[bsc_idx(n_col[nr],ks,sn8)] & vnm[nr];
}
const long kba=k+(long)kg*32;
#pragma unroll
for(int mr=0;mr<M_REP;++mr){
const bf16* Ap=A+mrow[mr]*K+kba;
int4 ai[4];
ai[0]=*reinterpret_cast<const int4*>(Ap);
ai[1]=*reinterpret_cast<const int4*>(Ap+8);
ai[2]=*reinterpret_cast<const int4*>(Ap+16);
ai[3]=*reinterpret_cast<const int4*>(Ap+24);
i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);
if(!vmm[mr])a_sc=0;
i32x8 a8=w8(a4);
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
acc[mr][nr]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a8,w8(bb[nr]),acc[mr][nr],4,4,0,a_sc,0,bsv[nr]);
}
}
extern __shared__ float red[];
#pragma unroll
for(int mr=0;mr<M_REP;++mr)
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
#pragma unroll
for(int i=0;i<4;++i)
red[(((long)w*M_REP+mr)*N_REP+nr)*256+L*4+i]=acc[mr][nr][i];
__syncthreads();
if(w!=0)return;
#pragma unroll
for(int mr=0;mr<M_REP;++mr)
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
#pragma unroll
for(int i=0;i<4;++i){
float s=0;
#pragma unroll
for(int ww=0;ww<WAVES;++ww)
s+=red[(((long)ww*M_REP+mr)*N_REP+nr)*256+L*4+i];
acc[mr][nr][i]=s;
}
#pragma unroll
for(int mr=0;mr<M_REP;++mr){
#pragma unroll
for(int i=0;i<4;++i){
int mo=(m_tile*M_REP+mr)*16+kg*4+i;
if(mo>=M)continue;
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
if(vnm[nr]) C[(long)mo*N+n_col[nr]]=(bf16)acc[mr][nr][i];
}
}
}
// ───── fgemm_fqn: v9 (proven small-M winner — unchanged) ─────
template<int WAVES,int N_REP>
__global__ __launch_bounds__(WAVES*64)
void fgemm_fqn(
const bf16* __restrict__ A,
const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,
bf16* __restrict__ C,int M,int N,int K,long sn8,int NT)
{
const int tid=threadIdx.x,L=tid&63,w=tid>>6;
const int m16=L&15,kg=L>>4;
const int bid=blockIdx.x;
const int NTG=(NT+N_REP-1)/N_REP;
const int m_tile=bid/NTG, ntg=bid%NTG;
long ksz=((K/128+WAVES-1)/WAVES)*128;
long k_lo=(long)w*ksz,k_hi=min(k_lo+ksz,(long)K);
long n_col[N_REP]; int vnm[N_REP]; const uint8_t* Bsh_t[N_REP];
#pragma unroll
for(int nr=0;nr<N_REP;++nr){
int nt=ntg*N_REP+nr; int v=(nt<NT);
vnm[nr]=v?0xFF:0;
long ntr=v?nt:0;
n_col[nr]=ntr*16+m16;
Bsh_t[nr]=Bsh+ntr*(long)K*8;
}
const int m_row=m_tile*16+m16;
const bool vm=m_row<M;
const long mrow=vm?m_row:0;
f32x4 acc[N_REP];
#pragma unroll
for(int nr=0;nr<N_REP;++nr) acc[nr]=(f32x4){0,0,0,0};
for(long k=k_lo;k<k_hi;k+=128){
long kb_=(k>>5)*256+L*16, ks=(k>>5)+kg;
i32x4 bb[N_REP]; int bsv[N_REP];
#pragma unroll
for(int nr=0;nr<N_REP;++nr){
bb[nr]=*reinterpret_cast<const i32x4*>(Bsh_t[nr]+kb_);
bsv[nr]=(int)Bsc[bsc_idx(n_col[nr],ks,sn8)] & vnm[nr];
}
const long kba=k+(long)kg*32;
const bf16* Ap=A+mrow*K+kba;
int4 ai[4];
ai[0]=*reinterpret_cast<const int4*>(Ap);
ai[1]=*reinterpret_cast<const int4*>(Ap+8);
ai[2]=*reinterpret_cast<const int4*>(Ap+16);
ai[3]=*reinterpret_cast<const int4*>(Ap+24);
i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);
if(!vm)a_sc=0;
i32x8 a8=w8(a4);
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
acc[nr]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a8,w8(bb[nr]),acc[nr],4,4,0,a_sc,0,bsv[nr]);
}
extern __shared__ float red[];
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
#pragma unroll
for(int i=0;i<4;++i)red[((long)w*N_REP+nr)*256+L*4+i]=acc[nr][i];
__syncthreads();
if(w!=0)return;
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
#pragma unroll
for(int i=0;i<4;++i){
float s=0;
#pragma unroll
for(int ww=0;ww<WAVES;++ww)s+=red[((long)ww*N_REP+nr)*256+L*4+i];
acc[nr][i]=s;
}
#pragma unroll
for(int i=0;i<4;++i){
int mo=m_tile*16+kg*4+i;
if(mo>=M)continue;
#pragma unroll
for(int nr=0;nr<N_REP;++nr)
if(vnm[nr]) C[(long)mo*N+n_col[nr]]=(bf16)acc[nr][i];
}
}
// ───── fgemm_fq: v5d (safety) ─────
template<int WAVES,int M_REP>
__global__ __launch_bounds__(WAVES*64)
void fgemm_fq(
const bf16* __restrict__ A,
const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,
bf16* __restrict__ C,int M,int N,int K,long sn8,int NT)
{
const int tid=threadIdx.x,L=tid&63,w=tid>>6;
const int m16=L&15,kg=L>>4;
const int bid=blockIdx.x;
int m_tile=bid/NT, n_tile=bid%NT;
long ksz=((K/128+WAVES-1)/WAVES)*128;
long k_lo=(long)w*ksz,k_hi=min(k_lo+ksz,(long)K);
const bool vn=n_tile<NT;
const long n_col=(long)n_tile*16+m16;
const uint8_t* Bsh_t=Bsh+(long)n_tile*(long)K*8;
f32x4 acc[M_REP];
#pragma unroll
for(int r=0;r<M_REP;++r)acc[r]=(f32x4){0,0,0,0};
for(long k=k_lo;k<k_hi;k+=128){
i32x4 b4={0,0,0,0};int b_sc=0;
if(vn){
b4=*reinterpret_cast<const i32x4*>(Bsh_t+(k>>5)*256+L*16);
b_sc=(int)Bsc[bsc_idx(n_col,(k>>5)+kg,sn8)];
}
i32x8 b8=w8(b4);
#pragma unroll
for(int r=0;r<M_REP;++r){
const int m_row=(m_tile*M_REP+r)*16+m16;
const bool vm=m_row<M;
const long kb=k+(long)kg*32;
const bf16* Ap=A+(long)(vm?m_row:0)*K+kb;
int4 ai[4];
ai[0]=*reinterpret_cast<const int4*>(Ap);
ai[1]=*reinterpret_cast<const int4*>(Ap+8);
ai[2]=*reinterpret_cast<const int4*>(Ap+16);
ai[3]=*reinterpret_cast<const int4*>(Ap+24);
i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);
if(!vm)a_sc=0;
acc[r]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
w8(a4),b8,acc[r],4,4,0,a_sc,0,b_sc);
}
}
extern __shared__ float red[];
#pragma unroll
for(int r=0;r<M_REP;++r)
#pragma unroll
for(int i=0;i<4;++i)red[((long)w*M_REP+r)*256+L*4+i]=acc[r][i];
__syncthreads();
if(w!=0)return;
#pragma unroll
for(int r=0;r<M_REP;++r)
#pragma unroll
for(int i=0;i<4;++i){
float s=0;
#pragma unroll
for(int ww=0;ww<WAVES;++ww)s+=red[((long)ww*M_REP+r)*256+L*4+i];
acc[r][i]=s;
}
if(!vn)return;
#pragma unroll
for(int r=0;r<M_REP;++r)
#pragma unroll
for(int i=0;i<4;++i){
int mo=(m_tile*M_REP+r)*16+kg*4+i;
if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[r][i];
}
}
#include <torch/extension.h>
void go_pq(torch::Tensor A,torch::Tensor Af,torch::Tensor As,
int64_t M,int64_t K){
int64_t g=(M*(K>>5)+255)/256;
prequant<<<dim3(g),dim3(256),0,0>>>(
reinterpret_cast<const bf16*>(A.data_ptr()),
Af.data_ptr<uint8_t>(),As.data_ptr<uint8_t>(),(int)M,(int)K);
}
template<int W,int MR,int KU>
static void _gku(torch::Tensor Af,torch::Tensor As,torch::Tensor Bsh,
torch::Tensor Bsc,torch::Tensor C,
int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t MT,int64_t NT){
int64_t gx=MT*NT,lds=(int64_t)W*MR*256*4;
static bool _s=false;
if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_ku<W,MR,KU>,
hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
fgemm_ku<W,MR,KU><<<dim3(gx),dim3(W*64),lds,0>>>(
Af.data_ptr<uint8_t>(),As.data_ptr<uint8_t>(),
Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
reinterpret_cast<bf16*>(C.data_ptr()),
(int)M,(int)N,(int)K,sn8,(int)NT);
}
template<int W,int MR,int NR>
static void _gfqmn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
int64_t MT,int64_t NT){
int64_t NTG=(NT+NR-1)/NR;
int64_t gx=MT*NTG,lds=(int64_t)W*MR*NR*256*4;
static bool _s=false;
if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_fqmn<W,MR,NR>,
hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
fgemm_fqmn<W,MR,NR><<<dim3(gx),dim3(W*64),lds,0>>>(
reinterpret_cast<const bf16*>(A.data_ptr()),
Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
reinterpret_cast<bf16*>(C.data_ptr()),
(int)M,(int)N,(int)K,sn8,(int)NT);
}
template<int W,int NR>
static void _gfqn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
int64_t MT,int64_t NT){
int64_t NTG=(NT+NR-1)/NR;
int64_t gx=MT*NTG,lds=(int64_t)W*NR*256*4;
static bool _s=false;
if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_fqn<W,NR>,
hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
fgemm_fqn<W,NR><<<dim3(gx),dim3(W*64),lds,0>>>(
reinterpret_cast<const bf16*>(A.data_ptr()),
Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
reinterpret_cast<bf16*>(C.data_ptr()),
(int)M,(int)N,(int)K,sn8,(int)NT);
}
template<int W,int MR>
static void _gfq(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
int64_t MT,int64_t NT){
int64_t gx=MT*NT,lds=(int64_t)W*MR*256*4;
static bool _s=false;
if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_fq<W,MR>,
hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
fgemm_fq<W,MR><<<dim3(gx),dim3(W*64),lds,0>>>(
reinterpret_cast<const bf16*>(A.data_ptr()),
Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
reinterpret_cast<bf16*>(C.data_ptr()),
(int)M,(int)N,(int)K,sn8,(int)NT);
}
int64_t launch_ku(torch::Tensor Af,torch::Tensor As,torch::Tensor Bsh,
torch::Tensor Bsc,torch::Tensor C,int64_t M,int64_t N,int64_t K,
int64_t sn8,int64_t MT,int64_t NT,int64_t W,int64_t MR,int64_t KU){
#define D(Ww,Rr,Uu) if(W==Ww&&MR==Rr&&KU==Uu){ \
_gku<Ww,Rr,Uu>(Af,As,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
D(2,1,2);D(2,1,4);D(2,1,7);D(2,2,2);D(2,2,4);D(2,4,2);D(2,4,4);
D(4,1,2);D(4,1,4);D(4,1,7);D(4,2,2);D(4,2,4);D(4,4,2);D(4,4,4);
D(4,8,2);D(4,16,2);D(4,16,4);
D(8,1,2);D(8,1,4);D(8,1,7);D(8,2,2);D(8,2,4);D(8,4,2);
D(16,1,2);D(16,1,4);
#undef D
return -1;
}
int64_t launch_fqmn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
int64_t MT,int64_t NT,int64_t W,int64_t MR,int64_t NR){
#define D(Ww,Rr,Nn) if(W==Ww&&MR==Rr&&NR==Nn){ \
_gfqmn<Ww,Rr,Nn>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
D(2,2,2);D(2,4,2);D(2,4,4);
D(4,2,2);D(4,2,4);D(4,4,2);D(4,4,4);
D(8,2,2);D(8,4,2);
#undef D
return -1;
}
int64_t launch_fqn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
int64_t MT,int64_t NT,int64_t W,int64_t NR){
#define D(Ww,Nn) if(W==Ww&&NR==Nn){ \
_gfqn<Ww,Nn>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
D(2,2);D(4,1);D(4,2);D(4,3);D(4,4);D(8,1);D(8,2);D(8,4);
#undef D
return -1;
}
int64_t launch_fq(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
int64_t MT,int64_t NT,int64_t W,int64_t MR){
#define D(Ww,Rr) if(W==Ww&&MR==Rr){ \
_gfq<Ww,Rr>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
D(2,1);D(4,1);D(4,2);D(8,1);D(16,1);
#undef D
return -1;
}
void probe(){
hipFuncAttributes a;
#define P(k,s) (void)hipFuncGetAttributes(&a,(const void*)k); \
printf("[v11] %-24s VGPR=%3d spill=%zu\n",s,a.numRegs,a.localSizeBytes);
P((fgemm_ku<4,1,4>),"ku<4,1,4>");
P((fgemm_ku<4,2,4>),"ku<4,2,4>");
P((fgemm_ku<4,4,4>),"ku<4,4,4>");
P((fgemm_ku<4,4,2>),"ku<4,4,2>");
P((fgemm_ku<8,1,4>),"ku<8,1,4>");
P((fgemm_ku<8,1,7>),"ku<8,1,7>");
P((fgemm_ku<4,16,2>),"ku<4,16,2>");
P((fgemm_fqmn<4,4,2>),"fqmn<4,4,2>");
P((fgemm_fqmn<4,4,4>),"fqmn<4,4,4>");
P((fgemm_fqmn<2,4,2>),"fqmn<2,4,2>");
P((fgemm_fqmn<4,2,4>),"fqmn<4,2,4>");
P((fgemm_fqn<4,2>),"fqn<4,2>");
#undef P
}
"""
_CPP = r"""
#include <torch/extension.h>
void go_pq(torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t);
int64_t launch_ku(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
torch::Tensor,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,
int64_t,int64_t,int64_t);
int64_t launch_fqmn(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
int64_t launch_fqn(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
int64_t launch_fq(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
void probe();
"""
_hip = None
try:
from torch.utils.cpp_extension import load_inline
_t0 = time.time()
_hip = load_inline(name="v11_ku", cpp_sources=_CPP,
cuda_sources=_HIP_SRC,
functions=["go_pq","launch_ku","launch_fqmn","launch_fqn","launch_fq","probe"],
with_cuda=True,
extra_cuda_cflags=["-O3","--offload-arch=gfx950","-ffast-math"],
verbose=False)
_L(f"[v11] HIP compiled {time.time()-_t0:.1f}s"); _hip.probe()
except Exception as ex:
import traceback
_L(f"[v11] HIP FAIL: {type(ex).__name__}: {str(ex)[:2000]}")
for ln in traceback.format_exc().splitlines()[-25:]:
_L(f" {ln[:200]}")
@triton.jit
def _sh_row(r,sn8):return (r//32)*(sn8*256)+(r%16)*4+(r//16)%2
@triton.jit
def _sh_col(c):return (c//8)*256+(c%4)*64+(c//4)%2*2
@triton.jit
def _gemm_k(A,Asc,Bq,Bsc,C,M,N,K,sAm,sAcm,sBn,sCk,sCm,sn8,
BM:tl.constexpr,BN:tl.constexpr,BK:tl.constexpr,
SK:tl.constexpr,EN:tl.constexpr,PQ:tl.constexpr):
pid=tl.program_id(0);nn=tl.cdiv(N,BN);nmn=tl.cdiv(M,BM)*nn
pk=pid//nmn;pmn=pid%nmn;pm=pmn//nn;pn=pmn%nn
om=pm*BM+tl.arange(0,BM);on=pn*BN+tl.arange(0,BN)
o64=on.to(tl.int64);mm=om<M;mn=on<N
rk=tl.arange(0,BK);r2=tl.arange(0,BK//2);r32=tl.arange(0,BK//32)
kp=tl.cdiv(tl.cdiv(K,BK),SK)*BK;kl=pk*kp;kh=min(kl+kp,K)
bp=Bq+o64[:,None]*sBn+(kl//2+r2)[None,:]
br=_sh_row(o64,sn8);acc=tl.zeros((BM,BN),dtype=tl.float32)
if PQ:
ap=A+om[:,None].to(tl.int64)*sAm+(kl//2+r2)[None,:]
asp=Asc+om[:,None].to(tl.int64)*sAcm+(kl//32+r32)[None,:]
else:
ap=A+om[:,None].to(tl.int64)*sAm+(kl+rk)[None,:]
for k in tl.range(kl,kh,BK):
if PQ:
af=tl.load(ap,mask=mm[:,None],other=0)
asc=tl.load(asp,mask=mm[:,None],other=0)
ap+=BK//2;asp+=BK//32
else:
ab=tl.load(ap,mask=mm[:,None],other=0.)
af,asc=_mxfp4_quant_op(ab.to(tl.float32),BK,BM,32);ap+=BK
if EN:
bf=tl.load(bp);bs=tl.load(Bsc+br[:,None]+_sh_col(k//32+r32)[None,:])
else:
bf=tl.load(bp,mask=mn[:,None],other=0)
bs=tl.load(Bsc+br[:,None]+_sh_col(k//32+r32)[None,:],mask=mn[:,None],other=0)
acc=tl.dot_scaled(af,asc,"e2m1",tl.trans(bf),bs,"e2m1",acc);bp+=BK//2
co=pk*sCk+om[:,None].to(tl.int64)*sCm+on[None,:];cm=mm[:,None]&mn[None,:]
if SK==1:tl.store(C+co,acc.to(tl.bfloat16),mask=cm)
else:tl.store(C+co,acc,mask=cm)
@triton.jit
def _reduce_k(W,C,SK,M,N,sWk,sWm,sCm,BLK:tl.constexpr,SKC:tl.constexpr):
p=tl.program_id(0);o=p*BLK+tl.arange(0,BLK)
om=o//N;on=o%N;m=om<M;b=om.to(tl.int64)*sWm+on
s=tl.zeros((BLK,),dtype=tl.float32)
for i in tl.static_range(SKC):s+=tl.load(W+i*sWk+b,mask=m&(i<SK),other=0.)
tl.store(C+om.to(tl.int64)*sCm+on,s.to(tl.bfloat16),mask=m)
def _hip_cfgs(m,n,k):
if _hip is None:return []
NT=-(-n//16);MT16=-(-m//16);K128=k//128;out=[]
# fq / fqn (1-launch proven)
for W in(4,8,2,16):
if W>K128:continue
out.append(("fq",W,1,0,MT16,NT))
if MT16<=2:
for W in(4,8,2):
if W>K128:continue
for NR in(1,2,3,4):
out.append(("fqn",W,1,NR,MT16,NT))
# fqmn (1-launch, MR>=2) — m>=32 only
if MT16>=2:
for W in(4,2,8):
if W>K128:continue
for MR in(2,4):
if MR>MT16:continue
for NR in(2,4):
MT=-(-MT16//MR);NTG=-(-NT//NR);gx=MT*NTG
lds=W*MR*NR*1024
if gx<8 or gx>8192 or lds>160*1024:continue
out.append(("fqmn",W,MR,NR,MT,NT))
# ku (2-launch, K-unrolled split-K) — m>=16
# KU chosen s.t. W*KU ~ K128 (1-2 outer iters) OR KU=4/2 for big K
for W in(4,8,2,16):
if W>K128:continue
for MR in(1,2,4,8,16):
if MR>MT16:continue
MT=-(-MT16//MR);gx=MT*NT;lds=W*MR*1024
if gx<16 or gx>8192 or lds>160*1024:continue
for KU in(7,4,2):
if W*KU>K128*2:continue # avoid mostly-empty waves
out.append(("ku",W,MR,KU,MT,NT))
return out
def _tri_cfgs(m,n,k):
out=[]
if m<=32:
BK=512 if k>=512 else 256
# baseline (v5d proven)
for BN in(32,64):
for nw in(4,8):out.append((16,BN,BK,1,nw,2,False))
# SK+PQ=False (m=16 winner)
if k>=2048:
for BK2 in(256,512):
for SK in(4,8,14):
if SK*BK2>k:continue
out.append((16,64,BK2,SK,4,2,False))
out.append((16,32,BK2,SK,4,2,False))
# SK+PQ=True (NEW: 3-launch but BK=512 possible)
for SK in(4,8):
out.append((16,64,512,SK,4,2,True))
out.append((16,32,512,SK,8,2,True))
else:
# m>=64: PQ=True baseline + expanded BK/ns
for BM in(32,64):
for BN in(32,64):
for BK in(512,1024):
if BK>k:continue
for nw in(4,8):
for ns in(2,3):
out.append((BM,BN,BK,1,nw,ns,True))
# SK+PQ at m>=64 (NEW)
if k>=2048:
for SK in(2,4):
out.append((64,32,512,SK,8,2,True))
return out
_L2=torch.empty(512*1024*1024,dtype=torch.int8,device="cuda")
def _tcold(fn,n=7):
for _ in range(2):fn()
torch.cuda.synchronize()
evs=[(torch.cuda.Event(True),torch.cuda.Event(True))for _ in range(n)]
for e0,e1 in evs:_L2.zero_();e0.record();fn();e1.record()
torch.cuda.synchronize()
ts=sorted(e0.elapsed_time(e1)for e0,e1 in evs)
return sum(ts[1:-1])*1000/(n-2)
def _ref(A,Bsh,Bsc):
Aq,As=dynamic_mxfp4_quant(A)
return aiter.gemm_a4w4(Aq.view(dtypes.fp4x2),Bsh,
e8m0_shuffle(As).view(dtypes.fp8_e8m0),Bsc,
dtype=dtypes.bf16,bpreshuffle=True)
_ST={}
def _build(data):
A,B,Bq_,Bsh_,Bsc_=data
m,k=A.shape;n=B.shape[0]
sn=Bsc_.shape[1];sn8=sn//8;dev=A.device;NT=-(-n//16)
Bq=Bq_.view(torch.uint8);Bsh=Bsh_.view(torch.uint8);Bsc=Bsc_.view(torch.uint8)
C=torch.empty((m,n),dtype=torch.bfloat16,device=dev)
W=torch.zeros((16,m,n),dtype=torch.float32,device=dev)
Af=torch.empty((m,k//2),dtype=torch.uint8,device=dev)
As=torch.empty((m,k//32),dtype=torch.uint8,device=dev)
rg=triton.cdiv(m*n,256)
def _rh(cfg,_A,_Bq,_Bsh,_Bsc):
kind,Wv,P2,P3,MT,_=cfg
if kind=="fq":
rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,1)
elif kind=="fqn":
rc=_hip.launch_fqn(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,P3)
elif kind=="fqmn":
rc=_hip.launch_fqmn(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,P2,P3)
elif kind=="ku":
_hip.go_pq(_A,Af,As,m,k)
rc=_hip.launch_ku(Af,As,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,P2,P3)
else:
rc=-1
if rc!=0:raise RuntimeError(f"dispatch miss {cfg}")
return C
def _rt(cfg,_A,_Bq,_Bsh,_Bsc):
BM,BN,BK,SK,nw,ns,PQ=cfg
gx=triton.cdiv(m,BM)*triton.cdiv(n,BN)*SK
Co=W if SK>1 else C;sCk=W.stride(0)if SK>1 else 0
sCm=W.stride(1)if SK>1 else n
if PQ:
_hip.go_pq(_A,Af,As,m,k);a,sAm=Af,k//2
else:a,sAm=_A,k
_gemm_k[(gx,)](a,As,_Bq,_Bsc,Co,m,n,k,sAm,k//32,k//2,sCk,sCm,sn8,
BM=BM,BN=BN,BK=BK,SK=SK,EN=(n%BN==0),PQ=PQ,
num_warps=nw,num_stages=ns,matrix_instr_nonkdim=16,waves_per_eu=0)
if SK>1:_reduce_k[(rg,)](W,C,SK,m,n,m*n,n,n,BLK=256,SKC=16,num_warps=4)
return C
_ref_cfg=(16,32,min(512,k),1,8,2,False)
rf=_rt(_ref_cfg,A,Bq,Bsh,Bsc).clone().float()
mag=rf.abs().mean().item()+1e-9
cand=[("hip",c,_rh)for c in _hip_cfgs(m,n,k)]
cand+=[("tri",c,_rt)for c in _tri_cfgs(m,n,k)]
_L(f"\n[v11 m={m} n={n} k={k}] {len(cand)}c hip={len(_hip_cfgs(m,n,k))}")
t0=time.time();best=None;bt=1e18;br=None;log=[];nerr=0
for tag,cfg,run in cand:
if time.time()-t0>50:_L(" [budget]");break
try:
C.fill_(float('nan'))
o=run(cfg,A,Bq,Bsh,Bsc);torch.cuda.synchronize()
err=((o.float()-rf).abs().mean()/mag).item()
if not (err<5e-3):
if nerr<4:_L(f" [{tag}]{cfg}:ERR{err:.2%}");nerr+=1
continue
t=_tcold(lambda c=cfg,r=run:r(c,A,Bq,Bsh,Bsc))
log.append((tag,cfg,t))
if t<bt:bt,best,br=t,(tag,cfg),run;_L(f" [{tag}]{cfg}:{t:.2f}us*")
except Exception as e:
if nerr<4:_L(f" [{tag}]{cfg}:EXC{type(e).__name__}:{str(e)[:100]}");nerr+=1
torch.cuda.synchronize()
if best is None:_L(" ->fb");return None
try:
A2=torch.randn_like(A)
rf2=_rt(_ref_cfg,A2,Bq,Bsh,Bsc).clone().float()
C.fill_(float('nan'))
o2=br(best[1],A2,Bq,Bsh,Bsc);torch.cuda.synchronize()
e2=((o2.float()-rf2).abs().mean()/(rf2.abs().mean()+1e-9)).item()
if not (e2<5e-3):
_L(f" RECHECK FAIL {best}: e2={e2:.2%} -> tri fallback")
best=("tri",_ref_cfg);br=_rt
except Exception as e:
_L(f" recheck exc {e}")
log.sort(key=lambda x:x[2])
for t,c,u in log[:12]:_L(f" top[{t}]{c}:{u:.2f}")
_L(f" ->best={best}@{bt:.2f}us")
return{"cfg":best[1],"run":br,"C":C,"Af":Af,"As":As}
def custom_kernel(data):
A=data[0];m,k=A.shape;n=data[2].shape[0]
S=_ST.get((m,n,k))
if S is None:
S=_build(data);_ST[(m,n,k)]=S if S is not None else False
if not S:return _ref(A,data[3],data[4])
return S["run"](S["cfg"],A,
data[2].view(torch.uint8),data[3].view(torch.uint8),
data[4].view(torch.uint8))
scrolls · 847 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 739530.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X"""- v5: 2-kernel HIP path — tiny prequant + lean fp4-GEMM.+ v11: K-unrolled split-K + 1-launch fqmn + expanded Triton.- RATIONALE: Floor analysis: event=4.2µs, +1.5µs/launch. 2 launches =- 7.2µs + work. At m≥64, hw_quant32 in-loop (80 serial ops × K/128 × M_REP)- is the wall. Separate prequant writes Afp4[M,K/2]+Asc[M,K/32] once (tiny),- GEMM reads 17B/lane/K-step (vs 64B+quant).+ ROOT CAUSE (v9/v10 analysis): HIP loops at BK=128 (16 K-iters @ k=2048,+ `#pragma unroll 1`). Triton at BK=512 (4 iters). sched_barrier can't fix+ trip-count. Need BK=512 in HIP too.- KERNEL 1 (prequant): grid=M*K/32 threads. Each does 1 quant-group.- KERNEL 2 (fgemm_pq): v1's SPLITK/SPLITN but A from Afp4 (16B) + Asc (1B).- Predicate-free (clamp m_row to 0, mask a_sc=0). M_REP=MT16 viable since- no quant VGPR pressure.+ NEW KERNELS:+ fgemm_ku<W,MR,KU>: split-K, KU-unrolled K-body. K-iters/(W*KU) outer loops.+ KU=4 @ k=2048,W=4 -> 4 outer iters (= Triton). KU mfma back-to-back hide+ latency without 2-stage bufs (keep VGPR low). LDS reduce (proven cheap+ at MR<=4).+ fgemm_fqmn<W,MR,NR>: 1-launch fused-quant, MR m-tiles x NR n-tiles.+ For m=64 (MR=4, NR=2): quant once, 8 mfma. 1-LAUNCH FLOOR = 5.9us.+ prequant_z: prequant + zero Cf in same grid (for ak path, 3->launch).- ALSO: v1's fused-quant SPLITK (proven winner m≤32) retained as arm.+ TRITON: +SK*PQ combos at m<=32; +BK=1024,ns=3 at m>=64.+ KEPT: fqn<4,2> (small-M), fq (safety), full Triton fallback."""import os, sys, timeos.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")⋯ 61 unchanged lineso=(i32x4){w0,w1,w2,w3};}- // ═══════ KERNEL 1: prequant A -> Afp4[M,K/2] + Asc[M,K/32] ═══════__global__ __launch_bounds__(256)void prequant(const bf16* __restrict__ A,uint8_t* __restrict__ Af,uint8_t* __restrict__ As,int M,int K){⋯ 13 unchanged linesAs[(long)m*K32+kb]=(uint8_t)e8;}- // ═══════ KERNEL 2a: GEMM with pre-quanted A (fp4) ═══════- // MODE 0=SPLITN (waves=n-tiles), 1=SPLITK (waves=K-slices).- template<int WAVES,int M_REP,int MODE>+ // ───── fgemm_ku: split-K, KU-unrolled body (BK_eff = KU*128) ─────+ // WAVES split K at KU*128 granularity. Each wave: outer loop × KU mfma.+ // Compiler sees KU independent load->mfma chains per body, can batch-issue.+ template<int WAVES,int M_REP,int KU>__global__ __launch_bounds__(WAVES*64)- void fgemm_pq(- const uint8_t* __restrict__ Af, // [M,K/2]- const uint8_t* __restrict__ As, // [M,K/32]- const uint8_t* __restrict__ Bsh,- const uint8_t* __restrict__ Bsc,- bf16* __restrict__ C,- int M,int N,int K,long sn8,int NT)+ void fgemm_ku(+ const uint8_t* __restrict__ Af,const uint8_t* __restrict__ As,+ const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,+ bf16* __restrict__ C,int M,int N,int K,long sn8,int NT){const int tid=threadIdx.x,L=tid&63,w=tid>>6;const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;- int m_tile,n_tile;long k_lo,k_hi;- if constexpr(MODE==0){- const int ntw=(NT+WAVES-1)/WAVES;- m_tile=bid/ntw; n_tile=(bid%ntw)*WAVES+w;- k_lo=0;k_hi=K;- } else {- m_tile=bid/NT; n_tile=bid%NT;- long ksz=((K/128+WAVES-1)/WAVES)*128;- k_lo=(long)w*ksz;k_hi=min(k_lo+ksz,(long)K);- }- const bool vn=n_tile<NT;+ const int m_tile=bid/NT, n_tile=bid%NT;+ const long K128=K>>7;+ const long ksz=((K128+(long)WAVES*KU-1)/((long)WAVES*KU))*KU;+ const long i_lo=(long)w*ksz, i_hi=min(i_lo+ksz,K128);+ const long Kh=K>>1,K32=K>>5;const long n_col=(long)n_tile*16+m16;const uint8_t* Bsh_t=Bsh+(long)n_tile*(long)K*8;- const long Kh=K>>1,K32=K>>5;f32x4 acc[M_REP];#pragma unrollfor(int r=0;r<M_REP;++r)acc[r]=(f32x4){0,0,0,0};- // Predicate-free m: clamp row to 0, mask scale to 0 (fp4 garbage * 2^-127 ~ 0)- int m_row[M_REP]; int m_msk[M_REP];+ long mrow[M_REP]; int mmsk[M_REP];#pragma unrollfor(int r=0;r<M_REP;++r){int mr=(m_tile*M_REP+r)*16+m16;- m_msk[r]=(mr<M)?0xFF:0;- m_row[r]=(mr<M)?mr:0;+ mmsk[r]=(mr<M)?0xFF:0;mrow[r]=(mr<M)?mr:0;}+ // Outer loop steps KU*128; body fully unrolls KU k-substeps.+ // Hoist all KU*(1+MR) global loads BEFORE all KU*MR mfma so compiler+ // can issue them together (s_waitcnt before each batch decreases).+ for(long ib=i_lo; ib<i_hi; ib+=KU){+ i32x4 bb[KU]; int bsv[KU];+ i32x4 ab[KU][M_REP]; int asv[KU][M_REP];+ #pragma unroll+ for(int u=0;u<KU;++u){+ long k=(ib+u)*128;+ bb[u]=*reinterpret_cast<const i32x4*>(Bsh_t+(k>>5)*256+L*16);+ bsv[u]=(int)Bsc[bsc_idx(n_col,(k>>5)+kg,sn8)];+ long kf=(k>>1)+(long)kg*16, ks=(k>>5)+kg;+ #pragma unroll+ for(int r=0;r<M_REP;++r){+ ab[u][r]=*reinterpret_cast<const i32x4*>(Af+mrow[r]*Kh+kf);+ asv[u][r]=(int)As[mrow[r]*K32+ks]&mmsk[r];+ }+ }+ __builtin_amdgcn_sched_barrier(0);+ #pragma unroll+ for(int u=0;u<KU;++u){+ i32x8 b8=w8(bb[u]);+ #pragma unroll+ for(int r=0;r<M_REP;++r)+ acc[r]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(+ w8(ab[u][r]),b8,acc[r],4,4,0,asv[u][r],0,bsv[u]);+ }+ }++ extern __shared__ float red[];+ #pragma unroll+ for(int r=0;r<M_REP;++r)+ #pragma unroll+ for(int i=0;i<4;++i)red[((long)w*M_REP+r)*256+L*4+i]=acc[r][i];+ __syncthreads();+ if(w!=0)return;+ #pragma unroll+ for(int r=0;r<M_REP;++r)+ #pragma unroll+ for(int i=0;i<4;++i){+ float s=0;+ #pragma unroll+ for(int ww=0;ww<WAVES;++ww)s+=red[((long)ww*M_REP+r)*256+L*4+i];+ acc[r][i]=s;+ }+ #pragma unroll+ for(int r=0;r<M_REP;++r)+ #pragma unroll+ for(int i=0;i<4;++i){+ int mo=(m_tile*M_REP+r)*16+kg*4+i;+ if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[r][i];+ }+ }++ // ───── fgemm_fqmn: 1-launch fused-quant, MR x NR ─────+ // Per K-step: NR× B-load + MR× (A-load+quant) -> MR*NR mfma.+ // Quant amortized NR× (quant once per m-tile, use for all NR n-tiles).+ template<int WAVES,int M_REP,int N_REP>+ __global__ __launch_bounds__(WAVES*64)+ void fgemm_fqmn(+ const bf16* __restrict__ A,+ const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,+ bf16* __restrict__ C,int M,int N,int K,long sn8,int NT)+ {+ const int tid=threadIdx.x,L=tid&63,w=tid>>6;+ const int m16=L&15,kg=L>>4;+ const int bid=blockIdx.x;+ const int NTG=(NT+N_REP-1)/N_REP;+ const int m_tile=bid/NTG, ntg=bid%NTG;+ long ksz=((K/128+WAVES-1)/WAVES)*128;+ long k_lo=(long)w*ksz,k_hi=min(k_lo+ksz,(long)K);++ long n_col[N_REP]; int vnm[N_REP]; const uint8_t* Bsh_t[N_REP];+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr){+ int nt=ntg*N_REP+nr; int v=(nt<NT);+ vnm[nr]=v?0xFF:0;+ long ntr=v?nt:0;+ n_col[nr]=ntr*16+m16;+ Bsh_t[nr]=Bsh+ntr*(long)K*8;+ }+ long mrow[M_REP]; int vmm[M_REP];+ #pragma unroll+ for(int mr=0;mr<M_REP;++mr){+ int m_r=(m_tile*M_REP+mr)*16+m16;+ vmm[mr]=(m_r<M)?1:0;+ mrow[mr]=(m_r<M)?m_r:0;+ }++ f32x4 acc[M_REP][N_REP];+ #pragma unroll+ for(int mr=0;mr<M_REP;++mr)+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr) acc[mr][nr]=(f32x4){0,0,0,0};+for(long k=k_lo;k<k_hi;k+=128){- i32x4 b4={0,0,0,0};int b_sc=0;- if(vn){- b4=*reinterpret_cast<const i32x4*>(Bsh_t+(k>>5)*256+L*16);- b_sc=(int)Bsc[bsc_idx(n_col,(k>>5)+kg,sn8)];+ long kb_=(k>>5)*256+L*16, ks=(k>>5)+kg;+ i32x4 bb[N_REP]; int bsv[N_REP];+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr){+ bb[nr]=*reinterpret_cast<const i32x4*>(Bsh_t[nr]+kb_);+ bsv[nr]=(int)Bsc[bsc_idx(n_col[nr],ks,sn8)] & vnm[nr];}- i32x8 b8=w8(b4);- const long kf=(k>>1)+kg*16, ks=(k>>5)+kg;+ const long kba=k+(long)kg*32;#pragma unroll- for(int r=0;r<M_REP;++r){- i32x4 a4=*reinterpret_cast<const i32x4*>(Af+(long)m_row[r]*Kh+kf);- int a_sc=(int)As[(long)m_row[r]*K32+ks] & m_msk[r];- acc[r]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(- w8(a4),b8,acc[r],4,4,0,a_sc,0,b_sc);+ for(int mr=0;mr<M_REP;++mr){+ const bf16* Ap=A+mrow[mr]*K+kba;+ int4 ai[4];+ ai[0]=*reinterpret_cast<const int4*>(Ap);+ ai[1]=*reinterpret_cast<const int4*>(Ap+8);+ ai[2]=*reinterpret_cast<const int4*>(Ap+16);+ ai[3]=*reinterpret_cast<const int4*>(Ap+24);+ i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);+ if(!vmm[mr])a_sc=0;+ i32x8 a8=w8(a4);+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr)+ acc[mr][nr]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(+ a8,w8(bb[nr]),acc[mr][nr],4,4,0,a_sc,0,bsv[nr]);}}- if constexpr(MODE==1){- extern __shared__ float red[];+ extern __shared__ float red[];+ #pragma unroll+ for(int mr=0;mr<M_REP;++mr)#pragma unroll- for(int r=0;r<M_REP;++r)+ for(int nr=0;nr<N_REP;++nr)#pragma unroll- for(int i=0;i<4;++i)red[((long)w*M_REP+r)*256+L*4+i]=acc[r][i];- __syncthreads();- if(w!=0)return;+ for(int i=0;i<4;++i)+ red[(((long)w*M_REP+mr)*N_REP+nr)*256+L*4+i]=acc[mr][nr][i];+ __syncthreads();+ if(w!=0)return;+ #pragma unroll+ for(int mr=0;mr<M_REP;++mr)#pragma unroll- for(int r=0;r<M_REP;++r)+ for(int nr=0;nr<N_REP;++nr)#pragma unrollfor(int i=0;i<4;++i){float s=0;#pragma unroll- for(int ww=0;ww<WAVES;++ww)s+=red[((long)ww*M_REP+r)*256+L*4+i];- acc[r][i]=s;+ for(int ww=0;ww<WAVES;++ww)+ s+=red[(((long)ww*M_REP+mr)*N_REP+nr)*256+L*4+i];+ acc[mr][nr][i]=s;}+ #pragma unroll+ for(int mr=0;mr<M_REP;++mr){+ #pragma unroll+ for(int i=0;i<4;++i){+ int mo=(m_tile*M_REP+mr)*16+kg*4+i;+ if(mo>=M)continue;+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr)+ if(vnm[nr]) C[(long)mo*N+n_col[nr]]=(bf16)acc[mr][nr][i];+ }}+ }- if(!vn)return;+ // ───── fgemm_fqn: v9 (proven small-M winner — unchanged) ─────+ template<int WAVES,int N_REP>+ __global__ __launch_bounds__(WAVES*64)+ void fgemm_fqn(+ const bf16* __restrict__ A,+ const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,+ bf16* __restrict__ C,int M,int N,int K,long sn8,int NT)+ {+ const int tid=threadIdx.x,L=tid&63,w=tid>>6;+ const int m16=L&15,kg=L>>4;+ const int bid=blockIdx.x;+ const int NTG=(NT+N_REP-1)/N_REP;+ const int m_tile=bid/NTG, ntg=bid%NTG;+ long ksz=((K/128+WAVES-1)/WAVES)*128;+ long k_lo=(long)w*ksz,k_hi=min(k_lo+ksz,(long)K);++ long n_col[N_REP]; int vnm[N_REP]; const uint8_t* Bsh_t[N_REP];#pragma unroll- for(int r=0;r<M_REP;++r)+ for(int nr=0;nr<N_REP;++nr){+ int nt=ntg*N_REP+nr; int v=(nt<NT);+ vnm[nr]=v?0xFF:0;+ long ntr=v?nt:0;+ n_col[nr]=ntr*16+m16;+ Bsh_t[nr]=Bsh+ntr*(long)K*8;+ }+ const int m_row=m_tile*16+m16;+ const bool vm=m_row<M;+ const long mrow=vm?m_row:0;++ f32x4 acc[N_REP];+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr) acc[nr]=(f32x4){0,0,0,0};++ for(long k=k_lo;k<k_hi;k+=128){+ long kb_=(k>>5)*256+L*16, ks=(k>>5)+kg;+ i32x4 bb[N_REP]; int bsv[N_REP];#pragma unroll+ for(int nr=0;nr<N_REP;++nr){+ bb[nr]=*reinterpret_cast<const i32x4*>(Bsh_t[nr]+kb_);+ bsv[nr]=(int)Bsc[bsc_idx(n_col[nr],ks,sn8)] & vnm[nr];+ }+ const long kba=k+(long)kg*32;+ const bf16* Ap=A+mrow*K+kba;+ int4 ai[4];+ ai[0]=*reinterpret_cast<const int4*>(Ap);+ ai[1]=*reinterpret_cast<const int4*>(Ap+8);+ ai[2]=*reinterpret_cast<const int4*>(Ap+16);+ ai[3]=*reinterpret_cast<const int4*>(Ap+24);+ i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);+ if(!vm)a_sc=0;+ i32x8 a8=w8(a4);+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr)+ acc[nr]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(+ a8,w8(bb[nr]),acc[nr],4,4,0,a_sc,0,bsv[nr]);+ }++ extern __shared__ float red[];+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr)+ #pragma unroll+ for(int i=0;i<4;++i)red[((long)w*N_REP+nr)*256+L*4+i]=acc[nr][i];+ __syncthreads();+ if(w!=0)return;+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr)+ #pragma unrollfor(int i=0;i<4;++i){- int mo=(m_tile*M_REP+r)*16+kg*4+i;- if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[r][i];+ float s=0;+ #pragma unroll+ for(int ww=0;ww<WAVES;++ww)s+=red[((long)ww*N_REP+nr)*256+L*4+i];+ acc[nr][i]=s;}+ #pragma unroll+ for(int i=0;i<4;++i){+ int mo=m_tile*16+kg*4+i;+ if(mo>=M)continue;+ #pragma unroll+ for(int nr=0;nr<N_REP;++nr)+ if(vnm[nr]) C[(long)mo*N+n_col[nr]]=(bf16)acc[nr][i];+ }}- // ═══════ KERNEL 2b: fused-quant SPLITK (EXACT v5a — proven 10.0us@m=64) ═══+ // ───── fgemm_fq: v5d (safety) ─────template<int WAVES,int M_REP>__global__ __launch_bounds__(WAVES*64)void fgemm_fq(- const bf16* __restrict__ A,- const uint8_t* __restrict__ Bsh,- const uint8_t* __restrict__ Bsc,- bf16* __restrict__ C,- int M,int N,int K,long sn8,int NT)+ const bf16* __restrict__ A,+ const uint8_t* __restrict__ Bsh,const uint8_t* __restrict__ Bsc,+ bf16* __restrict__ C,int M,int N,int K,long sn8,int NT){const int tid=threadIdx.x,L=tid&63,w=tid>>6;const int m16=L&15,kg=L>>4;⋯ 4 unchanged linesconst bool vn=n_tile<NT;const long n_col=(long)n_tile*16+m16;const uint8_t* Bsh_t=Bsh+(long)n_tile*(long)K*8;-f32x4 acc[M_REP];#pragma unrollfor(int r=0;r<M_REP;++r)acc[r]=(f32x4){0,0,0,0};-for(long k=k_lo;k<k_hi;k+=128){i32x4 b4={0,0,0,0};int b_sc=0;if(vn){⋯ 12 unchanged linesai[1]=*reinterpret_cast<const int4*>(Ap+8);ai[2]=*reinterpret_cast<const int4*>(Ap+16);ai[3]=*reinterpret_cast<const int4*>(Ap+24);- i32x4 a4;int a_sc;- hw_quant32(ai,a4,a_sc);+ i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);if(!vm)a_sc=0;acc[r]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(w8(a4),b8,acc[r],4,4,0,a_sc,0,b_sc);}}-extern __shared__ float red[];#pragma unrollfor(int r=0;r<M_REP;++r)⋯ 30 unchanged linesAf.data_ptr<uint8_t>(),As.data_ptr<uint8_t>(),(int)M,(int)K);}- template<int W,int MR,int MD>- static void _gpq(torch::Tensor Af,torch::Tensor As,torch::Tensor Bsh,+ template<int W,int MR,int KU>+ static void _gku(torch::Tensor Af,torch::Tensor As,torch::Tensor Bsh,torch::Tensor Bsc,torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t MT,int64_t NT){- int64_t gx=(MD==0)?MT*((NT+W-1)/W):MT*NT;- int64_t lds=(MD==1)?(int64_t)W*MR*256*4:0;+ int64_t gx=MT*NT,lds=(int64_t)W*MR*256*4;static bool _s=false;- if(!_s&&lds>65536){hipFuncSetAttribute((const void*)fgemm_pq<W,MR,MD>,+ if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_ku<W,MR,KU>,hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}- fgemm_pq<W,MR,MD><<<dim3(gx),dim3(W*64),lds,0>>>(+ fgemm_ku<W,MR,KU><<<dim3(gx),dim3(W*64),lds,0>>>(Af.data_ptr<uint8_t>(),As.data_ptr<uint8_t>(),Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),reinterpret_cast<bf16*>(C.data_ptr()),(int)M,(int)N,(int)K,sn8,(int)NT);}+ template<int W,int MR,int NR>+ static void _gfqmn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,+ torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,+ int64_t MT,int64_t NT){+ int64_t NTG=(NT+NR-1)/NR;+ int64_t gx=MT*NTG,lds=(int64_t)W*MR*NR*256*4;+ static bool _s=false;+ if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_fqmn<W,MR,NR>,+ hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}+ fgemm_fqmn<W,MR,NR><<<dim3(gx),dim3(W*64),lds,0>>>(+ reinterpret_cast<const bf16*>(A.data_ptr()),+ Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),+ reinterpret_cast<bf16*>(C.data_ptr()),+ (int)M,(int)N,(int)K,sn8,(int)NT);+ }++ template<int W,int NR>+ static void _gfqn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,+ torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,+ int64_t MT,int64_t NT){+ int64_t NTG=(NT+NR-1)/NR;+ int64_t gx=MT*NTG,lds=(int64_t)W*NR*256*4;+ static bool _s=false;+ if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_fqn<W,NR>,+ hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}+ fgemm_fqn<W,NR><<<dim3(gx),dim3(W*64),lds,0>>>(+ reinterpret_cast<const bf16*>(A.data_ptr()),+ Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),+ reinterpret_cast<bf16*>(C.data_ptr()),+ (int)M,(int)N,(int)K,sn8,(int)NT);+ }+template<int W,int MR>static void _gfq(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t MT,int64_t NT){int64_t gx=MT*NT,lds=(int64_t)W*MR*256*4;static bool _s=false;- if(!_s&&lds>65536){hipFuncSetAttribute((const void*)fgemm_fq<W,MR>,+ if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_fq<W,MR>,hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}fgemm_fq<W,MR><<<dim3(gx),dim3(W*64),lds,0>>>(reinterpret_cast<const bf16*>(A.data_ptr()),⋯ 2 unchanged lines(int)M,(int)N,(int)K,sn8,(int)NT);}- int64_t launch_pq(torch::Tensor Af,torch::Tensor As,torch::Tensor Bsh,+ int64_t launch_ku(torch::Tensor Af,torch::Tensor As,torch::Tensor Bsh,torch::Tensor Bsc,torch::Tensor C,int64_t M,int64_t N,int64_t K,- int64_t sn8,int64_t MT,int64_t NT,int64_t W,int64_t MR,int64_t MD){- #define D(Ww,Rr,Mm) if(W==Ww&&MR==Rr&&MD==Mm){ \- _gpq<Ww,Rr,Mm>(Af,As,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}- D(2,1,1);D(2,2,1);D(2,4,1);D(2,8,1);D(2,16,1);- D(4,1,1);D(4,2,1);D(4,4,1);D(4,8,1);D(4,16,1);- D(8,1,1);D(8,2,1);D(8,4,1);D(8,8,1);D(8,16,1);+ int64_t sn8,int64_t MT,int64_t NT,int64_t W,int64_t MR,int64_t KU){+ #define D(Ww,Rr,Uu) if(W==Ww&&MR==Rr&&KU==Uu){ \+ _gku<Ww,Rr,Uu>(Af,As,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}+ D(2,1,2);D(2,1,4);D(2,1,7);D(2,2,2);D(2,2,4);D(2,4,2);D(2,4,4);+ D(4,1,2);D(4,1,4);D(4,1,7);D(4,2,2);D(4,2,4);D(4,4,2);D(4,4,4);+ D(4,8,2);D(4,16,2);D(4,16,4);+ D(8,1,2);D(8,1,4);D(8,1,7);D(8,2,2);D(8,2,4);D(8,4,2);+ D(16,1,2);D(16,1,4);#undef Dreturn -1;}+ int64_t launch_fqmn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,+ torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,+ int64_t MT,int64_t NT,int64_t W,int64_t MR,int64_t NR){+ #define D(Ww,Rr,Nn) if(W==Ww&&MR==Rr&&NR==Nn){ \+ _gfqmn<Ww,Rr,Nn>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}+ D(2,2,2);D(2,4,2);D(2,4,4);+ D(4,2,2);D(4,2,4);D(4,4,2);D(4,4,4);+ D(8,2,2);D(8,4,2);+ #undef D+ return -1;+ }++ int64_t launch_fqn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,+ torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,+ int64_t MT,int64_t NT,int64_t W,int64_t NR){+ #define D(Ww,Nn) if(W==Ww&&NR==Nn){ \+ _gfqn<Ww,Nn>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}+ D(2,2);D(4,1);D(4,2);D(4,3);D(4,4);D(8,1);D(8,2);D(8,4);+ #undef D+ return -1;+ }+int64_t launch_fq(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,- int64_t MT,int64_t NT,int64_t W,int64_t MR,int64_t MD){- (void)MD;+ int64_t MT,int64_t NT,int64_t W,int64_t MR){#define D(Ww,Rr) if(W==Ww&&MR==Rr){ \_gfq<Ww,Rr>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}- D(2,1);D(2,2);D(4,1);D(4,2);D(8,1);D(8,2);D(16,1);+ D(2,1);D(4,1);D(4,2);D(8,1);D(16,1);#undef Dreturn -1;}void probe(){hipFuncAttributes a;- #define P(k,W,R,MD) hipFuncGetAttributes(&a,(const void*)k); \- printf("[v5] %s W=%d MR=%d VGPR=%d spill=%zu\n",#k,W,R,a.numRegs,a.localSizeBytes);- P((fgemm_pq<4,4,0>),4,4,0);P((fgemm_pq<4,16,0>),4,16,0);- P((fgemm_pq<4,4,1>),4,4,1);P((fgemm_pq<8,4,1>),8,4,1);- P((fgemm_fq<8,1>),8,1,0);P((fgemm_fq<4,2>),4,2,0);- P((fgemm_pq<8,16,1>),8,16,1);- P((prequant),0,0,0);+ #define P(k,s) (void)hipFuncGetAttributes(&a,(const void*)k); \+ printf("[v11] %-24s VGPR=%3d spill=%zu\n",s,a.numRegs,a.localSizeBytes);+ P((fgemm_ku<4,1,4>),"ku<4,1,4>");+ P((fgemm_ku<4,2,4>),"ku<4,2,4>");+ P((fgemm_ku<4,4,4>),"ku<4,4,4>");+ P((fgemm_ku<4,4,2>),"ku<4,4,2>");+ P((fgemm_ku<8,1,4>),"ku<8,1,4>");+ P((fgemm_ku<8,1,7>),"ku<8,1,7>");+ P((fgemm_ku<4,16,2>),"ku<4,16,2>");+ P((fgemm_fqmn<4,4,2>),"fqmn<4,4,2>");+ P((fgemm_fqmn<4,4,4>),"fqmn<4,4,4>");+ P((fgemm_fqmn<2,4,2>),"fqmn<2,4,2>");+ P((fgemm_fqmn<4,2,4>),"fqmn<4,2,4>");+ P((fgemm_fqn<4,2>),"fqn<4,2>");#undef P}"""⋯ 1 unchanged lines_CPP = r"""#include <torch/extension.h>void go_pq(torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t);- int64_t launch_pq(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,+ int64_t launch_ku(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);- int64_t launch_fq(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,+ int64_t launch_fqmn(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);+ int64_t launch_fqn(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,+ int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);+ int64_t launch_fq(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,+ int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);void probe();"""⋯ 1 unchanged linestry:from torch.utils.cpp_extension import load_inline_t0 = time.time()- _hip = load_inline(name="v5d_pq", cpp_sources=_CPP,- cuda_sources=_HIP_SRC, functions=["go_pq","launch_pq","launch_fq","probe"],+ _hip = load_inline(name="v11_ku", cpp_sources=_CPP,+ cuda_sources=_HIP_SRC,+ functions=["go_pq","launch_ku","launch_fqmn","launch_fqn","launch_fq","probe"],with_cuda=True,extra_cuda_cflags=["-O3","--offload-arch=gfx950","-ffast-math"],verbose=False)- _L(f"[v5] HIP compiled {time.time()-_t0:.1f}s"); _hip.probe()+ _L(f"[v11] HIP compiled {time.time()-_t0:.1f}s"); _hip.probe()except Exception as ex:import traceback- _L(f"[v5] HIP FAIL: {type(ex).__name__}: {str(ex)[:400]}")- for ln in traceback.format_exc().splitlines()[-20:]:- _L(f" {ln[:180]}")+ _L(f"[v11] HIP FAIL: {type(ex).__name__}: {str(ex)[:2000]}")+ for ln in traceback.format_exc().splitlines()[-25:]:+ _L(f" {ln[:200]}")- # ═════ Triton fallback (v10h) ═════@triton.jitdef _sh_row(r,sn8):return (r//32)*(sn8*256)+(r%16)*4+(r//16)%2@triton.jit⋯ 44 unchanged linesdef _hip_cfgs(m,n,k):if _hip is None:return []NT=-(-n//16);MT16=-(-m//16);K128=k//128;out=[]- # fq (proven, MR<=2 only — MR>2 serial-quants too much)+ # fq / fqn (1-launch proven)for W in(4,8,2,16):if W>K128:continue- for MR in(1,2):- if MR>MT16:continue- MT=-(-MT16//MR);gx=MT*NT- if gx<16 or gx>8192:continue- out.append(("fq",W,MR,0,MT,NT))- # pq (2-launch) — MR up to 16, all dispatch entries exist now- for W in(8,4,2):+ out.append(("fq",W,1,0,MT16,NT))+ if MT16<=2:+ for W in(4,8,2):+ if W>K128:continue+ for NR in(1,2,3,4):+ out.append(("fqn",W,1,NR,MT16,NT))+ # fqmn (1-launch, MR>=2) — m>=32 only+ if MT16>=2:+ for W in(4,2,8):+ if W>K128:continue+ for MR in(2,4):+ if MR>MT16:continue+ for NR in(2,4):+ MT=-(-MT16//MR);NTG=-(-NT//NR);gx=MT*NTG+ lds=W*MR*NR*1024+ if gx<8 or gx>8192 or lds>160*1024:continue+ out.append(("fqmn",W,MR,NR,MT,NT))+ # ku (2-launch, K-unrolled split-K) — m>=16+ # KU chosen s.t. W*KU ~ K128 (1-2 outer iters) OR KU=4/2 for big K+ for W in(4,8,2,16):if W>K128:continue- for MR in(16,8,4,2,1):+ for MR in(1,2,4,8,16):if MR>MT16:continueMT=-(-MT16//MR);gx=MT*NT;lds=W*MR*1024if gx<16 or gx>8192 or lds>160*1024:continue- out.append(("pq",W,MR,1,MT,NT))+ for KU in(7,4,2):+ if W*KU>K128*2:continue # avoid mostly-empty waves+ out.append(("ku",W,MR,KU,MT,NT))return outdef _tri_cfgs(m,n,k):- BK=512 if k>=512 else 256;out=[]+ out=[]if m<=32:+ BK=512 if k>=512 else 256+ # baseline (v5d proven)for BN in(32,64):- for nw in(4,8):out.append((16,BN,BK,1,nw,False))- if k>=2048:out.append((16,64,256,8,4,False))+ for nw in(4,8):out.append((16,BN,BK,1,nw,2,False))+ # SK+PQ=False (m=16 winner)+ if k>=2048:+ for BK2 in(256,512):+ for SK in(4,8,14):+ if SK*BK2>k:continue+ out.append((16,64,BK2,SK,4,2,False))+ out.append((16,32,BK2,SK,4,2,False))+ # SK+PQ=True (NEW: 3-launch but BK=512 possible)+ for SK in(4,8):+ out.append((16,64,512,SK,4,2,True))+ out.append((16,32,512,SK,8,2,True))else:- for BM in(32,64)if m>=64 else(32,):- for nw in(4,8):out.append((BM,32,BK,1,nw,True))+ # m>=64: PQ=True baseline + expanded BK/ns+ for BM in(32,64):+ for BN in(32,64):+ for BK in(512,1024):+ if BK>k:continue+ for nw in(4,8):+ for ns in(2,3):+ out.append((BM,BN,BK,1,nw,ns,True))+ # SK+PQ at m>=64 (NEW)+ if k>=2048:+ for SK in(2,4):+ out.append((64,32,512,SK,8,2,True))return out⋯ 22 unchanged linessn=Bsc_.shape[1];sn8=sn//8;dev=A.device;NT=-(-n//16)Bq=Bq_.view(torch.uint8);Bsh=Bsh_.view(torch.uint8);Bsc=Bsc_.view(torch.uint8)C=torch.empty((m,n),dtype=torch.bfloat16,device=dev)- W=torch.zeros((8,m,n),dtype=torch.float32,device=dev)+ W=torch.zeros((16,m,n),dtype=torch.float32,device=dev)Af=torch.empty((m,k//2),dtype=torch.uint8,device=dev)As=torch.empty((m,k//32),dtype=torch.uint8,device=dev)- rf=_ref(A,Bsh_,Bsc_).float();mag=rf.abs().mean().item()+1e-9rg=triton.cdiv(m*n,256)def _rh(cfg,_A,_Bq,_Bsh,_Bsc):- kind,Wv,MR,MD,MT,_=cfg+ kind,Wv,P2,P3,MT,_=cfgif kind=="fq":- rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,MR,MD)- else:+ rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,1)+ elif kind=="fqn":+ rc=_hip.launch_fqn(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,P3)+ elif kind=="fqmn":+ rc=_hip.launch_fqmn(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,P2,P3)+ elif kind=="ku":_hip.go_pq(_A,Af,As,m,k)- rc=_hip.launch_pq(Af,As,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,MR,MD)+ rc=_hip.launch_ku(Af,As,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,P2,P3)+ else:+ rc=-1if rc!=0:raise RuntimeError(f"dispatch miss {cfg}")return Cdef _rt(cfg,_A,_Bq,_Bsh,_Bsc):- BM,BN,BK,SK,nw,PQ=cfg+ BM,BN,BK,SK,nw,ns,PQ=cfggx=triton.cdiv(m,BM)*triton.cdiv(n,BN)*SKCo=W if SK>1 else C;sCk=W.stride(0)if SK>1 else 0sCm=W.stride(1)if SK>1 else n⋯ 2 unchanged lineselse:a,sAm=_A,k_gemm_k[(gx,)](a,As,_Bq,_Bsc,Co,m,n,k,sAm,k//32,k//2,sCk,sCm,sn8,BM=BM,BN=BN,BK=BK,SK=SK,EN=(n%BN==0),PQ=PQ,- num_warps=nw,num_stages=2,matrix_instr_nonkdim=16,waves_per_eu=0)- if SK>1:_reduce_k[(rg,)](W,C,SK,m,n,m*n,n,n,BLK=256,SKC=8,num_warps=4)+ num_warps=nw,num_stages=ns,matrix_instr_nonkdim=16,waves_per_eu=0)+ if SK>1:_reduce_k[(rg,)](W,C,SK,m,n,m*n,n,n,BLK=256,SKC=16,num_warps=4)return C++ _ref_cfg=(16,32,min(512,k),1,8,2,False)+ rf=_rt(_ref_cfg,A,Bq,Bsh,Bsc).clone().float()+ mag=rf.abs().mean().item()+1e-9+cand=[("hip",c,_rh)for c in _hip_cfgs(m,n,k)]cand+=[("tri",c,_rt)for c in _tri_cfgs(m,n,k)]- _L(f"\n[v5 m={m} n={n} k={k}] {len(cand)}c hip={len(_hip_cfgs(m,n,k))}")- t0=time.time();best=None;bt=1e18;br=None;log=[]+ _L(f"\n[v11 m={m} n={n} k={k}] {len(cand)}c hip={len(_hip_cfgs(m,n,k))}")+ t0=time.time();best=None;bt=1e18;br=None;log=[];nerr=0for tag,cfg,run in cand:- if time.time()-t0>35:_L(" [budget]");break+ if time.time()-t0>50:_L(" [budget]");breaktry:- C.fill_(float('nan')) # poison — catches partial-write dispatch bugs+ C.fill_(float('nan'))o=run(cfg,A,Bq,Bsh,Bsc);torch.cuda.synchronize()err=((o.float()-rf).abs().mean()/mag).item()if not (err<5e-3):- if len(log)<3:_L(f" [{tag}]{cfg}:ERR{err:.2%}")+ if nerr<4:_L(f" [{tag}]{cfg}:ERR{err:.2%}");nerr+=1continuet=_tcold(lambda c=cfg,r=run:r(c,A,Bq,Bsh,Bsc))log.append((tag,cfg,t))if t<bt:bt,best,br=t,(tag,cfg),run;_L(f" [{tag}]{cfg}:{t:.2f}us*")except Exception as e:- if best is None:_L(f" [{tag}]{cfg}:EXC{type(e).__name__}:{str(e)[:100]}")+ if nerr<4:_L(f" [{tag}]{cfg}:EXC{type(e).__name__}:{str(e)[:100]}");nerr+=1torch.cuda.synchronize()if best is None:_L(" ->fb");return None- # RECHECK-style self-test: run winner on FRESH random A (different data)try:A2=torch.randn_like(A)- rf2=_rt(list(_tri_cfgs(m,n,k))[0],A2,Bq,Bsh,Bsc).clone().float()+ rf2=_rt(_ref_cfg,A2,Bq,Bsh,Bsc).clone().float()C.fill_(float('nan'))o2=br(best[1],A2,Bq,Bsh,Bsc);torch.cuda.synchronize()e2=((o2.float()-rf2).abs().mean()/(rf2.abs().mean()+1e-9)).item()if not (e2<5e-3):- _L(f" RECHECK FAIL {best}: e2={e2:.2%} -> Triton fallback")- best=("tri",list(_tri_cfgs(m,n,k))[0]);br=_rt+ _L(f" RECHECK FAIL {best}: e2={e2:.2%} -> tri fallback")+ best=("tri",_ref_cfg);br=_rtexcept Exception as e:_L(f" recheck exc {e}")log.sort(key=lambda x:x[2])- for t,c,u in log[:6]:_L(f" top[{t}]{c}:{u:.2f}")+ for t,c,u in log[:12]:_L(f" top[{t}]{c}:{u:.2f}")_L(f" ->best={best}@{bt:.2f}us")return{"cfg":best[1],"run":br,"C":C,"Af":Af,"As":As}
scrolls · 795 diff lines total
Best evidence level for this revision: reported
JSON