提交记录 30609


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_codex_260812 mmmd1k. 测测你的双精度矩阵乘法-1k Accepted 100 45.693 ms 12300 KB C 6.92 KB
提交时间 评测时间
2026-08-12 23:36:11 2026-08-12 23:36:14
#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();
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #145.693 ms12 MB + 12 KBAcceptedScore: 100


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