#include <string.h>
#include <immintrin.h>
#ifndef KC
#define KC 256
#endif
#ifndef NC
#define NC 384
#endif
#ifndef MC
#define MC 16
#endif
#define MR 4
#define NR 8
static double Abuf[MC*KC+64];
static double Bbuf[NC*KC+64];
__attribute__((target("avx2,fma")))
static inline void mkernel4x8(int kc, const double*ap, const double*bp, double*c, int ldc){
__m256d c0=_mm256_setzero_pd(),c1=_mm256_setzero_pd(),c2=_mm256_setzero_pd(),c3=_mm256_setzero_pd();
__m256d c4=_mm256_setzero_pd(),c5=_mm256_setzero_pd(),c6=_mm256_setzero_pd(),c7=_mm256_setzero_pd();
int k=0;
for(;k+4<=kc;k+=4){
__m256d b0,b1,a0,a1,a2,a3;
#define STEP(OFF) \
b0=_mm256_loadu_pd(bp+(k+OFF)*NR); b1=_mm256_loadu_pd(bp+(k+OFF)*NR+4); \
a0=_mm256_broadcast_sd(ap+k+OFF); a1=_mm256_broadcast_sd(ap+KC+k+OFF); \
a2=_mm256_broadcast_sd(ap+2*KC+k+OFF); a3=_mm256_broadcast_sd(ap+3*KC+k+OFF); \
c0=_mm256_fmadd_pd(a0,b0,c0); c1=_mm256_fmadd_pd(a0,b1,c1); \
c2=_mm256_fmadd_pd(a1,b0,c2); c3=_mm256_fmadd_pd(a1,b1,c3); \
c4=_mm256_fmadd_pd(a2,b0,c4); c5=_mm256_fmadd_pd(a2,b1,c5); \
c6=_mm256_fmadd_pd(a3,b0,c6); c7=_mm256_fmadd_pd(a3,b1,c7);
STEP(0) STEP(1) STEP(2) STEP(3)
#undef STEP
}
for(;k<kc;k++){
__m256d b0=_mm256_loadu_pd(bp+k*NR), b1=_mm256_loadu_pd(bp+k*NR+4);
__m256d a0=_mm256_broadcast_sd(ap+k), a1=_mm256_broadcast_sd(ap+KC+k);
__m256d a2=_mm256_broadcast_sd(ap+2*KC+k), a3=_mm256_broadcast_sd(ap+3*KC+k);
c0=_mm256_fmadd_pd(a0,b0,c0); c1=_mm256_fmadd_pd(a0,b1,c1);
c2=_mm256_fmadd_pd(a1,b0,c2); c3=_mm256_fmadd_pd(a1,b1,c3);
c4=_mm256_fmadd_pd(a2,b0,c4); c5=_mm256_fmadd_pd(a2,b1,c5);
c6=_mm256_fmadd_pd(a3,b0,c6); c7=_mm256_fmadd_pd(a3,b1,c7);
}
_mm256_storeu_pd(c+0*ldc, _mm256_add_pd(_mm256_loadu_pd(c+0*ldc), c0));
_mm256_storeu_pd(c+0*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+0*ldc+4), c1));
_mm256_storeu_pd(c+1*ldc, _mm256_add_pd(_mm256_loadu_pd(c+1*ldc), c2));
_mm256_storeu_pd(c+1*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+1*ldc+4), c3));
_mm256_storeu_pd(c+2*ldc, _mm256_add_pd(_mm256_loadu_pd(c+2*ldc), c4));
_mm256_storeu_pd(c+2*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+2*ldc+4), c5));
_mm256_storeu_pd(c+3*ldc, _mm256_add_pd(_mm256_loadu_pd(c+3*ldc), c6));
_mm256_storeu_pd(c+3*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+3*ldc+4), c7));
}
__attribute__((target("avx2,fma")))
void matrix_multiply(int n, const double *A, const double *B, double *C){
for(int jc=0; jc<n; jc+=NC){
int nb=n-jc; if(nb>NC) nb=NC;
int nrb=(nb+NR-1)/NR;
for(int kk=0; kk<n; kk+=KC){
int kb=n-kk; if(kb>KC) kb=KC;
for(int jb=0; jb<nrb; jb++){
int j0=jc+jb*NR;
double* bp = Bbuf + (size_t)jb*KC*NR;
for(int k=0;k<kb;k++){
const double* src = B + (size_t)(kk+k)*n + j0;
double* dst = bp + k*NR;
dst[0]=src[0];dst[1]=src[1];dst[2]=src[2];dst[3]=src[3];
dst[4]=src[4];dst[5]=src[5];dst[6]=src[6];dst[7]=src[7];
}
}
for(int ii=0; ii<n; ii+=MC){
int mb=n-ii; if(mb>MC) mb=MC;
for(int i=0;i<mb;i++)
memcpy(Abuf+(size_t)i*KC, A+(size_t)(ii+i)*n+kk, (size_t)kb*sizeof(double));
int im=mb-((mb/MR)*MR);
for(int jb=0; jb<nrb; jb++){
double* bp = Bbuf + (size_t)jb*KC*NR;
double* cp = C + (size_t)ii*n + jc + jb*NR;
int i=0;
for(; i<mb-im; i+=MR) mkernel4x8(kb, Abuf+(size_t)i*KC, bp, cp+(size_t)i*n, n);
for(; i<mb; i++){
for(int k=0;k<kb;k++){
double a=Abuf[(size_t)i*KC+k];
double* c=cp+(size_t)i*n;
for(int j=0;j<NR;j++) c[j]+=a*bp[k*NR+j];
}
}
}
}
}
}
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 399.429 ms | 32 MB + 808 KB | Accepted | Score: 100 | 显示更多 |