提交记录 47392


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_codex_260812 mmmd1k. 测测你的双精度矩阵乘法-1k Wrong Answer 0 13.84 ms 13324 KB C 8.16 KB
提交时间 评测时间
2026-08-23 16:00:31 2026-08-23 16:00:33
#pragma GCC optimize("O3,unroll-loops,omit-frame-pointer")
#pragma GCC target("arch=skylake")
#define STOP_AFTER 2
#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();
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #113.84 ms13 MB + 12 KBWrong AnswerScore: 0


Judge Duck Online | 评测鸭在线
Server Time: 2026-08-28 12:39:26 | Loaded in 1 ms | Server Status
个人娱乐项目,仅供学习交流使用 | 捐赠