submission 739530
vuxml · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 556 lines, June 9 Researcher Reciprocity License v1.0.
submission_v5_pq.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-739530?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:fabf5e6fa4127851dcf33fbb753f2397f9e58c838963ac580f40363f87b3d1d2
license declaredunknown
license concludedunknown
authorsvuxml
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
v5: 2-kernel HIP path — tiny prequant + lean fp4-GEMM.num-warps = 4
if SK>1:_reduce_k[(rg,)](W,C,SK,m,n,m*n,n,n,BLK=256,SKC=8,num_warps=4)shared-memory
extern __shared__ float red[];split-k
KERNEL 2 (fgemm_pq): v1's SPLITK/SPLITN but A from Afp4 (16B) + Asc (1B).stages = 2
num_warps=nw,num_stages=2,matrix_instr_nonkdim=16,waves_per_eu=0)tile-k = 512
BK=512 if k>=512 else 256;out=[]vector-width = int4
void hw_quant32(const int4* __restrict__ a4,i32x4& o,int& e8){Kernel source
submission_v5_pq.py556 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v5: 2-kernel HIP path — tiny prequant + lean fp4-GEMM.
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).
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.
ALSO: v1's fused-quant SPLITK (proven winner m≤32) retained as arm.
"""
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};
}
// ═══════ 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){
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;
}
// ═══════ 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>
__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)
{
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 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 unroll
for(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];
#pragma unroll
for(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;
}
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);
const long kf=(k>>1)+kg*16, ks=(k>>5)+kg;
#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);
}
}
if constexpr(MODE==1){
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];
}
}
// ═══════ KERNEL 2b: fused-quant SPLITK (EXACT v5a — proven 10.0us@m=64) ═══
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 MD>
static void _gpq(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;
static bool _s=false;
if(!_s&&lds>65536){hipFuncSetAttribute((const void*)fgemm_pq<W,MR,MD>,
hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
fgemm_pq<W,MR,MD><<<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>
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>,
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_pq(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);
#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;
#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);
#undef D
return -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);
#undef P
}
"""
_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,
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,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="v5d_pq", cpp_sources=_CPP,
cuda_sources=_HIP_SRC, functions=["go_pq","launch_pq","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()
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]}")
# ═════ Triton fallback (v10h) ═════
@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 (proven, MR<=2 only — MR>2 serial-quants too much)
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):
if W>K128:continue
for MR in(16,8,4,2,1):
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
out.append(("pq",W,MR,1,MT,NT))
return out
def _tri_cfgs(m,n,k):
BK=512 if k>=512 else 256;out=[]
if m<=32:
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))
else:
for BM in(32,64)if m>=64 else(32,):
for nw in(4,8):out.append((BM,32,BK,1,nw,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((8,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-9
rg=triton.cdiv(m*n,256)
def _rh(cfg,_A,_Bq,_Bsh,_Bsc):
kind,Wv,MR,MD,MT,_=cfg
if kind=="fq":
rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,MR,MD)
else:
_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)
if rc!=0:raise RuntimeError(f"dispatch miss {cfg}")
return C
def _rt(cfg,_A,_Bq,_Bsh,_Bsc):
BM,BN,BK,SK,nw,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=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)
return C
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=[]
for tag,cfg,run in cand:
if time.time()-t0>35:_L(" [budget]");break
try:
C.fill_(float('nan')) # poison — catches partial-write dispatch bugs
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%}")
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 best is None:_L(f" [{tag}]{cfg}:EXC{type(e).__name__}:{str(e)[:100]}")
torch.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()
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
except 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}")
_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 · 556 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 712312.
- #!POPCORN leaderboard amd-mxfp4-mm- #!POPCORN gpu MI355X- """- v10h: clean, rules-compliant version of v10d.-- ARCHITECTURE:- One custom Triton kernel `_gemm_k` with two modes:- - FUSED (m≤32): load bf16 A-tile → in-register MXFP4 quant via- aiter's _mxfp4_quant_op → tl.dot_scaled vs B_q. 1 kernel launch.- - PREQUANT (m≥64): tiny quant kernel writes Afp4/Asc; GEMM reads- fp4 A. Amortizes quant cost across N-tiles.- + split-K (workspace + reduce) for thin grids (m≤16, k≥2048).-- KEY TECHNIQUES:- 1. In-register quant fused into GEMM K-loop (no separate quant launch- or intermediate HBM buffer for small-M).- 2. B_scale_sh read DIRECTLY from its e8m0_shuffle layout via the- closed-form forward index → no unshuffle preprocessing.- 3. L2-cold autotune (matches eval.py's clear_l2_cache semantics).- 4. Per-(m,n,k) config cache; preallocated C/W/Afp4/Asc.-- NOTE: eval.py measures GPU-event time AFTER a 16GB L2-flush, so Python- launch overhead (~15µs) is fully overlapped and never measured → plain- `kernel[grid](...)` is optimal; no low-level launch tricks needed.- """- import os, sys- 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)--- # e8m0_shuffle forward flat idx (from aiter/utility/fp4_utils.py):- # view(sm//32,2,16, sn//8,2,4).permute(0,3,5,2,4,1).reshape(sm,sn)- # Separable: flat = row_part(r) + col_part(c)- @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, sA_m, sAsc_m, sBq_n, sC_k, sC_m, sn8,- BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,- SPLIT_K: tl.constexpr, EVEN_N: tl.constexpr, PREQUANT: tl.constexpr,- ):- pid = tl.program_id(0)- num_n = tl.cdiv(N, BN)- num_mn = tl.cdiv(M, BM) * num_n- pid_k = pid // num_mn- pid_mn = pid % num_mn- pid_m = pid_mn // num_n- pid_n = pid_mn % num_n-- offs_m = pid_m * BM + tl.arange(0, BM)- offs_n = pid_n * BN + tl.arange(0, BN)- offs_n64 = offs_n.to(tl.int64)- mask_m = offs_m < M- mask_n = offs_n < N- rk = tl.arange(0, BK)- rk2 = tl.arange(0, BK // 2)- rk32 = tl.arange(0, BK // 32)-- k_per = tl.cdiv(tl.cdiv(K, BK), SPLIT_K) * BK- k_lo = pid_k * k_per- k_hi = min(k_lo + k_per, K)-- bq_ptrs = Bq + offs_n64[:, None] * sBq_n + (k_lo // 2 + rk2)[None, :]- bsc_row = _sh_row(offs_n64, sn8)- acc = tl.zeros((BM, BN), dtype=tl.float32)-- if PREQUANT:- a_ptrs = A + offs_m[:, None].to(tl.int64) * sA_m \- + (k_lo // 2 + rk2)[None, :]- asc_ptrs = Asc + offs_m[:, None].to(tl.int64) * sAsc_m \- + (k_lo // 32 + rk32)[None, :]- else:- a_ptrs = A + offs_m[:, None].to(tl.int64) * sA_m \- + (k_lo + rk)[None, :]-- for k in tl.range(k_lo, k_hi, BK):- if PREQUANT:- a_fp4 = tl.load(a_ptrs, mask=mask_m[:, None], other=0)- a_sc = tl.load(asc_ptrs, mask=mask_m[:, None], other=0)- a_ptrs += BK // 2; asc_ptrs += BK // 32- else:- a_bf = tl.load(a_ptrs, mask=mask_m[:, None], other=0.0)- a_fp4, a_sc = _mxfp4_quant_op(a_bf.to(tl.float32), BK, BM, 32)- a_ptrs += BK-- if EVEN_N:- b_fp4_t = tl.load(bq_ptrs)- b_sc = tl.load(Bsc + bsc_row[:, None]- + _sh_col(k // 32 + rk32)[None, :])- else:- b_fp4_t = tl.load(bq_ptrs, mask=mask_n[:, None], other=0)- b_sc = tl.load(Bsc + bsc_row[:, None]- + _sh_col(k // 32 + rk32)[None, :],- mask=mask_n[:, None], other=0)- acc = tl.dot_scaled(a_fp4, a_sc, "e2m1",- tl.trans(b_fp4_t), b_sc, "e2m1", acc)- bq_ptrs += BK // 2-- c_off = (pid_k * sC_k + offs_m[:, None].to(tl.int64) * sC_m- + offs_n[None, :])- cmask = mask_m[:, None] & mask_n[None, :]- if SPLIT_K == 1:- tl.store(C + c_off, acc.to(tl.bfloat16), mask=cmask)- else:- tl.store(C + c_off, acc, mask=cmask)--- @triton.jit- def _reduce_k(W, C, SK, M, N, sW_k, sW_m, sC_m,- BLK: tl.constexpr, SKC: tl.constexpr):- pid = tl.program_id(0)- off = pid * BLK + tl.arange(0, BLK)- om = off // N; on = off % N; mask = om < M- base = om.to(tl.int64) * sW_m + on- s = tl.zeros((BLK,), dtype=tl.float32)- for i in tl.static_range(SKC):- s += tl.load(W + i * sW_k + base, mask=mask & (i < SK), other=0.0)- tl.store(C + om.to(tl.int64) * sC_m + on, s.to(tl.bfloat16), mask=mask)--- @triton.jit- def _quant_a_k(A, Afp4, Asc, M, K, sA_m, sAf_m, sAs_m,- BM: tl.constexpr, BK: tl.constexpr):- pid = tl.program_id(0)- nk = tl.cdiv(K, BK)- pm = pid // nk; pk = pid % nk- offs_m = pm * BM + tl.arange(0, BM)- offs_k = pk * BK + tl.arange(0, BK)- mask_m = offs_m < M- a = tl.load(A + offs_m[:, None].to(tl.int64) * sA_m + offs_k[None, :],- mask=mask_m[:, None], other=0.0).to(tl.float32)- af, asc = _mxfp4_quant_op(a, BK, BM, 32)- tl.store(Afp4 + offs_m[:, None].to(tl.int64) * sAf_m- + (pk * (BK // 2) + tl.arange(0, BK // 2))[None, :],- af, mask=mask_m[:, None])- tl.store(Asc + offs_m[:, None].to(tl.int64) * sAs_m- + (pk * (BK // 32) + tl.arange(0, BK // 32))[None, :],- asc, mask=mask_m[:, None])--- def _cfgs(m, n, k):- """(BM, BN, BK, SPLIT_K, num_warps, nonK, num_stages, PREQUANT)"""- out = []- PQs = (False,) if m <= 32 else (True,) if m >= 128 else (False, True)- for PQ in PQs:- BMs = (16,) if not PQ else tuple(b for b in (16, 32, 64, 128) if b <= m)- for BM in BMs:- for BN in (32, 64, 128, 256):- if BN > n: continue- bt = -(-m // BM) * -(-n // BN)- for BK in (256, 512):- if BK > k: continue- nkit = k // BK- SKs = [1]- if bt < 200 and nkit >= 2:- tgt = max(1, 256 // bt)- for s in (2, 4, 8):- if s <= nkit and s <= tgt * 2: SKs.append(s)- for SK in SKs:- for nw in (4, 8):- for nK in ((16,) if BM == 16 else (16, 32)):- out.append((BM, BN, BK, SK, nw, nK, 2, PQ))- seen, r = set(), []- for c in out:- if c not in seen: seen.add(c); r.append(c)- return r--- _L2FLUSH = torch.empty(512 * 1024 * 1024, dtype=torch.int8, device="cuda")--- def _gpu_time_cold(fn, n_iter=6):- for _ in range(2): fn()- torch.cuda.synchronize()- evs = [(torch.cuda.Event(True), torch.cuda.Event(True)) for _ in range(n_iter)]- for e0, e1 in evs:- _L2FLUSH.zero_()- e0.record(); fn(); e1.record()- torch.cuda.synchronize()- ts = sorted(e0.elapsed_time(e1) for e0, e1 in evs)- return sum(ts[:n_iter - 1]) * 1000.0 / (n_iter - 1)--- def _ref(A, B_shuffle, B_scale_sh):- Aq, As = dynamic_mxfp4_quant(A)- return aiter.gemm_a4w4(Aq.view(dtypes.fp4x2), B_shuffle,- e8m0_shuffle(As).view(dtypes.fp8_e8m0),- B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)--- _STATE: dict = {}--- def _build(data):- A, B, B_q, B_shuffle, B_scale_sh = data- m, k = A.shape; n = B.shape[0]- sn = B_scale_sh.shape[1]; sn8 = sn // 8- dev = A.device-- Bq = B_q.contiguous().view(torch.uint8)- Bsc = B_scale_sh.contiguous().view(torch.uint8)- C_bf = torch.empty((m, n), dtype=torch.bfloat16, device=dev)- W = torch.zeros((8, m, n), dtype=torch.float32, device=dev)- Afp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=dev)- Asc = torch.empty((m, k // 32), dtype=torch.uint8, device=dev)-- ref_f = _ref(A, B_shuffle, B_scale_sh).float()- mag = ref_f.abs().mean().item() + 1e-9-- QBM, QBK = 16, min(256, k)- q_gx = triton.cdiv(m, QBM) * (k // QBK)- r_gx = triton.cdiv(m * n, 256)-- def _do_quant(A_in):- _quant_a_k[(q_gx,)](A_in, Afp4, Asc, m, k, k, k // 2, k // 32,- BM=QBM, BK=QBK, num_warps=4)-- def _do_reduce(SK):- _reduce_k[(r_gx,)](W, C_bf, SK, m, n, m * n, n, n,- BLK=256, SKC=8, num_warps=4)-- _do_quant(A); _do_reduce(1); torch.cuda.synchronize()-- cfgs = _cfgs(m, n, k)- _L(f"\n[v10h m={m} n={n} k={k}] {len(cfgs)} cfgs")-- best, best_t, best_go = None, float("inf"), None- for cfg in cfgs:- BM, BN, BK, SK, nw, nK, ns, PQ = cfg- C_out = W if SK > 1 else C_bf- sC_k = W.stride(0) if SK > 1 else 0- sC_m = W.stride(1) if SK > 1 else n- sAm = (k // 2) if PQ else k- even_n = (n % BN == 0)- gx = triton.cdiv(m, BM) * triton.cdiv(n, BN) * SK-- def _go(_A=A, _Bq=Bq, _Bsc=Bsc, _cfg=cfg, _gx=gx, _C=C_out,- _sAm=sAm, _sCk=sC_k, _sCm=sC_m, _even=even_n):- BM, BN, BK, SK, nw, nK, ns, PQ = _cfg- if PQ:- _do_quant(_A)- a_src = Afp4- else:- a_src = _A- _gemm_k[(_gx,)](- a_src, Asc, _Bq, _Bsc, _C, m, n, k,- _sAm, k // 32, k // 2, _sCk, _sCm, sn8,- BM=BM, BN=BN, BK=BK, SPLIT_K=SK, EVEN_N=_even,- PREQUANT=PQ, num_warps=nw, num_stages=ns,- matrix_instr_nonkdim=nK, waves_per_eu=0)- if SK > 1:- _do_reduce(SK)- return C_bf-- try:- out = _go()- torch.cuda.synchronize()- err = ((out.float() - ref_f).abs().mean() / mag).item()- if err > 5e-3:- if best is None: _L(f" {cfg}: ERR {err:.2%}")- continue- t = _gpu_time_cold(_go)- if t < best_t:- best_t, best, best_go = t, cfg, _go- _L(f" {cfg}: {t:.2f}us grid={gx} *")- except Exception as e:- if best is None:- _L(f" {cfg}: {type(e).__name__}: {str(e)[:100]}")-- if best is None:- _L(f" → fallback"); return {"hot": None}- _L(f" → best={best} @ {best_t:.2f}us")- return {"hot": best_go}--- def custom_kernel(data):- A = data[0]; Bq = data[2]- key = (A.shape[0], Bq.shape[0], A.shape[1])- S = _STATE.get(key)- if S is None:- S = _build(data); _STATE[key] = S- hot = S["hot"]- if hot is None:- return _ref(A, data[3], data[4])- return hot(A, Bq.view(torch.uint8), data[4].view(torch.uint8))+ #!POPCORN leaderboard amd-mxfp4-mm+ #!POPCORN gpu MI355X+ """+ v5: 2-kernel HIP path — tiny prequant + lean fp4-GEMM.++ 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).++ 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.++ ALSO: v1's fused-quant SPLITK (proven winner m≤32) retained as arm.+ """+ 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};+ }++ // ═══════ 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){+ 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;+ }++ // ═══════ 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>+ __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)+ {+ 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 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 unroll+ for(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];+ #pragma unroll+ for(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;+ }++ 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);+ const long kf=(k>>1)+kg*16, ks=(k>>5)+kg;+ #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);+ }+ }++ if constexpr(MODE==1){+ 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];+ }+ }++ // ═══════ KERNEL 2b: fused-quant SPLITK (EXACT v5a — proven 10.0us@m=64) ═══+ 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 MD>+ static void _gpq(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;+ static bool _s=false;+ if(!_s&&lds>65536){hipFuncSetAttribute((const void*)fgemm_pq<W,MR,MD>,+ hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}+ fgemm_pq<W,MR,MD><<<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>+ 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>,+ 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_pq(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);+ #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;+ #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);+ #undef D+ return -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);+ #undef P+ }+ """++ _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,+ 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,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="v5d_pq", cpp_sources=_CPP,+ cuda_sources=_HIP_SRC, functions=["go_pq","launch_pq","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()+ 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]}")+++ # ═════ Triton fallback (v10h) ═════+ @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 (proven, MR<=2 only — MR>2 serial-quants too much)+ 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):+ if W>K128:continue+ for MR in(16,8,4,2,1):+ 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+ out.append(("pq",W,MR,1,MT,NT))+ return out+++ def _tri_cfgs(m,n,k):+ BK=512 if k>=512 else 256;out=[]+ if m<=32:+ 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))+ else:+ for BM in(32,64)if m>=64 else(32,):+ for nw in(4,8):out.append((BM,32,BK,1,nw,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((8,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-9+ rg=triton.cdiv(m*n,256)++ def _rh(cfg,_A,_Bq,_Bsh,_Bsc):+ kind,Wv,MR,MD,MT,_=cfg+ if kind=="fq":+ rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT,NT,Wv,MR,MD)+ else:+ _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)+ if rc!=0:raise RuntimeError(f"dispatch miss {cfg}")+ return C+ def _rt(cfg,_A,_Bq,_Bsh,_Bsc):+ BM,BN,BK,SK,nw,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=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)+ return C+ 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=[]+ for tag,cfg,run in cand:+ if time.time()-t0>35:_L(" [budget]");break+ try:+ C.fill_(float('nan')) # poison — catches partial-write dispatch bugs+ 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%}")+ 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 best is None:_L(f" [{tag}]{cfg}:EXC{type(e).__name__}:{str(e)[:100]}")+ torch.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()+ 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+ except 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}")+ _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 · 860 diff lines total
Best evidence level for this revision: reported
JSON