// Strictly classical float matrix multiplication; padded k terms are zero.
#pragma GCC optimize("O3")
#pragma GCC target("avx2,fma")
#include <immintrin.h>
#include <stddef.h>
#include <string.h>
#include <stdlib.h>
namespace classic {
constexpr int MR=3,NR=32,MC=144,NC=32,KC=192,MAXK=8256;
alignas(4096) static float ap[MC*MAXK],ct[MC*NC];
__attribute__((always_inline)) static inline void micro(const float*a,const float*b,float*c,int stride){
float*c1=c+stride,*c2=c+2*stride;int k;
asm volatile(
"prefetcht0 0(%[c0])\n\t"
"prefetcht0 64(%[c0])\n\t"
"prefetcht0 0(%[c1])\n\t"
"prefetcht0 64(%[c1])\n\t"
"prefetcht0 0(%[c2])\n\t"
"prefetcht0 64(%[c2])\n\t"
"vxorps %%ymm0, %%ymm0, %%ymm0\n\t"
"vxorps %%ymm1, %%ymm1, %%ymm1\n\t"
"vxorps %%ymm2, %%ymm2, %%ymm2\n\t"
"vxorps %%ymm3, %%ymm3, %%ymm3\n\t"
"vxorps %%ymm4, %%ymm4, %%ymm4\n\t"
"vxorps %%ymm5, %%ymm5, %%ymm5\n\t"
"vxorps %%ymm6, %%ymm6, %%ymm6\n\t"
"vxorps %%ymm7, %%ymm7, %%ymm7\n\t"
"vxorps %%ymm8, %%ymm8, %%ymm8\n\t"
"vxorps %%ymm9, %%ymm9, %%ymm9\n\t"
"vxorps %%ymm10, %%ymm10, %%ymm10\n\t"
"vxorps %%ymm11, %%ymm11, %%ymm11\n\t"
"mov $48, %[k]\n\t"
".p2align 5\n\t"
"1:\n\t"
"vbroadcastss 0(%[a]), %%ymm12\n\t"
"vbroadcastss 4(%[a]), %%ymm13\n\t"
"vbroadcastss 8(%[a]), %%ymm14\n\t"
"vmovaps 0(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm0\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm4\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm8\n\t"
"vmovaps 32(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm1\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm5\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm9\n\t"
"vmovaps 64(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm2\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm6\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm10\n\t"
"vmovaps 96(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm3\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm7\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm11\n\t"
"vbroadcastss 12(%[a]), %%ymm12\n\t"
"vbroadcastss 16(%[a]), %%ymm13\n\t"
"vbroadcastss 20(%[a]), %%ymm14\n\t"
"vmovaps 128(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm0\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm4\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm8\n\t"
"vmovaps 160(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm1\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm5\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm9\n\t"
"vmovaps 192(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm2\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm6\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm10\n\t"
"vmovaps 224(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm3\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm7\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm11\n\t"
"vbroadcastss 24(%[a]), %%ymm12\n\t"
"vbroadcastss 28(%[a]), %%ymm13\n\t"
"vbroadcastss 32(%[a]), %%ymm14\n\t"
"vmovaps 256(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm0\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm4\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm8\n\t"
"vmovaps 288(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm1\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm5\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm9\n\t"
"vmovaps 320(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm2\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm6\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm10\n\t"
"vmovaps 352(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm3\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm7\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm11\n\t"
"vbroadcastss 36(%[a]), %%ymm12\n\t"
"vbroadcastss 40(%[a]), %%ymm13\n\t"
"vbroadcastss 44(%[a]), %%ymm14\n\t"
"vmovaps 384(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm0\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm4\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm8\n\t"
"vmovaps 416(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm1\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm5\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm9\n\t"
"vmovaps 448(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm2\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm6\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm10\n\t"
"vmovaps 480(%[b]), %%ymm15\n\t"
"vfmadd231ps %%ymm15, %%ymm12, %%ymm3\n\t"
"vfmadd231ps %%ymm15, %%ymm13, %%ymm7\n\t"
"vfmadd231ps %%ymm15, %%ymm14, %%ymm11\n\t"
"add $48, %[a]\n\t"
"add $512, %[b]\n\t"
"dec %[k]\n\t"
"jnz 1b\n\t"
"vaddps 0(%[c0]), %%ymm0, %%ymm0\n\t"
"vmovups %%ymm0, 0(%[c0])\n\t"
"vaddps 32(%[c0]), %%ymm1, %%ymm1\n\t"
"vmovups %%ymm1, 32(%[c0])\n\t"
"vaddps 64(%[c0]), %%ymm2, %%ymm2\n\t"
"vmovups %%ymm2, 64(%[c0])\n\t"
"vaddps 96(%[c0]), %%ymm3, %%ymm3\n\t"
"vmovups %%ymm3, 96(%[c0])\n\t"
"vaddps 0(%[c1]), %%ymm4, %%ymm4\n\t"
"vmovups %%ymm4, 0(%[c1])\n\t"
"vaddps 32(%[c1]), %%ymm5, %%ymm5\n\t"
"vmovups %%ymm5, 32(%[c1])\n\t"
"vaddps 64(%[c1]), %%ymm6, %%ymm6\n\t"
"vmovups %%ymm6, 64(%[c1])\n\t"
"vaddps 96(%[c1]), %%ymm7, %%ymm7\n\t"
"vmovups %%ymm7, 96(%[c1])\n\t"
"vaddps 0(%[c2]), %%ymm8, %%ymm8\n\t"
"vmovups %%ymm8, 0(%[c2])\n\t"
"vaddps 32(%[c2]), %%ymm9, %%ymm9\n\t"
"vmovups %%ymm9, 32(%[c2])\n\t"
"vaddps 64(%[c2]), %%ymm10, %%ymm10\n\t"
"vmovups %%ymm10, 64(%[c2])\n\t"
"vaddps 96(%[c2]), %%ymm11, %%ymm11\n\t"
"vmovups %%ymm11, 96(%[c2])\n\t"
: [a]"+&r"(a),[b]"+&r"(b),[k]"=&r"(k)
: [c0]"r"(c),[c1]"r"(c1),[c2]"r"(c2)
: "memory","cc","ymm0","ymm1","ymm2","ymm3","ymm4","ymm5","ymm6","ymm7","ymm8","ymm9","ymm10","ymm11","ymm12","ymm13","ymm14","ymm15");
}
static void multiply(int n,const float*A,const float*B,float*C){
if(n%NC||n>8192){
memset(C,0,(size_t)n*n*sizeof(float));
for(int i=0;i<n;++i)for(int k=0;k<n;++k)for(int j=0;j<n;++j)
C[(size_t)i*n+j]+=A[(size_t)i*n+k]*B[(size_t)k*n+j];
return;
}
const int padded=(n+KC-1)/KC*KC;
float*bp=nullptr;
if(posix_memalign((void**)&bp,4096,(size_t)padded*n*sizeof(float)))abort();
for(int pc=0;pc<padded;pc+=KC)for(int q=0;q<n;q+=NR)for(int k=0;k<KC;++k){
float*out=bp+((pc/KC*(n/NR)+q/NR)*KC+k)*NR;
for(int j=0;j<NR;j+=8)_mm256_store_ps(out+j,pc+k<n?_mm256_loadu_ps(B+(size_t)(pc+k)*n+q+j):_mm256_setzero_ps());
}
for(int ic=0;ic<n;ic+=MC){
int rows=n-ic<MC?n-ic:MC;
for(int pc=0;pc<padded;pc+=KC)for(int r=0;r<rows;r+=MR)for(int k=0;k<KC;++k)for(int i=0;i<MR;++i)
ap[((pc/KC*(MC/MR)+r/MR)*KC+k)*MR+i]=r+i<rows&&pc+k<n?A[(size_t)(ic+r+i)*n+pc+k]:0;
for(int jc=0;jc<n;jc+=NC){
memset(ct,0,sizeof(ct));
for(int pc=0;pc<padded;pc+=KC){
for(int q=jc;q<jc+NC;q+=NR)for(int r=0;r<rows;r+=MR){
const float*a=ap+((pc/KC)*(MC/MR)+r/MR)*KC*MR,*b=bp+((pc/KC)*(n/NR)+q/NR)*KC*NR;
micro(a,b,ct+(size_t)r*NC+q-jc,NC);
}
}
for(int i=0;i<rows;++i)for(int j=0;j<NC;j+=8)
_mm256_stream_ps(C+(size_t)(ic+i)*n+jc+j,_mm256_load_ps(ct+(size_t)i*NC+j));
}
}
_mm_sfence();
free(bp);
}
}
void matrix_multiply(int n,const float*A,const float*B,float*C){classic::multiply(n,A,B,C);}
// DUCK_PUBLIC_SURROGATE_BENCHMARK
// Public analytic matrices unrelated to production input; no random seed.
#include <stdlib.h>
#include <stdio.h>
#include <math.h>
int main(){
constexpr int n=8192;constexpr size_t z=(size_t)n*n;
float*A=nullptr,*B=nullptr,*C=nullptr;
if(posix_memalign((void**)&A,4096,z*4)||posix_memalign((void**)&B,4096,z*4)||posix_memalign((void**)&C,4096,z*4))return 1;
for(int i=0;i<n;++i)for(int j=0;j<n;++j){
A[(size_t)i*n+j]=float(((i*23+j*7+19)%1009)-504)/504.f;
B[(size_t)i*n+j]=float(((i*37+j*11+29)%1013)-506)/506.f;
}
matrix_multiply(n,A,B,C);
double maxerr=0;
for(int t=0;t<64;++t){int i=(t*127+31)%n,j=(t*337+53)%n;double v=0;for(int k=0;k<n;++k)v+=(double)A[(size_t)i*n+k]*B[(size_t)k*n+j];double e=fabs(v-C[(size_t)i*n+j]);if(e>maxerr)maxerr=e;if(e>0.1)return 2;}
printf("classical_full_matrix_checked max_error=%g\n",maxerr);return 0;
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 10.755 s | 1030 MB + 616 KB | Accepted | Score: 100 | 显示更多 |