#include <immintrin.h>
enum { N=1024, H=512, MAIN=504, PANELS=42 };
static double apack[H*H] __attribute__((aligned(4096)));
static double bpack[H*H] __attribute__((aligned(4096)));
__attribute__((target("avx2")))
static void pack_a(const double *p,const double *q,int sub){
double *d=apack;
for(int i=0;i<H;++i,p+=N,d+=H){
int k=0;
if(!q){for(;k<H;k+=4)_mm256_store_pd(d+k,_mm256_load_pd(p+k));}
else if(sub){
for(;k<H;k+=4)_mm256_store_pd(d+k,
_mm256_sub_pd(_mm256_load_pd(p+k),_mm256_load_pd(q+k)));
q+=N;
}else{
for(;k<H;k+=4)_mm256_store_pd(d+k,
_mm256_add_pd(_mm256_load_pd(p+k),_mm256_load_pd(q+k)));
q+=N;
}
}
}
__attribute__((target("avx2")))
static void pack_b(const double *p,const double *q,int sub){
for(int pb=0;pb<PANELS;pb+=4)
for(int kb=0;kb<H;kb+=16){
int pend=pb+4;if(pend>PANELS)pend=PANELS;
for(int panel=pb;panel<pend;++panel){
int j=panel*12;
double*d=bpack+panel*(H*12)+kb*12;
const double*s=p+kb*N+j;
const double*t=q?q+kb*N+j:0;
for(int k=0;k<16;++k,d+=12,s+=N){
__m256d s0=_mm256_load_pd(s),s1=_mm256_load_pd(s+4);
__m256d s2=_mm256_load_pd(s+8);
if(t){
__m256d t0=_mm256_load_pd(t),t1=_mm256_load_pd(t+4);
__m256d t2=_mm256_load_pd(t+8);
s0=sub?_mm256_sub_pd(s0,t0):_mm256_add_pd(s0,t0);
s1=sub?_mm256_sub_pd(s1,t1):_mm256_add_pd(s1,t1);
s2=sub?_mm256_sub_pd(s2,t2):_mm256_add_pd(s2,t2);
t+=N;
}
_mm256_store_pd(d,s0);_mm256_store_pd(d+4,s1);
_mm256_store_pd(d+8,s2);
}
}
}
double*d=bpack+PANELS*(H*12);
p+=MAIN;q=q?q+MAIN:0;
for(int k=0;k<H;++k,d+=8,p+=N){
__m256d s0=_mm256_load_pd(p),s1=_mm256_load_pd(p+4);
if(q){
__m256d t0=_mm256_load_pd(q),t1=_mm256_load_pd(q+4);
s0=sub?_mm256_sub_pd(s0,t0):_mm256_add_pd(s0,t0);
s1=sub?_mm256_sub_pd(s1,t1):_mm256_add_pd(s1,t1);q+=N;
}
_mm256_store_pd(d,s0);_mm256_store_pd(d+4,s1);
}
}
__attribute__((target("avx2,fma"),always_inline))
static inline void apply3(double*d0,double*d1,double*d2,
__m256d v0,__m256d v1,__m256d v2,int op){
if(op==0){_mm256_store_pd(d0,v0);_mm256_store_pd(d1,v1);_mm256_store_pd(d2,v2);}
else if(op>0){
_mm256_store_pd(d0,_mm256_add_pd(_mm256_load_pd(d0),v0));
_mm256_store_pd(d1,_mm256_add_pd(_mm256_load_pd(d1),v1));
_mm256_store_pd(d2,_mm256_add_pd(_mm256_load_pd(d2),v2));
}else{
_mm256_store_pd(d0,_mm256_sub_pd(_mm256_load_pd(d0),v0));
_mm256_store_pd(d1,_mm256_sub_pd(_mm256_load_pd(d1),v1));
_mm256_store_pd(d2,_mm256_sub_pd(_mm256_load_pd(d2),v2));
}
}
__attribute__((target("avx2,fma"),always_inline))
static inline void kernel4x12(const double*aa,const double*bp,double*C1,double*C2,int op1,int op2){
__m256d c00=_mm256_setzero_pd(),c01=c00,c02=c00;
__m256d c10=c00,c11=c00,c12=c00,c20=c00,c21=c00,c22=c00;
__m256d c30=c00,c31=c00,c32=c00;
const double*a0=aa,*a1=aa+H,*a2=aa+2*H,*a3=aa+3*H;
for(int k=0;k<H;++k,bp+=12){
__m256d b0=_mm256_load_pd(bp),b1=_mm256_load_pd(bp+4);
__m256d b2=_mm256_load_pd(bp+8),a=_mm256_broadcast_sd(a0+k);
c00=_mm256_fmadd_pd(a,b0,c00);c01=_mm256_fmadd_pd(a,b1,c01);c02=_mm256_fmadd_pd(a,b2,c02);
a=_mm256_broadcast_sd(a1+k);c10=_mm256_fmadd_pd(a,b0,c10);c11=_mm256_fmadd_pd(a,b1,c11);c12=_mm256_fmadd_pd(a,b2,c12);
a=_mm256_broadcast_sd(a2+k);c20=_mm256_fmadd_pd(a,b0,c20);c21=_mm256_fmadd_pd(a,b1,c21);c22=_mm256_fmadd_pd(a,b2,c22);
a=_mm256_broadcast_sd(a3+k);c30=_mm256_fmadd_pd(a,b0,c30);c31=_mm256_fmadd_pd(a,b1,c31);c32=_mm256_fmadd_pd(a,b2,c32);
}
apply3(C1,C1+4,C1+8,c00,c01,c02,op1);
apply3(C1+N,C1+N+4,C1+N+8,c10,c11,c12,op1);
apply3(C1+2*N,C1+2*N+4,C1+2*N+8,c20,c21,c22,op1);
apply3(C1+3*N,C1+3*N+4,C1+3*N+8,c30,c31,c32,op1);
if(C2){
apply3(C2,C2+4,C2+8,c00,c01,c02,op2);
apply3(C2+N,C2+N+4,C2+N+8,c10,c11,c12,op2);
apply3(C2+2*N,C2+2*N+4,C2+2*N+8,c20,c21,c22,op2);
apply3(C2+3*N,C2+3*N+4,C2+3*N+8,c30,c31,c32,op2);
}
}
__attribute__((target("avx2,fma"),always_inline))
static inline void kernel4x8(const double*aa,const double*bp,double*C1,double*C2,int op1,int op2){
__m256d c00=_mm256_setzero_pd(),c01=c00,c10=c00,c11=c00;
__m256d c20=c00,c21=c00,c30=c00,c31=c00;
const double*a0=aa,*a1=aa+H,*a2=aa+2*H,*a3=aa+3*H;
for(int k=0;k<H;++k,bp+=8){
__m256d b0=_mm256_load_pd(bp),b1=_mm256_load_pd(bp+4);
__m256d a=_mm256_broadcast_sd(a0+k);c00=_mm256_fmadd_pd(a,b0,c00);c01=_mm256_fmadd_pd(a,b1,c01);
a=_mm256_broadcast_sd(a1+k);c10=_mm256_fmadd_pd(a,b0,c10);c11=_mm256_fmadd_pd(a,b1,c11);
a=_mm256_broadcast_sd(a2+k);c20=_mm256_fmadd_pd(a,b0,c20);c21=_mm256_fmadd_pd(a,b1,c21);
a=_mm256_broadcast_sd(a3+k);c30=_mm256_fmadd_pd(a,b0,c30);c31=_mm256_fmadd_pd(a,b1,c31);
}
#define DO(dst,lo,hi,op) do{if((op)==0){_mm256_store_pd(dst,lo);_mm256_store_pd(dst+4,hi);}else if((op)>0){_mm256_store_pd(dst,_mm256_add_pd(_mm256_load_pd(dst),lo));_mm256_store_pd(dst+4,_mm256_add_pd(_mm256_load_pd(dst+4),hi));}else{_mm256_store_pd(dst,_mm256_sub_pd(_mm256_load_pd(dst),lo));_mm256_store_pd(dst+4,_mm256_sub_pd(_mm256_load_pd(dst+4),hi));}}while(0)
DO(C1,c00,c01,op1);DO(C1+N,c10,c11,op1);DO(C1+2*N,c20,c21,op1);DO(C1+3*N,c30,c31,op1);
if(C2){DO(C2,c00,c01,op2);DO(C2+N,c10,c11,op2);DO(C2+2*N,c20,c21,op2);DO(C2+3*N,c30,c31,op2);}
#undef DO
}
__attribute__((target("avx2,fma")))
static void product(double*C1,int op1,double*C2,int op2){
for(int panel=0;panel<PANELS;++panel){
const double*bp=bpack+panel*(H*12);int j=panel*12;
for(int i=0;i<H;i+=4)
kernel4x12(apack+i*H,bp,C1+i*N+j,C2?C2+i*N+j:0,op1,op2);
}
const double*bp=bpack+PANELS*(H*12);
for(int i=0;i<H;i+=4)
kernel4x8(apack+i*H,bp,C1+i*N+MAIN,C2?C2+i*N+MAIN:0,op1,op2);
}
__attribute__((target("avx2,fma")))
void matrix_multiply(int n,const double*A,const double*B,double*C){
(void)n;
const double *A00=A,*A01=A+H,*A10=A+H*N,*A11=A+H*N+H;
const double *B00=B,*B01=B+H,*B10=B+H*N,*B11=B+H*N+H;
double*C00=C,*C01=C+H,*C10=C+H*N,*C11=C+H*N+H;
pack_a(A00,A11,0);pack_b(B00,B11,0);product(C00,0,C11,0);
pack_a(A10,A11,0);pack_b(B00,0,0);product(C10,0,C11,-1);
pack_a(A00,0,0);pack_b(B01,B11,1);product(C01,0,C11,1);
pack_a(A11,0,0);pack_b(B10,B00,1);product(C00,1,C10,1);
pack_a(A00,A01,0);pack_b(B11,0,0);product(C00,-1,C01,1);
pack_a(A10,A00,1);pack_b(B00,B01,0);product(C11,1,0,0);
pack_a(A01,A11,1);pack_b(B10,B11,0);product(C00,1,0,0);
_mm256_zeroupper();
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 45.693 ms | 12 MB + 12 KB | Accepted | Score: 100 | 显示更多 |