Skip to content
KernelIndex
Search⌘K

submission 741643

vuxml · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 765 lines, June 9 Researcher Reciprocity License v1.0.

submission_v18_lb.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-741643?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.21µs
#44 of 1143
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a8ffe7da5d9db3ce06cdec1edac00a73b5c343271a8f9f8f46399d9cc7a48fa5
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 = 8num_warps=8,num_stages=2,matrix_instr_nonkdim=16)
shared-memoryextern __shared__ uint8_t _sh[];
stages = 2num_warps=8,num_stages=2,matrix_instr_nonkdim=16)
tile-m = 16BM=16,BN=32,BK=BK,EN=(n%32==0),
tile-n = 32BM=16,BN=32,BK=BK,EN=(n%32==0),
vector-width = int4void hw_quant32(const int4* __restrict__ a4,i32x4& o,int& e8){

Kernel source

submission_v18_lb.py765 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v18: HARDCODED v17 winners. ALL-HIP hot path. LB-ready.

v17 bench GM = 8.10us. All 6 shapes on my HIP kernels:
  fqn<4,2>      small-M 1-launch    6.15-6.28us
  alds_sk<8,8,7> m=16   2-launch    11.06us  (BEATS tri SK=14 @ 11.8)
  alds3<8,1,8>  m=64    1-launch    10.48us
  alds3<8,2,4>  m=256   1-launch    10.40us  (MR=2 WINS — 192WGs, B/=2)

v18 = v17 kernels + v16-style hardcode logic + Bq-fixed.
  alds_sk ping-pong Cf: cast zeros NEXT buffer → 2 launches not 3.
  Secret shapes: same pattern-match as v16 (tested 4/4 pass).
"""
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

_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__ 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};
}

// ═══════ fgemm_alds_sk: grid-level split-K (m=16 occ fix) ═══════
template<int WAVES,int KU>
__global__ __launch_bounds__(WAVES*64)
void fgemm_alds_sk(
    const bf16*    __restrict__ A,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,
    float*         __restrict__ Cf,
    int M,int N,int K,long sn8,int NT,int SK)
{
  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 NTW=(NT+WAVES-1)/WAVES;
  const int pk=bid/NTW, ntg=bid%NTW;
  const int n_tile=ntg*WAVES+w;
  const bool vn=(n_tile<NT);
  const int K32=K>>5, K128=K>>7;
  const int Kps128=(K128+SK-1)/SK;
  const int ks_lo=pk*Kps128;
  const int ks_hi=min(ks_lo+Kps128,K128);
  const int nsl=ks_hi-ks_lo;
  if(nsl<=0)return;
  extern __shared__ uint8_t _sh[];
  uint8_t* Ald=_sh;
  uint8_t* Asd=_sh+(long)nsl*1024;
  {
    const int nthr=WAVES*64;
    const int nkg=nsl*4;
    const int ngrp=16*nkg;
    const int kb_lo=ks_lo*4;
    for(int g=tid; g<ngrp; g+=nthr){
      const int r=g/nkg;
      const int kbs=g%nkg;
      const int kb=kb_lo+kbs;
      const int k128s=kbs>>2;
      const int kg4=kbs&3;
      const int Lw=kg4*16+r;
      i32x4 o={0,0,0,0}; int e8=0;
      if(r<M){
        const bf16* Ap=A+(long)r*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);
        hw_quant32(ai,o,e8);
      }
      *reinterpret_cast<i32x4*>(Ald+(long)k128s*1024+Lw*16)=o;
      Asd[(long)k128s*64+Lw]=(uint8_t)e8;
    }
  }
  __syncthreads();
  if(!vn)return;
  const long n_col=(long)n_tile*16+m16;
  const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
  const uint8_t* Bsc_r=Bsc+bsc_row;
  const uint8_t* Bsh_L=Bsh+(long)n_tile*(long)K*8+L*16;
  const uint8_t* Ald_L=Ald+L*16;
  const uint8_t* Asd_L=Asd+L;
  f32x4 acc={0,0,0,0};
  const int m_row=m16;
  const int mmsk=(m_row<M)?0xFF:0;
  for(int ks_l=0;ks_l<nsl;ks_l+=KU){
    i32x4 bb[KU]; int bsv[KU];
    i32x4 ab[KU]; int asv[KU];
    const int lim=min(KU,nsl-ks_l);
    #pragma unroll
    for(int u=0;u<KU;++u){
      const int ksi_l=ks_l+u;
      const int ksi_g=ks_lo+ksi_l;
      if(u<lim){
        bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi_g*1024);
        bsv[u]=(int)Bsc_r[(long)(ksi_g>>1)*256+(ksi_g&1)*2];
        ab[u]=*reinterpret_cast<const i32x4*>(Ald_L+(long)ksi_l*1024);
        asv[u]=(int)Asd_L[(long)ksi_l*64] & mmsk;
      } else {
        bb[u]=(i32x4){0,0,0,0};bsv[u]=0;ab[u]=(i32x4){0,0,0,0};asv[u]=0;
      }
    }
    __builtin_amdgcn_sched_barrier(0);
    #pragma unroll
    for(int u=0;u<KU;++u)
      acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
          w8(ab[u]),w8(bb[u]),acc,4,4,0,asv[u],0,bsv[u]);
  }
  #pragma unroll
  for(int i=0;i<4;++i){
    int mo=kg*4+i;
    if(mo<M) atomicAdd(&Cf[(long)mo*N+n_col],acc[i]);
  }
}

__global__ __launch_bounds__(256)
void cast_f32_bf16_z(const float* __restrict__ S,bf16* __restrict__ C,
                     float* __restrict__ Sz,long N){
  long g=(long)blockIdx.x*256+threadIdx.x;
  if(g<N){ C[g]=(bf16)S[g]; Sz[g]=0.0f; }
}

// ═══════ fgemm_alds3: v16-opt (m>=64) ═══════
template<int WAVES,int M_REP,int KU>
__global__ __launch_bounds__(WAVES*64)
void fgemm_alds3(
    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 NTW=(NT+WAVES-1)/WAVES;
  const int m_tile=bid/NTW, ntg=bid%NTW;
  const int n_tile=ntg*WAVES+w;
  const bool vn=(n_tile<NT);
  const int K32=K>>5, K128=K>>7;
  const long Ald_stride=(long)K128*1024;
  const long Asd_stride=(long)K128*64;
  const long Asd_base=(long)M_REP*Ald_stride;
  extern __shared__ uint8_t _sh[];
  uint8_t* Ald=_sh; uint8_t* Asd=_sh+Asd_base;
  {
    const int nthr=WAVES*64;
    const int ngrp=M_REP*16*K32;
    const int m_base=m_tile*M_REP*16;
    for(int g=tid; g<ngrp; g+=nthr){
      const int r=g/K32; const int kb=g%K32;
      const int r_tile=r>>4; const int r16=r&15;
      const int k128s=kb>>2; const int kg4=kb&3;
      const int Lw=kg4*16+r16;
      const int m=m_base+r;
      i32x4 o={0,0,0,0}; int e8=0;
      if(m<M){
        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);
        hw_quant32(ai,o,e8);
      }
      *reinterpret_cast<i32x4*>(Ald+(long)r_tile*Ald_stride+(long)k128s*1024+Lw*16)=o;
      Asd[(long)r_tile*Asd_stride+(long)k128s*64+Lw]=(uint8_t)e8;
    }
  }
  __syncthreads();
  if(!vn) return;
  const long n_col=(long)n_tile*16+m16;
  const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
  const uint8_t* Bsc_r=Bsc+bsc_row;
  const uint8_t* Bsh_L=Bsh+(long)n_tile*(long)K*8+L*16;
  const uint8_t* Ald_L=Ald+L*16;
  const uint8_t* Asd_L=Asd+L;
  f32x4 acc[M_REP];
  #pragma unroll
  for(int r=0;r<M_REP;++r) acc[r]=(f32x4){0,0,0,0};
  int mmsk[M_REP];
  #pragma unroll
  for(int r=0;r<M_REP;++r){
    int m=(m_tile*M_REP+r)*16+m16;
    mmsk[r]=(m<M)?0xFF:0;
  }
  for(int ks=0;ks<K128;ks+=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){
      const int ksi=ks+u;
      bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi*1024);
      bsv[u]=(int)Bsc_r[(long)(ksi>>1)*256+(ksi&1)*2];
      #pragma unroll
      for(int r=0;r<M_REP;++r){
        ab[u][r]=*reinterpret_cast<const i32x4*>(Ald_L+(long)r*Ald_stride+(long)ksi*1024);
        asv[u][r]=(int)Asd_L[(long)r*Asd_stride+(long)ksi*64] & 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]);
    }
  }
  #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_fqn / fgemm_fq (small-M / safety) ═══
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_);
      long c=ks;
      bsv[nr]=(int)Bsc[(n_col[nr]>>5)*(sn8*256)+(n_col[nr]&15)*4+((n_col[nr]>>4)&1)
                      +(c>>3)*256+(c&3)*64+((c>>2)&1)*2] & 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];
  }
}

template<int WAVES>
__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 long n_col=(long)n_tile*16+m16;
  const uint8_t* Bsh_t=Bsh+(long)n_tile*(long)K*8;
  f32x4 acc={0,0,0,0};
  const int m_row=m_tile*16+m16;
  const bool vm=m_row<M;
  for(long k=k_lo;k<k_hi;k+=128){
    i32x4 b4=*reinterpret_cast<const i32x4*>(Bsh_t+(k>>5)*256+L*16);
    long c=(k>>5)+kg;
    int b_sc=(int)Bsc[(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)
                     +(c>>3)*256+(c&3)*64+((c>>2)&1)*2];
    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=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        w8(a4),w8(b4),acc,4,4,0,a_sc,0,b_sc);
  }
  extern __shared__ float red[];
  #pragma unroll
  for(int i=0;i<4;++i)red[(long)w*256+L*4+i]=acc[i];
  __syncthreads();
  if(w!=0)return;
  #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*256+L*4+i];
    acc[i]=s;
  }
  #pragma unroll
  for(int i=0;i<4;++i){
    int mo=m_tile*16+kg*4+i;
    if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[i];
  }
}

#include <torch/extension.h>

template<int W,int KU>
static void _galdsk(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t NT,int64_t SK){
  const int64_t K128=K>>7;
  const int64_t Kps128=(K128+SK-1)/SK;
  const int64_t NTW=(NT+W-1)/W;
  const int64_t gx=NTW*SK;
  const int64_t lds=Kps128*1088;
  static bool _s=false;
  if(!_s){(void)hipFuncSetAttribute((const void*)fgemm_alds_sk<W,KU>,
    hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
  fgemm_alds_sk<W,KU><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    Cf.data_ptr<float>(),(int)M,(int)N,(int)K,sn8,(int)NT,(int)SK);
}

void go_cast(torch::Tensor Cf,torch::Tensor C,torch::Tensor Cfz,int64_t Ne){
  int64_t g=(Ne+255)/256;
  cast_f32_bf16_z<<<dim3(g),dim3(256),0,0>>>(
    Cf.data_ptr<float>(),reinterpret_cast<bf16*>(C.data_ptr()),
    Cfz.data_ptr<float>(),Ne);
}

template<int W,int MR,int KU>
static void _galds3(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 NT){
  const int64_t K128=K>>7;
  const int64_t MT=(M+16*MR-1)/(16*MR);
  const int64_t NTW=(NT+W-1)/W;
  const int64_t gx=MT*NTW;
  const int64_t lds=(int64_t)MR*K128*1088;
  static bool _s=false;
  if(!_s){(void)hipFuncSetAttribute((const void*)fgemm_alds3<W,MR,KU>,
    hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
  fgemm_alds3<W,MR,KU><<<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;
  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>
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*256*4;
  fgemm_fq<W><<<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_alds_sk(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,int64_t M,int64_t N,int64_t K,int64_t sn8,
    int64_t NT,int64_t SK,int64_t W,int64_t KU){
  int64_t K128=K>>7;
  int64_t Kps128=(K128+SK-1)/SK;
  if(Kps128*1088>160*1024) return -2;
  #define D(Ww,Uu) if(W==Ww&&KU==Uu){ \
      _galdsk<Ww,Uu>(A,Bsh,Bsc,Cf,M,N,K,sn8,NT,SK);return 0;}
  D(4,4);D(4,7);D(8,4);D(8,7);
  #undef D
  return -1;
}

int64_t launch_alds3(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 NT,int64_t W,int64_t MR,int64_t KU){
  int64_t K128=K>>7;
  if((int64_t)MR*K128*1088>160*1024) return -2;
  if(K128%KU!=0) return -3;
  #define D(Ww,Rr,Uu) if(W==Ww&&MR==Rr&&KU==Uu){ \
      _galds3<Ww,Rr,Uu>(A,Bsh,Bsc,C,M,N,K,sn8,NT);return 0;}
  D(8,1,4);D(8,1,6);D(8,1,8);D(8,1,12);D(8,1,16);
  D(4,1,4);D(4,1,6);D(4,1,8);D(4,1,12);D(4,1,16);
  D(2,1,4);D(2,1,8);
  D(8,2,4);D(8,2,6);D(4,2,4);D(4,2,6);
  #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(4,1);D(4,2);D(2,2);D(8,1);D(8,2);
  #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){
  #define D(Ww) if(W==Ww){_gfq<Ww>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
  D(4);D(8);
  #undef D
  return -1;
}

void probe(){
  hipFuncAttributes a;
  #define P(k,s) (void)hipFuncGetAttributes(&a,(const void*)k); \
    printf("[v18] %-24s VGPR=%3d spill=%zu\n",s,a.numRegs,a.localSizeBytes);
  P((fgemm_alds_sk<8,7>),"alds_sk<8,7>");
  P((fgemm_alds3<8,1,8>),"alds3<8,1,8>");
  P((fgemm_alds3<8,2,4>),"alds3<8,2,4>");
  P((fgemm_fqn<4,2>),"fqn<4,2>");
  #undef P
}
"""

_CPP = r"""
#include <torch/extension.h>
int64_t launch_alds_sk(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 go_cast(torch::Tensor,torch::Tensor,torch::Tensor,int64_t);
int64_t launch_alds3(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_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);
void probe();
"""

_hip = None
try:
    from torch.utils.cpp_extension import load_inline
    _t0 = time.time()
    _hip = load_inline(name="v18_lb", cpp_sources=_CPP,
        cuda_sources=_HIP_SRC,
        functions=["launch_alds_sk","go_cast","launch_alds3",
                   "launch_fqn","launch_fq","probe"],
        with_cuda=True,
        extra_cuda_cflags=["-O3","--offload-arch=gfx950","-ffast-math"],
        verbose=False)
    _L(f"[v18] HIP compiled {time.time()-_t0:.1f}s"); _hip.probe()
except Exception as ex:
    import traceback
    _L(f"[v18] HIP FAIL: {type(ex).__name__}: {str(ex)[:2000]}")
    for ln in traceback.format_exc().splitlines()[-25:]:
        _L(f"   {ln[:200]}")


# Triton: reference ONLY (correctness check in _build). Never in hot path.
@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
def _make_gemm_ref():
    from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
    @triton.jit
    def _k(A,Bq,Bsc,C,M,N,K,sAm,sBn,sn8,
           BM:tl.constexpr,BN:tl.constexpr,BK:tl.constexpr,EN:tl.constexpr):
        pid=tl.program_id(0);nn=tl.cdiv(N,BN);pm=pid//nn;pn=pid%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)
        bp=Bq+o64[:,None]*sBn+r2[None,:]
        br=_sh_row(o64,sn8);acc=tl.zeros((BM,BN),dtype=tl.float32)
        ap=A+om[:,None].to(tl.int64)*sAm+rk[None,:]
        for kk in tl.range(0,K,BK):
            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(kk//32+r32)[None,:])
            else:
                bf=tl.load(bp,mask=mn[:,None],other=0)
                bs=tl.load(Bsc+br[:,None]+_sh_col(kk//32+r32)[None,:],mask=mn[:,None],other=0)
            acc=tl.dot_scaled(af,asc,"e2m1",tl.trans(bf),bs,"e2m1",acc);bp+=BK//2
        cm=mm[:,None]&mn[None,:]
        tl.store(C+om[:,None].to(tl.int64)*N+on[None,:],acc.to(tl.bfloat16),mask=cm)
    return _k
_gemm_ref=_make_gemm_ref()


# ═══════ HARDCODED TABLE (v17 empirical winners) ═══════
_HARD = {
    (4,   2880, 512 ): ("fqn", 4, 2),
    (16,  2112, 7168): ("aldsk", 8, 8, 7),      # 2-launch. Beats tri SK=14.
    (32,  4096, 512 ): ("fqn", 4, 2),
    (32,  2880, 512 ): ("fqn", 4, 2),
    (64,  7168, 2048): ("alds3", 8, 1, 8),      # 1-launch.
    (256, 3072, 1536): ("alds3", 8, 2, 4),      # 1-launch. MR=2 wins.
}

def _pick_cands(m,n,k):
    NT=-(-n//16);MT16=-(-m//16);K128=k//128
    h=_HARD.get((m,n,k))
    if h: return [h]
    out=[]
    divs=[d for d in(4,6,8,12,16) if K128%d==0 and K128*1088<=160*1024]
    # alds3 for m>=32
    if MT16>=2 and divs:
        for W in(8,4):
            for KU in divs: out.append(("alds3",W,1,KU))
        # MR=2 when grid will be >=~CUs
        if MT16>=8 and 2*K128*1088<=160*1024:
            divs2=[d for d in(4,6) if K128%d==0]
            for KU in divs2: out.append(("alds3",8,2,KU))
    # aldsk for small-M large-K
    if MT16<=2 and K128>=16:
        for SK in(8,14,7):
            if SK>K128:continue
            out.append(("aldsk",8,SK,7))
            out.append(("aldsk",4,SK,4))
    # fqn for small-M small-K
    if MT16<=2 and K128<=16:
        out.append(("fqn",4,2))
        out.append(("fqn",2,2))
    # safety
    out.append(("fqn",4,2))
    out.append(("fq",min(8,K128)))
    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)


_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);MT16=-(-m//16);K128=k//128
    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)
    mn=m*n
    # aldsk ping-pong buffers (pre-zeroed; cast re-zeros for next)
    Cf=torch.zeros((m,n),dtype=torch.float32,device=dev)
    Cf2=torch.zeros((m,n),dtype=torch.float32,device=dev)
    _pp=[Cf,Cf2]

    def _run(cfg,_A,_Bq,_Bsh,_Bsc):
        kind=cfg[0]
        if kind=="alds3":
            _,Wv,MR,KU=cfg
            rc=_hip.launch_alds3(_A,_Bsh,_Bsc,C,m,n,k,sn8,NT,Wv,MR,KU)
            if rc!=0:raise RuntimeError(f"alds3 rc={rc}")
            return C
        if kind=="aldsk":
            _,Wv,SK,KU=cfg
            rc=_hip.launch_alds_sk(_A,_Bsh,_Bsc,_pp[0],m,n,k,sn8,NT,SK,Wv,KU)
            if rc!=0:raise RuntimeError(f"aldsk rc={rc}")
            _hip.go_cast(_pp[0],C,_pp[1],mn)
            _pp[0],_pp[1]=_pp[1],_pp[0]
            return C
        if kind=="fqn":
            _,Wv,NR=cfg
            rc=_hip.launch_fqn(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT16,NT,Wv,NR)
            if rc!=0:raise RuntimeError(f"fqn rc={rc}")
            return C
        if kind=="fq":
            _,Wv=cfg
            rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT16,NT,Wv)
            if rc!=0:raise RuntimeError(f"fq rc={rc}")
            return C
        raise RuntimeError(f"?{cfg}")

    # Triton reference (build-time only, fresh inputs used — NOT closure-captured)
    def _ref_tri(_A,_Bq,_Bsc):
        Cref=torch.empty_like(C)
        BK=min(512,k);gx=MT16*triton.cdiv(n,32)
        _gemm_ref[(gx,)](_A,_Bq,_Bsc,Cref,m,n,k,k,k//2,sn8,
            BM=16,BN=32,BK=BK,EN=(n%32==0),
            num_warps=8,num_stages=2,matrix_instr_nonkdim=16)
        return Cref

    rf=_ref_tri(A,Bq,Bsc).float()
    mag=rf.abs().mean().item()+1e-9

    cands=_pick_cands(m,n,k)
    is_hard=(m,n,k) in _HARD
    _L(f"\n[v18 m={m} n={n} k={k}] {'HARD' if is_hard else 'PICK'}: {len(cands)}c")

    if is_hard:
        cfg=cands[0]
        try:
            C.fill_(float('nan'))
            o=_run(cfg,A,Bq,Bsh,Bsc);torch.cuda.synchronize()
            err=((o.float()-rf).abs().mean()/mag).item()
            if err<5e-3:
                A2=torch.randn_like(A)
                rf2=_ref_tri(A2,Bq,Bsc).float()
                C.fill_(float('nan'))
                o2=_run(cfg,A2,Bq,Bsh,Bsc);torch.cuda.synchronize()
                e2=((o2.float()-rf2).abs().mean()/(rf2.abs().mean()+1e-9)).item()
                if e2<5e-3:
                    _L(f"  {cfg} chk={err:.3%},{e2:.3%} OK")
                    return {"cfg":cfg,"run":_run,"C":C}
                _L(f"  RECHECK FAIL {cfg} e2={e2:.2%}")
            else:
                _L(f"  ERR {cfg} {err:.2%}")
        except Exception as e:
            _L(f"  EXC {cfg} {type(e).__name__}:{e}")

    # Pick mode
    best=None;bt=1e18;log=[];t0=time.time()
    for cfg in cands:
        if time.time()-t0>40: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):_L(f"  {cfg}:ERR{err:.2%}");continue
            t=_tcold(lambda c=cfg:_run(c,A,Bq,Bsh,Bsc))
            log.append((cfg,t))
            if t<bt:bt,best=t,cfg;_L(f"  {cfg}:{t:.2f}us*")
        except Exception as e:
            _L(f"  {cfg}:EXC{type(e).__name__}:{str(e)[:100]}")
            torch.cuda.synchronize()
    if best is None:
        _L("  ->fq fallback");best=("fq",min(8,K128))
    try:
        A2=torch.randn_like(A)
        rf2=_ref_tri(A2,Bq,Bsc).float()
        C.fill_(float('nan'))
        o2=_run(best,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:.2%}");best=("fq",min(8,K128))
    except Exception as e:_L(f"  recheck exc {e}")
    log.sort(key=lambda x:x[1])
    for c,t in log[:8]:_L(f"  top{c}:{t:.2f}")
    _L(f"  ->best={best}@{bt:.2f}us")
    return {"cfg":best,"run":_run,"C":C}


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
    return S["run"](S["cfg"],A,
        data[2].view(torch.uint8),
        data[3].view(torch.uint8),data[4].view(torch.uint8))
scrolls · 765 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 740264.

#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
- v11: K-unrolled split-K + 1-launch fqmn + expanded Triton.
+ v18: HARDCODED v17 winners. ALL-HIP hot path. LB-ready.
- 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.
+ v17 bench GM = 8.10us. All 6 shapes on my HIP kernels:
+ fqn<4,2> small-M 1-launch 6.15-6.28us
+ alds_sk<8,8,7> m=16 2-launch 11.06us (BEATS tri SK=14 @ 11.8)
+ alds3<8,1,8> m=64 1-launch 10.48us
+ alds3<8,2,4> m=256 1-launch 10.40us (MR=2 WINS — 192WGs, B/=2)
- 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.
+ v18 = v17 kernels + v16-style hardcode logic + Bq-fixed.
+ alds_sk ping-pong Cf: cast zeros NEXT buffer → 2 launches not 3.
+ Secret shapes: same pattern-match as v16 (tested 4/4 pass).
"""
import os, sys, time
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
⋯ 2 unchanged lines
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)
⋯ 13 unchanged lines
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;}
⋯ 24 unchanged lines
o=(i32x4){w0,w1,w2,w3};
}
+ // ═══════ fgemm_alds_sk: grid-level split-K (m=16 occ fix) ═══════
+ template<int WAVES,int KU>
+ __global__ __launch_bounds__(WAVES*64)
+ void fgemm_alds_sk(
+ const bf16* __restrict__ A,
+ const uint8_t* __restrict__ Bsh,
+ const uint8_t* __restrict__ Bsc,
+ float* __restrict__ Cf,
+ int M,int N,int K,long sn8,int NT,int SK)
+ {
+ 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 NTW=(NT+WAVES-1)/WAVES;
+ const int pk=bid/NTW, ntg=bid%NTW;
+ const int n_tile=ntg*WAVES+w;
+ const bool vn=(n_tile<NT);
+ const int K32=K>>5, K128=K>>7;
+ const int Kps128=(K128+SK-1)/SK;
+ const int ks_lo=pk*Kps128;
+ const int ks_hi=min(ks_lo+Kps128,K128);
+ const int nsl=ks_hi-ks_lo;
+ if(nsl<=0)return;
+ extern __shared__ uint8_t _sh[];
+ uint8_t* Ald=_sh;
+ uint8_t* Asd=_sh+(long)nsl*1024;
+ {
+ const int nthr=WAVES*64;
+ const int nkg=nsl*4;
+ const int ngrp=16*nkg;
+ const int kb_lo=ks_lo*4;
+ for(int g=tid; g<ngrp; g+=nthr){
+ const int r=g/nkg;
+ const int kbs=g%nkg;
+ const int kb=kb_lo+kbs;
+ const int k128s=kbs>>2;
+ const int kg4=kbs&3;
+ const int Lw=kg4*16+r;
+ i32x4 o={0,0,0,0}; int e8=0;
+ if(r<M){
+ const bf16* Ap=A+(long)r*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);
+ hw_quant32(ai,o,e8);
+ }
+ *reinterpret_cast<i32x4*>(Ald+(long)k128s*1024+Lw*16)=o;
+ Asd[(long)k128s*64+Lw]=(uint8_t)e8;
+ }
+ }
+ __syncthreads();
+ if(!vn)return;
+ const long n_col=(long)n_tile*16+m16;
+ const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
+ const uint8_t* Bsc_r=Bsc+bsc_row;
+ const uint8_t* Bsh_L=Bsh+(long)n_tile*(long)K*8+L*16;
+ const uint8_t* Ald_L=Ald+L*16;
+ const uint8_t* Asd_L=Asd+L;
+ f32x4 acc={0,0,0,0};
+ const int m_row=m16;
+ const int mmsk=(m_row<M)?0xFF:0;
+ for(int ks_l=0;ks_l<nsl;ks_l+=KU){
+ i32x4 bb[KU]; int bsv[KU];
+ i32x4 ab[KU]; int asv[KU];
+ const int lim=min(KU,nsl-ks_l);
+ #pragma unroll
+ for(int u=0;u<KU;++u){
+ const int ksi_l=ks_l+u;
+ const int ksi_g=ks_lo+ksi_l;
+ if(u<lim){
+ bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi_g*1024);
+ bsv[u]=(int)Bsc_r[(long)(ksi_g>>1)*256+(ksi_g&1)*2];
+ ab[u]=*reinterpret_cast<const i32x4*>(Ald_L+(long)ksi_l*1024);
+ asv[u]=(int)Asd_L[(long)ksi_l*64] & mmsk;
+ } else {
+ bb[u]=(i32x4){0,0,0,0};bsv[u]=0;ab[u]=(i32x4){0,0,0,0};asv[u]=0;
+ }
+ }
+ __builtin_amdgcn_sched_barrier(0);
+ #pragma unroll
+ for(int u=0;u<KU;++u)
+ acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
+ w8(ab[u]),w8(bb[u]),acc,4,4,0,asv[u],0,bsv[u]);
+ }
+ #pragma unroll
+ for(int i=0;i<4;++i){
+ int mo=kg*4+i;
+ if(mo<M) atomicAdd(&Cf[(long)mo*N+n_col],acc[i]);
+ }
+ }
+
__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;
+ void cast_f32_bf16_z(const float* __restrict__ S,bf16* __restrict__ C,
+ float* __restrict__ Sz,long N){
+ long g=(long)blockIdx.x*256+threadIdx.x;
+ if(g<N){ C[g]=(bf16)S[g]; Sz[g]=0.0f; }
}
- // ───── 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.
+ // ═══════ fgemm_alds3: v16-opt (m>=64) ═══════
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)
+ void fgemm_alds3(
+ 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 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 int NTW=(NT+WAVES-1)/WAVES;
+ const int m_tile=bid/NTW, ntg=bid%NTW;
+ const int n_tile=ntg*WAVES+w;
+ const bool vn=(n_tile<NT);
+ const int K32=K>>5, K128=K>>7;
+ const long Ald_stride=(long)K128*1024;
+ const long Asd_stride=(long)K128*64;
+ const long Asd_base=(long)M_REP*Ald_stride;
+ extern __shared__ uint8_t _sh[];
+ uint8_t* Ald=_sh; uint8_t* Asd=_sh+Asd_base;
+ {
+ const int nthr=WAVES*64;
+ const int ngrp=M_REP*16*K32;
+ const int m_base=m_tile*M_REP*16;
+ for(int g=tid; g<ngrp; g+=nthr){
+ const int r=g/K32; const int kb=g%K32;
+ const int r_tile=r>>4; const int r16=r&15;
+ const int k128s=kb>>2; const int kg4=kb&3;
+ const int Lw=kg4*16+r16;
+ const int m=m_base+r;
+ i32x4 o={0,0,0,0}; int e8=0;
+ if(m<M){
+ 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);
+ hw_quant32(ai,o,e8);
+ }
+ *reinterpret_cast<i32x4*>(Ald+(long)r_tile*Ald_stride+(long)k128s*1024+Lw*16)=o;
+ Asd[(long)r_tile*Asd_stride+(long)k128s*64+Lw]=(uint8_t)e8;
+ }
+ }
+ __syncthreads();
+ if(!vn) return;
const long n_col=(long)n_tile*16+m16;
- const uint8_t* Bsh_t=Bsh+(long)n_tile*(long)K*8;
-
+ const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
+ const uint8_t* Bsc_r=Bsc+bsc_row;
+ const uint8_t* Bsh_L=Bsh+(long)n_tile*(long)K*8+L*16;
+ const uint8_t* Ald_L=Ald+L*16;
+ const uint8_t* Asd_L=Asd+L;
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];
+ for(int r=0;r<M_REP;++r) acc[r]=(f32x4){0,0,0,0};
+ 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;
+ int m=(m_tile*M_REP+r)*16+m16;
+ mmsk[r]=(m<M)?0xFF: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){
+ for(int ks=0;ks<K128;ks+=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;
+ const int ksi=ks+u;
+ bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi*1024);
+ bsv[u]=(int)Bsc_r[(long)(ksi>>1)*256+(ksi&1)*2];
#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];
+ ab[u][r]=*reinterpret_cast<const i32x4*>(Ald_L+(long)r*Ald_stride+(long)ksi*1024);
+ asv[u][r]=(int)Asd_L[(long)r*Asd_stride+(long)ksi*64] & mmsk[r];
}
}
__builtin_amdgcn_sched_barrier(0);
⋯ 6 unchanged lines
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];
+ 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) ─────
+ // ═══ fgemm_fqn / fgemm_fq (small-M / safety) ═══
template<int WAVES,int N_REP>
__global__ __launch_bounds__(WAVES*64)
void fgemm_fqn(
⋯ 8 unchanged lines
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){
⋯ 6 unchanged lines
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];
+ long c=ks;
+ bsv[nr]=(int)Bsc[(n_col[nr]>>5)*(sn8*256)+(n_col[nr]&15)*4+((n_col[nr]>>4)&1)
+ +(c>>3)*256+(c&3)*64+((c>>2)&1)*2] & vnm[nr];
}
const long kba=k+(long)kg*32;
const bf16* Ap=A+mrow*K+kba;
⋯ 10 unchanged lines
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)
⋯ 20 unchanged lines
}
}
- // ───── fgemm_fq: v5d (safety) ─────
- template<int WAVES,int M_REP>
+ template<int WAVES>
__global__ __launch_bounds__(WAVES*64)
void fgemm_fq(
const bf16* __restrict__ A,
⋯ 6 unchanged lines
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};
+ f32x4 acc={0,0,0,0};
+ const int m_row=m_tile*16+m16;
+ const bool vm=m_row<M;
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);
- }
+ i32x4 b4=*reinterpret_cast<const i32x4*>(Bsh_t+(k>>5)*256+L*16);
+ long c=(k>>5)+kg;
+ int b_sc=(int)Bsc[(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)
+ +(c>>3)*256+(c&3)*64+((c>>2)&1)*2];
+ 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=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
+ w8(a4),w8(b4),acc,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];
+ for(int i=0;i<4;++i)red[(long)w*256+L*4+i]=acc[i];
__syncthreads();
if(w!=0)return;
#pragma unroll
- for(int r=0;r<M_REP;++r)
+ for(int i=0;i<4;++i){
+ float s=0;
#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;
+ for(int ww=0;ww<WAVES;++ww)s+=red[(long)ww*256+L*4+i];
+ acc[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];
- }
+ for(int i=0;i<4;++i){
+ int mo=m_tile*16+kg*4+i;
+ if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[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;
+ template<int W,int KU>
+ static void _galdsk(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
+ torch::Tensor Cf,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t NT,int64_t SK){
+ const int64_t K128=K>>7;
+ const int64_t Kps128=(K128+SK-1)/SK;
+ const int64_t NTW=(NT+W-1)/W;
+ const int64_t gx=NTW*SK;
+ const int64_t lds=Kps128*1088;
static bool _s=false;
- if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_ku<W,MR,KU>,
+ if(!_s){(void)hipFuncSetAttribute((const void*)fgemm_alds_sk<W,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>(),
+ fgemm_alds_sk<W,KU><<<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);
+ Cf.data_ptr<float>(),(int)M,(int)N,(int)K,sn8,(int)NT,(int)SK);
}
- 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;
+ void go_cast(torch::Tensor Cf,torch::Tensor C,torch::Tensor Cfz,int64_t Ne){
+ int64_t g=(Ne+255)/256;
+ cast_f32_bf16_z<<<dim3(g),dim3(256),0,0>>>(
+ Cf.data_ptr<float>(),reinterpret_cast<bf16*>(C.data_ptr()),
+ Cfz.data_ptr<float>(),Ne);
+ }
+
+ template<int W,int MR,int KU>
+ static void _galds3(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 NT){
+ const int64_t K128=K>>7;
+ const int64_t MT=(M+16*MR-1)/(16*MR);
+ const int64_t NTW=(NT+W-1)/W;
+ const int64_t gx=MT*NTW;
+ const int64_t lds=(int64_t)MR*K128*1088;
static bool _s=false;
- if(!_s&&lds>65536){(void)hipFuncSetAttribute((const void*)fgemm_fqmn<W,MR,NR>,
+ if(!_s){(void)hipFuncSetAttribute((const void*)fgemm_alds3<W,MR,KU>,
hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
- fgemm_fqmn<W,MR,NR><<<dim3(gx),dim3(W*64),lds,0>>>(
+ fgemm_alds3<W,MR,KU><<<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()),
⋯ 6 unchanged lines
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>(),
⋯ 1 unchanged lines
(int)M,(int)N,(int)K,sn8,(int)NT);
}
- template<int W,int MR>
+ template<int W>
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>>>(
+ int64_t gx=MT*NT,lds=(int64_t)W*256*4;
+ fgemm_fq<W><<<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);
+ int64_t launch_alds_sk(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
+ torch::Tensor Cf,int64_t M,int64_t N,int64_t K,int64_t sn8,
+ int64_t NT,int64_t SK,int64_t W,int64_t KU){
+ int64_t K128=K>>7;
+ int64_t Kps128=(K128+SK-1)/SK;
+ if(Kps128*1088>160*1024) return -2;
+ #define D(Ww,Uu) if(W==Ww&&KU==Uu){ \
+ _galdsk<Ww,Uu>(A,Bsh,Bsc,Cf,M,N,K,sn8,NT,SK);return 0;}
+ D(4,4);D(4,7);D(8,4);D(8,7);
#undef D
return -1;
}
- int64_t launch_fqmn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
+ int64_t launch_alds3(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);
+ int64_t NT,int64_t W,int64_t MR,int64_t KU){
+ int64_t K128=K>>7;
+ if((int64_t)MR*K128*1088>160*1024) return -2;
+ if(K128%KU!=0) return -3;
+ #define D(Ww,Rr,Uu) if(W==Ww&&MR==Rr&&KU==Uu){ \
+ _galds3<Ww,Rr,Uu>(A,Bsh,Bsc,C,M,N,K,sn8,NT);return 0;}
+ D(8,1,4);D(8,1,6);D(8,1,8);D(8,1,12);D(8,1,16);
+ D(4,1,4);D(4,1,6);D(4,1,8);D(4,1,12);D(4,1,16);
+ D(2,1,4);D(2,1,8);
+ D(8,2,4);D(8,2,6);D(4,2,4);D(4,2,6);
#undef D
return -1;
}
⋯ 3 unchanged lines
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);
+ D(4,1);D(4,2);D(2,2);D(8,1);D(8,2);
#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);
+ int64_t MT,int64_t NT,int64_t W){
+ #define D(Ww) if(W==Ww){_gfq<Ww>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
+ D(4);D(8);
#undef D
return -1;
}
⋯ 1 unchanged lines
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>");
+ printf("[v18] %-24s VGPR=%3d spill=%zu\n",s,a.numRegs,a.localSizeBytes);
+ P((fgemm_alds_sk<8,7>),"alds_sk<8,7>");
+ P((fgemm_alds3<8,1,8>),"alds3<8,1,8>");
+ P((fgemm_alds3<8,2,4>),"alds3<8,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_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_alds_sk(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 go_cast(torch::Tensor,torch::Tensor,torch::Tensor,int64_t);
+ int64_t launch_alds3(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_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);
+ int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
void probe();
"""
⋯ 1 unchanged lines
try:
from torch.utils.cpp_extension import load_inline
_t0 = time.time()
- _hip = load_inline(name="v11_ku", cpp_sources=_CPP,
+ _hip = load_inline(name="v18_lb", cpp_sources=_CPP,
cuda_sources=_HIP_SRC,
- functions=["go_pq","launch_ku","launch_fqmn","launch_fqn","launch_fq","probe"],
+ functions=["launch_alds_sk","go_cast","launch_alds3",
+ "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()
+ _L(f"[v18] 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]}")
+ _L(f"[v18] HIP FAIL: {type(ex).__name__}: {str(ex)[:2000]}")
for ln in traceback.format_exc().splitlines()[-25:]:
_L(f" {ln[:200]}")
+ # Triton: reference ONLY (correctness check in _build). Never in hot path.
@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:
+ def _make_gemm_ref():
+ from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
+ @triton.jit
+ def _k(A,Bq,Bsc,C,M,N,K,sAm,sBn,sn8,
+ BM:tl.constexpr,BN:tl.constexpr,BK:tl.constexpr,EN:tl.constexpr):
+ pid=tl.program_id(0);nn=tl.cdiv(N,BN);pm=pid//nn;pn=pid%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)
+ bp=Bq+o64[:,None]*sBn+r2[None,:]
+ br=_sh_row(o64,sn8);acc=tl.zeros((BM,BN),dtype=tl.float32)
+ ap=A+om[:,None].to(tl.int64)*sAm+rk[None,:]
+ for kk in tl.range(0,K,BK):
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)
+ if EN:
+ bf=tl.load(bp);bs=tl.load(Bsc+br[:,None]+_sh_col(kk//32+r32)[None,:])
+ else:
+ bf=tl.load(bp,mask=mn[:,None],other=0)
+ bs=tl.load(Bsc+br[:,None]+_sh_col(kk//32+r32)[None,:],mask=mn[:,None],other=0)
+ acc=tl.dot_scaled(af,asc,"e2m1",tl.trans(bf),bs,"e2m1",acc);bp+=BK//2
+ cm=mm[:,None]&mn[None,:]
+ tl.store(C+om[:,None].to(tl.int64)*N+on[None,:],acc.to(tl.bfloat16),mask=cm)
+ return _k
+ _gemm_ref=_make_gemm_ref()
- 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
+ # ═══════ HARDCODED TABLE (v17 empirical winners) ═══════
+ _HARD = {
+ (4, 2880, 512 ): ("fqn", 4, 2),
+ (16, 2112, 7168): ("aldsk", 8, 8, 7), # 2-launch. Beats tri SK=14.
+ (32, 4096, 512 ): ("fqn", 4, 2),
+ (32, 2880, 512 ): ("fqn", 4, 2),
+ (64, 7168, 2048): ("alds3", 8, 1, 8), # 1-launch.
+ (256, 3072, 1536): ("alds3", 8, 2, 4), # 1-launch. MR=2 wins.
+ }
-
- def _tri_cfgs(m,n,k):
+ def _pick_cands(m,n,k):
+ NT=-(-n//16);MT16=-(-m//16);K128=k//128
+ h=_HARD.get((m,n,k))
+ if h: return [h]
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))
+ divs=[d for d in(4,6,8,12,16) if K128%d==0 and K128*1088<=160*1024]
+ # alds3 for m>=32
+ if MT16>=2 and divs:
+ for W in(8,4):
+ for KU in divs: out.append(("alds3",W,1,KU))
+ # MR=2 when grid will be >=~CUs
+ if MT16>=8 and 2*K128*1088<=160*1024:
+ divs2=[d for d in(4,6) if K128%d==0]
+ for KU in divs2: out.append(("alds3",8,2,KU))
+ # aldsk for small-M large-K
+ if MT16<=2 and K128>=16:
+ for SK in(8,14,7):
+ if SK>K128:continue
+ out.append(("aldsk",8,SK,7))
+ out.append(("aldsk",4,SK,4))
+ # fqn for small-M small-K
+ if MT16<=2 and K128<=16:
+ out.append(("fqn",4,2))
+ out.append(("fqn",2,2))
+ # safety
+ out.append(("fqn",4,2))
+ out.append(("fq",min(8,K128)))
return out
⋯ 8 unchanged lines
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)
+ sn=Bsc_.shape[1];sn8=sn//8;dev=A.device
+ NT=-(-n//16);MT16=-(-m//16);K128=k//128
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)
+ mn=m*n
+ # aldsk ping-pong buffers (pre-zeroed; cast re-zeros for next)
+ Cf=torch.zeros((m,n),dtype=torch.float32,device=dev)
+ Cf2=torch.zeros((m,n),dtype=torch.float32,device=dev)
+ _pp=[Cf,Cf2]
- def _rh(cfg,_A,_Bq,_Bsh,_Bsc):
- kind,Wv,P2,P3,MT,_=cfg
+ def _run(cfg,_A,_Bq,_Bsh,_Bsc):
+ kind=cfg[0]
+ if kind=="alds3":
+ _,Wv,MR,KU=cfg
+ rc=_hip.launch_alds3(_A,_Bsh,_Bsc,C,m,n,k,sn8,NT,Wv,MR,KU)
+ if rc!=0:raise RuntimeError(f"alds3 rc={rc}")
+ return C
+ if kind=="aldsk":
+ _,Wv,SK,KU=cfg
+ rc=_hip.launch_alds_sk(_A,_Bsh,_Bsc,_pp[0],m,n,k,sn8,NT,SK,Wv,KU)
+ if rc!=0:raise RuntimeError(f"aldsk rc={rc}")
+ _hip.go_cast(_pp[0],C,_pp[1],mn)
+ _pp[0],_pp[1]=_pp[1],_pp[0]
+ return C
+ if kind=="fqn":
+ _,Wv,NR=cfg
+ rc=_hip.launch_fqn(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT16,NT,Wv,NR)
+ if rc!=0:raise RuntimeError(f"fqn rc={rc}")
+ return C
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
+ _,Wv=cfg
+ rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT16,NT,Wv)
+ if rc!=0:raise RuntimeError(f"fq rc={rc}")
+ return C
+ raise RuntimeError(f"?{cfg}")
- _ref_cfg=(16,32,min(512,k),1,8,2,False)
- rf=_rt(_ref_cfg,A,Bq,Bsh,Bsc).clone().float()
+ # Triton reference (build-time only, fresh inputs used — NOT closure-captured)
+ def _ref_tri(_A,_Bq,_Bsc):
+ Cref=torch.empty_like(C)
+ BK=min(512,k);gx=MT16*triton.cdiv(n,32)
+ _gemm_ref[(gx,)](_A,_Bq,_Bsc,Cref,m,n,k,k,k//2,sn8,
+ BM=16,BN=32,BK=BK,EN=(n%32==0),
+ num_warps=8,num_stages=2,matrix_instr_nonkdim=16)
+ return Cref
+
+ rf=_ref_tri(A,Bq,Bsc).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
+ cands=_pick_cands(m,n,k)
+ is_hard=(m,n,k) in _HARD
+ _L(f"\n[v18 m={m} n={n} k={k}] {'HARD' if is_hard else 'PICK'}: {len(cands)}c")
+
+ if is_hard:
+ cfg=cands[0]
try:
C.fill_(float('nan'))
- o=run(cfg,A,Bq,Bsh,Bsc);torch.cuda.synchronize()
+ 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*")
+ if err<5e-3:
+ A2=torch.randn_like(A)
+ rf2=_ref_tri(A2,Bq,Bsc).float()
+ C.fill_(float('nan'))
+ o2=_run(cfg,A2,Bq,Bsh,Bsc);torch.cuda.synchronize()
+ e2=((o2.float()-rf2).abs().mean()/(rf2.abs().mean()+1e-9)).item()
+ if e2<5e-3:
+ _L(f" {cfg} chk={err:.3%},{e2:.3%} OK")
+ return {"cfg":cfg,"run":_run,"C":C}
+ _L(f" RECHECK FAIL {cfg} e2={e2:.2%}")
+ else:
+ _L(f" ERR {cfg} {err:.2%}")
except Exception as e:
- if nerr<4:_L(f" [{tag}]{cfg}:EXC{type(e).__name__}:{str(e)[:100]}");nerr+=1
+ _L(f" EXC {cfg} {type(e).__name__}:{e}")
+
+ # Pick mode
+ best=None;bt=1e18;log=[];t0=time.time()
+ for cfg in cands:
+ if time.time()-t0>40: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):_L(f" {cfg}:ERR{err:.2%}");continue
+ t=_tcold(lambda c=cfg:_run(c,A,Bq,Bsh,Bsc))
+ log.append((cfg,t))
+ if t<bt:bt,best=t,cfg;_L(f" {cfg}:{t:.2f}us*")
+ except Exception as e:
+ _L(f" {cfg}:EXC{type(e).__name__}:{str(e)[:100]}")
torch.cuda.synchronize()
- if best is None:_L(" ->fb");return None
+ if best is None:
+ _L(" ->fq fallback");best=("fq",min(8,K128))
try:
A2=torch.randn_like(A)
- rf2=_rt(_ref_cfg,A2,Bq,Bsh,Bsc).clone().float()
+ rf2=_ref_tri(A2,Bq,Bsc).float()
C.fill_(float('nan'))
- o2=br(best[1],A2,Bq,Bsh,Bsc);torch.cuda.synchronize()
+ o2=_run(best,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}
+ if not(e2<5e-3):_L(f" RECHECK FAIL {best} {e2:.2%}");best=("fq",min(8,K128))
+ except Exception as e:_L(f" recheck exc {e}")
+ log.sort(key=lambda x:x[1])
+ for c,t in log[:8]:_L(f" top{c}:{t:.2f}")
+ _L(f" ->best={best}@{bt:.2f}us")
+ return {"cfg":best,"run":_run,"C":C}
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])
+ S=_build(data);_ST[(m,n,k)]=S
return S["run"](S["cfg"],A,
- data[2].view(torch.uint8),data[3].view(torch.uint8),
- data[4].view(torch.uint8))
+ data[2].view(torch.uint8),
+ data[3].view(torch.uint8),data[4].view(torch.uint8))
scrolls · 1163 diff lines total

Best evidence level for this revision: reported

JSON