提交记录 85050


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_codex_maker test. 自定义测试 Accepted 100 258.426 ms 49316 KB C++17 8.54 KB
提交时间 评测时间
2026-09-23 00:17:19 2026-09-23 00:17:22
#define MATRIX_TYPE float
#define MATRIX_KIND 1
// Original blocked AVX2 kernel. All integer arithmetic has unsigned modular semantics.
#include <immintrin.h>
#include <stdint.h>
#include <stdlib.h>
#include <string.h>
#pragma GCC target("avx2,fma")
namespace matrix_impl {
using T = MATRIX_TYPE;
#if MATRIX_KIND == 0
using U=double; using P=double; using V=__m256d;
constexpr int VL=4;
static inline V load(const P*p){return _mm256_loadu_pd(p);}
static inline void store(P*p,V v){_mm256_storeu_pd(p,v);}
static inline V broadcast(P x){return _mm256_set1_pd(x);}
static inline V madd(V a,V b,V c){return _mm256_fmadd_pd(a,b,c);}
#elif MATRIX_KIND == 1
using U=float; using P=float; using V=__m256;
constexpr int VL=8;
static inline V load(const P*p){return _mm256_loadu_ps(p);}
static inline void store(P*p,V v){_mm256_storeu_ps(p,v);}
static inline V broadcast(P x){return _mm256_set1_ps(x);}
static inline V madd(V a,V b,V c){return _mm256_fmadd_ps(a,b,c);}
#else
using V=__m256i;
#if MATRIX_KIND == 2
using U=unsigned int; using P=U;
constexpr int VL=8;
static inline V broadcast(P x){int y;memcpy(&y,&x,4);return _mm256_set1_epi32(y);}
static inline V madd(V a,V b,V c){return _mm256_add_epi32(c,_mm256_mullo_epi32(a,b));}
#elif MATRIX_KIND == 3
using U=unsigned long long; using P=U;
constexpr int VL=4;
static inline V broadcast(P x){long long y;memcpy(&y,&x,8);return _mm256_set1_epi64x(y);}
static inline V madd(V a,V b,V c){
    V lo=_mm256_mul_epu32(a,b);
    V hi=_mm256_add_epi64(_mm256_mul_epu32(_mm256_srli_epi64(a,32),b),
                          _mm256_mul_epu32(a,_mm256_srli_epi64(b,32)));
    return _mm256_add_epi64(c,_mm256_add_epi64(lo,_mm256_slli_epi64(hi,32)));
}
#else
#if MATRIX_KIND == 4
using U=unsigned short;
#else
using U=unsigned char;
#endif
using P=unsigned short;
constexpr int VL=16;
static inline V broadcast(P x){short y;memcpy(&y,&x,2);return _mm256_set1_epi16(y);}
static inline V madd(V a,V b,V c){return _mm256_add_epi16(c,_mm256_mullo_epi16(a,b));}
#endif
static inline V load(const P*p){return _mm256_loadu_si256((const __m256i*)p);}
static inline void store(P*p,V v){_mm256_storeu_si256((__m256i*)p,v);}
#endif
constexpr int MR=4, NR=2*VL, MC=128, NC=128, KC=128;
alignas(4096) static P apack[MC*KC],bpack[NC*KC];
#if MATRIX_KIND == 5
static inline V loadc(const U*p){return _mm256_cvtepu8_epi16(_mm_loadu_si128((const __m128i*)p));}
static inline void storec(U*p,V v){alignas(32) P tmp[VL];store(tmp,v);for(int j=0;j<VL;++j)p[j]=(U)tmp[j];}
#else
static inline V loadc(const U*p){return load(p);}
static inline void storec(U*p,V v){store(p,v);}
#endif
static inline void micro(const P*a,const P*b,U*c,int stride) {
    V c00=loadc(c),c01=loadc(c+VL);
    V c10=loadc(c+stride),c11=loadc(c+stride+VL);
    V c20=loadc(c+2*stride),c21=loadc(c+2*stride+VL);
    V c30=loadc(c+3*stride),c31=loadc(c+3*stride+VL);
    for(int k=0;k<KC;++k) {
        V b0=load(b),b1=load(b+VL);
        V a0=broadcast(a[0]); c00=madd(a0,b0,c00); c01=madd(a0,b1,c01);
        V a1=broadcast(a[1]); c10=madd(a1,b0,c10); c11=madd(a1,b1,c11);
        V a2=broadcast(a[2]); c20=madd(a2,b0,c20); c21=madd(a2,b1,c21);
        V a3=broadcast(a[3]); c30=madd(a3,b0,c30); c31=madd(a3,b1,c31);
        a+=MR;b+=NR;
    }
    storec(c,c00);storec(c+VL,c01);
    storec(c+stride,c10);storec(c+stride+VL,c11);
    storec(c+2*stride,c20);storec(c+2*stride+VL,c21);
    storec(c+3*stride,c30);storec(c+3*stride+VL,c31);
}
static void classical(int n,const U*A,int sa,const U*B,int sb,U*C,int sc){
    for(int i=0;i<n;++i)memset(C+i*sc,0,n*sizeof(U));
    if(n%MC) {
        for(int i=0;i<n;++i)for(int k=0;k<n;++k)for(int j=0;j<n;++j)
#if MATRIX_KIND < 2
            C[i*sc+j]+=A[i*sa+k]*B[k*sb+j];
#else
            C[i*sc+j]=(U)(C[i*sc+j]+(unsigned long long)A[i*sa+k]*B[k*sb+j]);
#endif
        return;
    }
    for(int ic=0;ic<n;ic+=MC)for(int pc=0;pc<n;pc+=KC){
        for(int r=0;r<MC;r+=MR)for(int k=0;k<KC;++k)for(int i=0;i<MR;++i)
            apack[(r/MR*KC+k)*MR+i]=(P)A[(ic+r+i)*sa+pc+k];
        for(int jc=0;jc<n;jc+=NC){
            for(int q=0;q<NC;q+=NR)for(int k=0;k<KC;++k)for(int j=0;j<NR;++j)
                bpack[(q/NR*KC+k)*NR+j]=(P)B[(pc+k)*sb+jc+q+j];
            for(int r=0;r<MC;r+=MR)for(int q=0;q<NC;q+=NR)
                micro(apack+r*KC,bpack+q*KC,C+(ic+r)*sc+jc+q,sc);
        }
    }
}
#if MATRIX_KIND == 3
// Strassen's identities are exact in Z/(2^64). Workspace < n*n elements.
static void combine(int n,const U*a,int sa,const U*b,int sb,U*out,int sign){
    for(int i=0;i<n;++i)for(int j=0;j<n;++j)
        out[i*n+j]=sign>0?a[i*sa+j]+b[i*sb+j]:a[i*sa+j]-b[i*sb+j];
}
static void accumulate(int n,U*out,int so,const U*p,int sign){
    for(int i=0;i<n;++i)for(int j=0;j<n;++j)
        if(sign>0)out[i*so+j]+=p[i*n+j];else out[i*so+j]-=p[i*n+j];
}
static void recurse(int n,const U*a,int sa,const U*b,int sb,U*c,int sc,U*work){
    if(n<=128){classical(n,a,sa,b,sb,c,sc);return;}
    int h=n/2;size_t count=(size_t)h*h;
    U*s=work,*t=s+count,*p=t+count,*next=p+count;
    const U*a00=a,*a01=a+h,*a10=a+h*sa,*a11=a10+h;
    const U*b00=b,*b01=b+h,*b10=b+h*sb,*b11=b10+h;
    U*c00=c,*c01=c+h,*c10=c+h*sc,*c11=c10+h;
    for(int i=0;i<n;++i)memset(c+i*sc,0,n*sizeof(U));
    combine(h,a00,sa,a11,sa,s,1);combine(h,b00,sb,b11,sb,t,1);
    recurse(h,s,h,t,h,p,h,next);accumulate(h,c00,sc,p,1);accumulate(h,c11,sc,p,1);
    combine(h,a10,sa,a11,sa,s,1);
    recurse(h,s,h,b00,sb,p,h,next);accumulate(h,c10,sc,p,1);accumulate(h,c11,sc,p,-1);
    combine(h,b01,sb,b11,sb,t,-1);
    recurse(h,a00,sa,t,h,p,h,next);accumulate(h,c01,sc,p,1);accumulate(h,c11,sc,p,1);
    combine(h,b10,sb,b00,sb,t,-1);
    recurse(h,a11,sa,t,h,p,h,next);accumulate(h,c00,sc,p,1);accumulate(h,c10,sc,p,1);
    combine(h,a00,sa,a01,sa,s,1);
    recurse(h,s,h,b11,sb,p,h,next);accumulate(h,c00,sc,p,-1);accumulate(h,c01,sc,p,1);
    combine(h,a10,sa,a00,sa,s,-1);combine(h,b00,sb,b01,sb,t,1);
    recurse(h,s,h,t,h,p,h,next);accumulate(h,c11,sc,p,1);
    combine(h,a01,sa,a11,sa,s,-1);combine(h,b10,sb,b11,sb,t,1);
    recurse(h,s,h,t,h,p,h,next);accumulate(h,c00,sc,p,1);
}
#endif
} // namespace matrix_impl
void matrix_multiply(int n,const MATRIX_TYPE*A,const MATRIX_TYPE*B,MATRIX_TYPE*C){
    using namespace matrix_impl;
#if MATRIX_KIND == 3
    if(n>=256&&(n&(n-1))==0){
        U*work=nullptr;
        if(posix_memalign((void**)&work,4096,(size_t)n*n*sizeof(U)))abort();
        recurse(n,(const U*)A,n,(const U*)B,n,(U*)C,n,work);
        free(work);
    }else
#endif
    classical(n,(const U*)A,n,(const U*)B,n,(U*)C,n);
}

