// DUCK_PUBLIC_SURROGATE_BENCHMARK
// Analytic public matrices, independent of production data and all generators.
#pragma GCC optimize("O3")
#pragma GCC target("avx2,fma")
#include <immintrin.h>
#include <stdint.h>
#include <stddef.h>
#include <stdlib.h>
#include <stdio.h>
#include <string.h>
#include <math.h>
static constexpr int N=8192, ROWS=384, COLS=8160, NC=96, KC=256, MC=48;
alignas(4096) static float bp[NC*KC];
static uint64_t tick(){uint32_t lo,hi;asm volatile("lfence;rdtsc;lfence":"=a"(lo),"=d"(hi)::"memory");return ((uint64_t)hi<<32)|lo;}
namespace m6n16u2 {
constexpr int MR=6,NR=16;
static inline void micro(const float*a,const float*b,float*c,int stride){
__m256 c00=_mm256_loadu_ps(c+0*stride+0);
__m256 c01=_mm256_loadu_ps(c+0*stride+8);
__m256 c10=_mm256_loadu_ps(c+1*stride+0);
__m256 c11=_mm256_loadu_ps(c+1*stride+8);
__m256 c20=_mm256_loadu_ps(c+2*stride+0);
__m256 c21=_mm256_loadu_ps(c+2*stride+8);
__m256 c30=_mm256_loadu_ps(c+3*stride+0);
__m256 c31=_mm256_loadu_ps(c+3*stride+8);
__m256 c40=_mm256_loadu_ps(c+4*stride+0);
__m256 c41=_mm256_loadu_ps(c+4*stride+8);
__m256 c50=_mm256_loadu_ps(c+5*stride+0);
__m256 c51=_mm256_loadu_ps(c+5*stride+8);
#pragma GCC unroll 2
for(int k=0;k<KC;++k){
__m256 b0=_mm256_load_ps(b+0);
__m256 b1=_mm256_load_ps(b+8);
{ __m256 a0=_mm256_broadcast_ss(a+0);c00=_mm256_fmadd_ps(a0,b0,c00);c01=_mm256_fmadd_ps(a0,b1,c01);}
{ __m256 a1=_mm256_broadcast_ss(a+1);c10=_mm256_fmadd_ps(a1,b0,c10);c11=_mm256_fmadd_ps(a1,b1,c11);}
{ __m256 a2=_mm256_broadcast_ss(a+2);c20=_mm256_fmadd_ps(a2,b0,c20);c21=_mm256_fmadd_ps(a2,b1,c21);}
{ __m256 a3=_mm256_broadcast_ss(a+3);c30=_mm256_fmadd_ps(a3,b0,c30);c31=_mm256_fmadd_ps(a3,b1,c31);}
{ __m256 a4=_mm256_broadcast_ss(a+4);c40=_mm256_fmadd_ps(a4,b0,c40);c41=_mm256_fmadd_ps(a4,b1,c41);}
{ __m256 a5=_mm256_broadcast_ss(a+5);c50=_mm256_fmadd_ps(a5,b0,c50);c51=_mm256_fmadd_ps(a5,b1,c51);}
a+=MR;b+=NR;
}
_mm256_storeu_ps(c+0*stride+0,c00);
_mm256_storeu_ps(c+0*stride+8,c01);
_mm256_storeu_ps(c+1*stride+0,c10);
_mm256_storeu_ps(c+1*stride+8,c11);
_mm256_storeu_ps(c+2*stride+0,c20);
_mm256_storeu_ps(c+2*stride+8,c21);
_mm256_storeu_ps(c+3*stride+0,c30);
_mm256_storeu_ps(c+3*stride+8,c31);
_mm256_storeu_ps(c+4*stride+0,c40);
_mm256_storeu_ps(c+4*stride+8,c41);
_mm256_storeu_ps(c+5*stride+0,c50);
_mm256_storeu_ps(c+5*stride+8,c51);
}
static void run(const float*A,const float*B,float*C){
memset(C,0,(size_t)ROWS*N*4);
float*ap=nullptr; if(posix_memalign((void**)&ap,4096,(size_t)ROWS*N*4))abort();
constexpr int groups=ROWS/MR;
for(int pc=0;pc<N;pc+=KC)for(int g=0;g<groups;++g)for(int k=0;k<KC;++k)for(int i=0;i<MR;++i)
ap[(((size_t)(pc/KC)*groups+g)*KC+k)*MR+i]=A[(size_t)(g*MR+i)*N+pc+k];
uint64_t begin=tick(),packing=0;
for(int jc=0;jc<COLS;jc+=NC)for(int pc=0;pc<N;pc+=KC){
uint64_t ps=tick();
for(int q=0;q<NC;q+=NR)for(int k=0;k<KC;++k)for(int j=0;j<NR;j+=8)
_mm256_store_ps(bp+((q/NR*KC+k)*NR)+j,_mm256_loadu_ps(B+(size_t)(pc+k)*N+jc+q+j));
packing+=tick()-ps;
for(int ic=0;ic<ROWS;ic+=MC)for(int q=0;q<NC;q+=NR)for(int r=0;r<MC;r+=MR)
micro(ap+((size_t)(pc/KC)*groups+(ic+r)/MR)*KC*MR,bp+q*KC,C+(size_t)(ic+r)*N+jc+q,N);
}
uint64_t total=tick()-begin;
for(int t=0;t<4;++t){int i=t*91,j=t*1999;double expected=0;
for(int k=0;k<N;++k)expected+=(double)A[(size_t)i*N+k]*B[(size_t)k*N+j];
if(fabs(expected-C[(size_t)i*N+j])>0.1)abort();
}
printf("base=%llu\n",(unsigned long long)(total-packing));
free(ap);
}
}
namespace mem2x48 {
constexpr int MR=2,NR=48;
static inline void micro(const float*a,const float*b,float*c,int stride){
float*c1=c+stride;int k;
asm volatile(
"vmovups 0(%[c0]), %%ymm0\n\t"
"vmovups 32(%[c0]), %%ymm1\n\t"
"vmovups 64(%[c0]), %%ymm2\n\t"
"vmovups 96(%[c0]), %%ymm3\n\t"
"vmovups 128(%[c0]), %%ymm4\n\t"
"vmovups 160(%[c0]), %%ymm5\n\t"
"vmovups 0(%[c1]), %%ymm6\n\t"
"vmovups 32(%[c1]), %%ymm7\n\t"
"vmovups 64(%[c1]), %%ymm8\n\t"
"vmovups 96(%[c1]), %%ymm9\n\t"
"vmovups 128(%[c1]), %%ymm10\n\t"
"vmovups 160(%[c1]), %%ymm11\n\t"
"mov $256, %[k]\n\t"
".p2align 5\n\t"
"1:\n\t"
"vbroadcastss 0(%[a]), %%ymm12\n\t"
"vbroadcastss 4(%[a]), %%ymm13\n\t"
"vfmadd231ps 0(%[b]), %%ymm12, %%ymm0\n\t"
"vfmadd231ps 0(%[b]), %%ymm13, %%ymm6\n\t"
"vfmadd231ps 32(%[b]), %%ymm12, %%ymm1\n\t"
"vfmadd231ps 32(%[b]), %%ymm13, %%ymm7\n\t"
"vfmadd231ps 64(%[b]), %%ymm12, %%ymm2\n\t"
"vfmadd231ps 64(%[b]), %%ymm13, %%ymm8\n\t"
"vfmadd231ps 96(%[b]), %%ymm12, %%ymm3\n\t"
"vfmadd231ps 96(%[b]), %%ymm13, %%ymm9\n\t"
"vfmadd231ps 128(%[b]), %%ymm12, %%ymm4\n\t"
"vfmadd231ps 128(%[b]), %%ymm13, %%ymm10\n\t"
"vfmadd231ps 160(%[b]), %%ymm12, %%ymm5\n\t"
"vfmadd231ps 160(%[b]), %%ymm13, %%ymm11\n\t"
"add $8, %[a]\n\t"
"add $192, %[b]\n\t"
"dec %[k]\n\t"
"jnz 1b\n\t"
"vmovups %%ymm0, 0(%[c0])\n\t"
"vmovups %%ymm1, 32(%[c0])\n\t"
"vmovups %%ymm2, 64(%[c0])\n\t"
"vmovups %%ymm3, 96(%[c0])\n\t"
"vmovups %%ymm4, 128(%[c0])\n\t"
"vmovups %%ymm5, 160(%[c0])\n\t"
"vmovups %%ymm6, 0(%[c1])\n\t"
"vmovups %%ymm7, 32(%[c1])\n\t"
"vmovups %%ymm8, 64(%[c1])\n\t"
"vmovups %%ymm9, 96(%[c1])\n\t"
"vmovups %%ymm10, 128(%[c1])\n\t"
"vmovups %%ymm11, 160(%[c1])\n\t"
: [a]"+&r"(a),[b]"+&r"(b),[k]"=&r"(k)
: [c0]"r"(c),[c1]"r"(c1)
: "memory","cc","ymm0","ymm1","ymm2","ymm3","ymm4","ymm5","ymm6","ymm7","ymm8","ymm9","ymm10","ymm11","ymm12","ymm13");
}
static void run(const float*A,const float*B,float*C){
memset(C,0,(size_t)ROWS*N*4);
float*ap=nullptr; if(posix_memalign((void**)&ap,4096,(size_t)ROWS*N*4))abort();
constexpr int groups=ROWS/MR;
for(int pc=0;pc<N;pc+=KC)for(int g=0;g<groups;++g)for(int k=0;k<KC;++k)for(int i=0;i<MR;++i)
ap[(((size_t)(pc/KC)*groups+g)*KC+k)*MR+i]=A[(size_t)(g*MR+i)*N+pc+k];
uint64_t begin=tick(),packing=0;
for(int jc=0;jc<COLS;jc+=NC)for(int pc=0;pc<N;pc+=KC){
uint64_t ps=tick();
for(int q=0;q<NC;q+=NR)for(int k=0;k<KC;++k)for(int j=0;j<NR;j+=8)
_mm256_store_ps(bp+((q/NR*KC+k)*NR)+j,_mm256_loadu_ps(B+(size_t)(pc+k)*N+jc+q+j));
packing+=tick()-ps;
for(int ic=0;ic<ROWS;ic+=MC)for(int q=0;q<NC;q+=NR)for(int r=0;r<MC;r+=MR)
micro(ap+((size_t)(pc/KC)*groups+(ic+r)/MR)*KC*MR,bp+q*KC,C+(size_t)(ic+r)*N+jc+q,N);
}
uint64_t total=tick()-begin;
for(int t=0;t<4;++t){int i=t*91,j=t*1999;double expected=0;
for(int k=0;k<N;++k)expected+=(double)A[(size_t)i*N+k]*B[(size_t)k*N+j];
if(fabs(expected-C[(size_t)i*N+j])>0.1)abort();
}
printf("mem=%llu\n",(unsigned long long)(total-packing));
free(ap);
}
}
int main(){
float *A=nullptr,*B=nullptr,*C=nullptr;
if(posix_memalign((void**)&A,4096,(size_t)ROWS*N*4)||posix_memalign((void**)&B,4096,(size_t)N*N*4)||posix_memalign((void**)&C,4096,(size_t)ROWS*N*4))abort();
for(size_t i=0;i<(size_t)ROWS*N;++i)A[i]=(float)((int)((i*17+43)%257)-128)*(1.0f/128);
for(size_t i=0;i<(size_t)N*N;++i)B[i]=(float)((int)((i*23+71)%263)-131)*(1.0f/131);
m6n16u2::run(A,B,C);
mem2x48::run(A,B,C);
free(C);free(B);free(A);return 0;
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 1.784 s | 292 MB + 136 KB | Accepted | Score: 100 | 显示更多 |