#pragma GCC optimize("O3,unroll-loops,omit-frame-pointer")
#pragma GCC target("arch=skylake")
#define STOP_AFTER 4
#define MAYBE_STOP(K) do{if(STOP_AFTER==(K)){_mm256_zeroupper();return;}}while(0)
#include <immintrin.h>
enum { N=1024, H=512, Q=256, MAIN=252, PANELS=21 };
static double amat[H*H] __attribute__((aligned(4096)));
static double bmat[H*H] __attribute__((aligned(4096)));
static double suba[Q*Q] __attribute__((aligned(4096)));
static double subb[Q*Q] __attribute__((aligned(4096)));
static double product512[H*H] __attribute__((aligned(4096)));
static void pack_top(double *dst,const double *p,const double *q,int sub){
for(int i=0;i<H;++i,p+=N,dst+=H){
if(!q){
for(int j=0;j<H;j+=4)_mm256_store_pd(dst+j,_mm256_load_pd(p+j));
}else if(sub){
for(int j=0;j<H;j+=4)_mm256_store_pd(dst+j,
_mm256_sub_pd(_mm256_load_pd(p+j),_mm256_load_pd(q+j)));
q+=N;
}else{
for(int j=0;j<H;j+=4)_mm256_store_pd(dst+j,
_mm256_add_pd(_mm256_load_pd(p+j),_mm256_load_pd(q+j)));
q+=N;
}
}
}
static void pack_sub_a(const double *p,const double *q,int sub){
double *d=suba;
for(int i=0;i<Q;++i,p+=H,d+=Q){
if(!q){
for(int j=0;j<Q;j+=4)_mm256_store_pd(d+j,_mm256_load_pd(p+j));
}else if(sub){
for(int j=0;j<Q;j+=4)_mm256_store_pd(d+j,
_mm256_sub_pd(_mm256_load_pd(p+j),_mm256_load_pd(q+j)));
q+=H;
}else{
for(int j=0;j<Q;j+=4)_mm256_store_pd(d+j,
_mm256_add_pd(_mm256_load_pd(p+j),_mm256_load_pd(q+j)));
q+=H;
}
}
}
static void pack_sub_b(const double *p,const double *q,int sub){
for(int panel=0;panel<PANELS;++panel){
int j=panel*12;
double *d=subb+panel*(Q*12);
const double *s=p+j,*t=q?q+j:0;
for(int k=0;k<Q;++k,d+=12,s+=H){
__m256d x0=_mm256_load_pd(s),x1=_mm256_load_pd(s+4),x2=_mm256_load_pd(s+8);
if(t){
__m256d y0=_mm256_load_pd(t),y1=_mm256_load_pd(t+4),y2=_mm256_load_pd(t+8);
x0=sub?_mm256_sub_pd(x0,y0):_mm256_add_pd(x0,y0);
x1=sub?_mm256_sub_pd(x1,y1):_mm256_add_pd(x1,y1);
x2=sub?_mm256_sub_pd(x2,y2):_mm256_add_pd(x2,y2);
t+=H;
}
_mm256_store_pd(d,x0);_mm256_store_pd(d+4,x1);_mm256_store_pd(d+8,x2);
}
}
double *d=subb+PANELS*(Q*12);
p+=MAIN;q=q?q+MAIN:0;
for(int k=0;k<Q;++k,d+=4,p+=H){
__m256d x=_mm256_load_pd(p);
if(q){__m256d y=_mm256_load_pd(q);x=sub?_mm256_sub_pd(x,y):_mm256_add_pd(x,y);q+=H;}
_mm256_store_pd(d,x);
}
}
static inline void put3(double*d,__m256d x0,__m256d x1,__m256d x2,int op){
if(op==0){_mm256_store_pd(d,x0);_mm256_store_pd(d+4,x1);_mm256_store_pd(d+8,x2);}
else if(op>0){
_mm256_store_pd(d,_mm256_add_pd(_mm256_load_pd(d),x0));
_mm256_store_pd(d+4,_mm256_add_pd(_mm256_load_pd(d+4),x1));
_mm256_store_pd(d+8,_mm256_add_pd(_mm256_load_pd(d+8),x2));
}else{
_mm256_store_pd(d,_mm256_sub_pd(_mm256_load_pd(d),x0));
_mm256_store_pd(d+4,_mm256_sub_pd(_mm256_load_pd(d+4),x1));
_mm256_store_pd(d+8,_mm256_sub_pd(_mm256_load_pd(d+8),x2));
}
}
static inline void put1(double*d,__m256d x,int op){
if(op==0)_mm256_store_pd(d,x);
else if(op>0)_mm256_store_pd(d,_mm256_add_pd(_mm256_load_pd(d),x));
else _mm256_store_pd(d,_mm256_sub_pd(_mm256_load_pd(d),x));
}
static inline void kernel4x12(const double *aa,const double *bp,
double *c1,double *c2,int op1,int op2){
__m256d z=_mm256_setzero_pd();
__m256d c00=z,c01=z,c02=z,c10=z,c11=z,c12=z;
__m256d c20=z,c21=z,c22=z,c30=z,c31=z,c32=z;
const double *a0=aa,*a1=aa+Q,*a2=aa+2*Q,*a3=aa+3*Q;
for(int k=0;k<Q;++k,bp+=12){
__m256d b0=_mm256_load_pd(bp),b1=_mm256_load_pd(bp+4),b2=_mm256_load_pd(bp+8);
__m256d 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);
}
put3(c1,c00,c01,c02,op1);put3(c1+H,c10,c11,c12,op1);
put3(c1+2*H,c20,c21,c22,op1);put3(c1+3*H,c30,c31,c32,op1);
if(c2){put3(c2,c00,c01,c02,op2);put3(c2+H,c10,c11,c12,op2);
put3(c2+2*H,c20,c21,c22,op2);put3(c2+3*H,c30,c31,c32,op2);}
}
static inline void kernel4x4(const double *aa,const double *bp,
double *c1,double *c2,int op1,int op2){
__m256d z=_mm256_setzero_pd(),c0=z,c1v=z,c2v=z,c3=z;
const double *a0=aa,*a1=aa+Q,*a2=aa+2*Q,*a3=aa+3*Q;
for(int k=0;k<Q;++k,bp+=4){
__m256d b=_mm256_load_pd(bp),a=_mm256_broadcast_sd(a0+k);
c0=_mm256_fmadd_pd(a,b,c0);a=_mm256_broadcast_sd(a1+k);
c1v=_mm256_fmadd_pd(a,b,c1v);a=_mm256_broadcast_sd(a2+k);
c2v=_mm256_fmadd_pd(a,b,c2v);a=_mm256_broadcast_sd(a3+k);
c3=_mm256_fmadd_pd(a,b,c3);
}
put1(c1,c0,op1);put1(c1+H,c1v,op1);put1(c1+2*H,c2v,op1);put1(c1+3*H,c3,op1);
if(c2){put1(c2,c0,op2);put1(c2+H,c1v,op2);put1(c2+2*H,c2v,op2);put1(c2+3*H,c3,op2);}
}
static void direct256(double *c1,int op1,double *c2,int op2){
for(int panel=0;panel<PANELS;++panel){
const double *bp=subb+panel*(Q*12);int j=panel*12;
for(int i=0;i<Q;i+=4)
kernel4x12(suba+i*Q,bp,c1+i*H+j,c2?c2+i*H+j:0,op1,op2);
}
const double *bp=subb+PANELS*(Q*12);
for(int i=0;i<Q;i+=4)
kernel4x4(suba+i*Q,bp,c1+i*H+MAIN,c2?c2+i*H+MAIN:0,op1,op2);
}
static void inner_strassen(void){
const double *A00=amat,*A01=amat+Q,*A10=amat+Q*H,*A11=A10+Q;
const double *B00=bmat,*B01=bmat+Q,*B10=bmat+Q*H,*B11=B10+Q;
double *C00=product512,*C01=product512+Q,*C10=product512+Q*H,*C11=C10+Q;
pack_sub_a(A00,A11,0);pack_sub_b(B00,B11,0);direct256(C00,0,C11,0);
pack_sub_a(A10,A11,0);pack_sub_b(B00,0,0);direct256(C10,0,C11,-1);
pack_sub_a(A00,0,0);pack_sub_b(B01,B11,1);direct256(C01,0,C11,1);
pack_sub_a(A11,0,0);pack_sub_b(B10,B00,1);direct256(C00,1,C10,1);
pack_sub_a(A00,A01,0);pack_sub_b(B11,0,0);direct256(C00,-1,C01,1);
pack_sub_a(A10,A00,1);pack_sub_b(B00,B01,0);direct256(C11,1,0,0);
pack_sub_a(A01,A11,1);pack_sub_b(B10,B11,0);direct256(C00,1,0,0);
}
static void apply_top(double *c1,int op1,double *c2,int op2){
for(int i=0;i<H;++i){
const double *s=product512+i*H;double *d1=c1+i*N,*d2=c2?c2+i*N:0;
for(int j=0;j<H;j+=4){
__m256d x=_mm256_load_pd(s+j);
if(op1==0)_mm256_store_pd(d1+j,x);
else if(op1>0)_mm256_store_pd(d1+j,_mm256_add_pd(_mm256_load_pd(d1+j),x));
else _mm256_store_pd(d1+j,_mm256_sub_pd(_mm256_load_pd(d1+j),x));
if(d2){
if(op2==0)_mm256_store_pd(d2+j,x);
else if(op2>0)_mm256_store_pd(d2+j,_mm256_add_pd(_mm256_load_pd(d2+j),x));
else _mm256_store_pd(d2+j,_mm256_sub_pd(_mm256_load_pd(d2+j),x));
}
}
}
}
static void top_product(const double *a,const double *aq,int asub,
const double *b,const double *bq,int bsub,
double *c1,int op1,double *c2,int op2){
pack_top(amat,a,aq,asub);pack_top(bmat,b,bq,bsub);
inner_strassen();apply_top(c1,op1,c2,op2);
}
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=A10+H;
const double *B00=B,*B01=B+H,*B10=B+H*N,*B11=B10+H;
double *C00=C,*C01=C+H,*C10=C+H*N,*C11=C10+H;
top_product(A00,A11,0,B00,B11,0,C00,0,C11,0);MAYBE_STOP(1);
top_product(A10,A11,0,B00,0,0,C10,0,C11,-1);MAYBE_STOP(2);
top_product(A00,0,0,B01,B11,1,C01,0,C11,1);MAYBE_STOP(3);
top_product(A11,0,0,B10,B00,1,C00,1,C10,1);MAYBE_STOP(4);
top_product(A00,A01,0,B11,0,0,C00,-1,C01,1);MAYBE_STOP(5);
top_product(A10,A00,1,B00,B01,0,C11,1,0,0);MAYBE_STOP(6);
top_product(A01,A11,1,B10,B11,0,C00,1,0,0);
_mm256_zeroupper();
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 27.011 ms | 15 MB + 12 KB | Wrong Answer | Score: 0 | 显示更多 |