static uint64_t next_random(uint64_t &s) {
    s+=UINT64_C(0x9e3779b97f4a7c15);
    uint64_t z=s;
    z=(z^(z>>30))*UINT64_C(0xbf58476d1ce4e5b9);
    z=(z^(z>>27))*UINT64_C(0x94d049bb133111eb);
    return z^(z>>31);
}


#include <stdio.h>
#include <float.h>
#include <math.h>
#include <assert.h>
static MATRIX_TYPE *allocate_matrix(int n) {
    MATRIX_TYPE*p=nullptr;
    if(posix_memalign((void**)&p,4096,(size_t)n*n*sizeof(MATRIX_TYPE)))abort();
    memset(p,0,(size_t)n*n*sizeof(MATRIX_TYPE));
    return p;
}
static void check_one(int n,const MATRIX_TYPE*A,const MATRIX_TYPE*B,const MATRIX_TYPE*C,int i,int j) {
#if MATRIX_KIND<2
    long double expected=0;
    for(int k=0;k<n;++k)expected+=(long double)A[i*n+k]*B[k*n+j];
    long double eps=MATRIX_KIND==0?DBL_EPSILON:FLT_EPSILON;
    if(!isfinite(C[i*n+j])||fabsl((long double)C[i*n+j]-expected)>3.L*n*n*eps)abort();
#else
    using U=matrix_impl::U;
    U expected=0;
    for(int k=0;k<n;++k)
        expected=(U)(expected+(unsigned long long)(U)A[i*n+k]*(U)B[k*n+j]);
    U got;memcpy(&got,C+i*n+j,sizeof(U));
    if(got!=expected)abort();
#endif
}
static void run(int n,bool exhaustive) {
    MATRIX_TYPE*A=allocate_matrix(n),*B=allocate_matrix(n),*C=allocate_matrix(n);
    uint64_t state=UINT64_C(0x123456789abcdef0) ^ n ^ ((uint64_t)102<<48);
    for(int which=0;which<2;++which){
        MATRIX_TYPE*a=which?B:A;
        for(size_t i=0;i<(size_t)n*n;++i){a[i]=(float)((next_random(state)>>40)*(2.0/(16777216.0-1.0))-1.0);}
    }
    matrix_multiply(n,A,B,C);
    if(exhaustive) {
        for(int i=0;i<n;++i)for(int j=0;j<n;++j)check_one(n,A,B,C,i,j);
    }else{
        for(int t=0;t<256;++t){
            int i=t<4?(t&1?n-1:0):(int)(next_random(state)%n);
            int j=t<4?(t&2?n-1:0):(int)(next_random(state)%n);
            check_one(n,A,B,C,i,j);
        }
    }
    uint64_t digest=UINT64_C(14695981039346656037);
    const unsigned char*p=(const unsigned char*)C;
    for(size_t i=0;i<(size_t)n*n*sizeof(MATRIX_TYPE);++i)
        digest=(digest^p[i])*UINT64_C(1099511628211);
    printf("n=%d checksum=%016llx correctness=ok\n",n,(unsigned long long)digest);
    free(A);free(B);free(C);
}

int main(){run(2048,false);}

CompilationN/AN/ACompile OKScore: N/A

Testcase #1258.426 ms48 MB + 164 KBAcceptedScore: 100


Judge Duck Online | 评测鸭在线
Server Time: 2026-09-24 03:58:55 | Loaded in 1 ms | Server Status
个人娱乐项目,仅供学习交流使用 | 捐赠