#define MATRIX_TYPE long long
#define MATRIX_KIND 3
// 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=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 c0=loadc(c),c1=loadc(c+stride),c2=loadc(c+2*stride),c3=loadc(c+3*stride);
for(int k=0;k<KC;++k) {
V bv=load(b);
c0=madd(broadcast(a[0]),bv,c0);
c1=madd(broadcast(a[1]),bv,c1);
c2=madd(broadcast(a[2]),bv,c2);
c3=madd(broadcast(a[3]),bv,c3);
a+=MR;b+=NR;
}
storec(c,c0);storec(c+stride,c1);storec(c+2*stride,c2);storec(c+3*stride,c3);
}
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)108<<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){uint64_t bits=(uint64_t)next_random(state); memcpy(a+i,&bits,sizeof(bits));}
}
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(4096,false);}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 10.324 s | 512 MB + 168 KB | Accepted | Score: 100 | 显示更多 |