提交记录 125725


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_codex_maker test. 自定义测试 Accepted 100 12.845 s 917744 KB C++17 4.85 KB
提交时间 评测时间
2026-10-06 01:31:24 2026-10-06 01:31:40
// Classical float matrix multiplication, AVX2/FMA only.
#pragma GCC optimize("O3")
#pragma GCC target("avx2,fma")
#include <immintrin.h>
#include <stddef.h>
#include <string.h>
#include <stdlib.h>
namespace classic {
constexpr int MR=6,NR=16,MC=48,NC=128,KC=256;
alignas(4096) static float bp[NC*KC];
template<int R> static inline void micro(const float*a,const float*b,float*c,int stride){
    __m256 c00,c01;if(R>0){c00=_mm256_loadu_ps(c+0*stride);c01=_mm256_loadu_ps(c+0*stride+8);}
    __m256 c10,c11;if(R>1){c10=_mm256_loadu_ps(c+1*stride);c11=_mm256_loadu_ps(c+1*stride+8);}
    __m256 c20,c21;if(R>2){c20=_mm256_loadu_ps(c+2*stride);c21=_mm256_loadu_ps(c+2*stride+8);}
    __m256 c30,c31;if(R>3){c30=_mm256_loadu_ps(c+3*stride);c31=_mm256_loadu_ps(c+3*stride+8);}
    __m256 c40,c41;if(R>4){c40=_mm256_loadu_ps(c+4*stride);c41=_mm256_loadu_ps(c+4*stride+8);}
    __m256 c50,c51;if(R>5){c50=_mm256_loadu_ps(c+5*stride);c51=_mm256_loadu_ps(c+5*stride+8);}
#pragma GCC unroll 4
    for(int k=0;k<KC;++k){
        __m256 b0=_mm256_load_ps(b),b1=_mm256_load_ps(b+8);
        if(R>0){__m256 v=_mm256_broadcast_ss(a+0);c00=_mm256_fmadd_ps(v,b0,c00);c01=_mm256_fmadd_ps(v,b1,c01);}
        if(R>1){__m256 v=_mm256_broadcast_ss(a+1);c10=_mm256_fmadd_ps(v,b0,c10);c11=_mm256_fmadd_ps(v,b1,c11);}
        if(R>2){__m256 v=_mm256_broadcast_ss(a+2);c20=_mm256_fmadd_ps(v,b0,c20);c21=_mm256_fmadd_ps(v,b1,c21);}
        if(R>3){__m256 v=_mm256_broadcast_ss(a+3);c30=_mm256_fmadd_ps(v,b0,c30);c31=_mm256_fmadd_ps(v,b1,c31);}
        if(R>4){__m256 v=_mm256_broadcast_ss(a+4);c40=_mm256_fmadd_ps(v,b0,c40);c41=_mm256_fmadd_ps(v,b1,c41);}
        if(R>5){__m256 v=_mm256_broadcast_ss(a+5);c50=_mm256_fmadd_ps(v,b0,c50);c51=_mm256_fmadd_ps(v,b1,c51);}
        a+=MR;b+=NR;
    }
    if(R>0){_mm256_storeu_ps(c+0*stride,c00);_mm256_storeu_ps(c+0*stride+8,c01);}
    if(R>1){_mm256_storeu_ps(c+1*stride,c10);_mm256_storeu_ps(c+1*stride+8,c11);}
    if(R>2){_mm256_storeu_ps(c+2*stride,c20);_mm256_storeu_ps(c+2*stride+8,c21);}
    if(R>3){_mm256_storeu_ps(c+3*stride,c30);_mm256_storeu_ps(c+3*stride+8,c31);}
    if(R>4){_mm256_storeu_ps(c+4*stride,c40);_mm256_storeu_ps(c+4*stride+8,c41);}
    if(R>5){_mm256_storeu_ps(c+5*stride,c50);_mm256_storeu_ps(c+5*stride+8,c51);}
}
static void multiply_rows(int n,int row_count,const float*A,const float*B,float*C){
    memset(C,0,(size_t)row_count*n*sizeof(float));
    if(n%NC||n%KC){
        for(int i=0;i<row_count;++i)for(int k=0;k<n;++k)for(int j=0;j<n;++j)
            C[(size_t)i*n+j]+=A[(size_t)i*n+k]*B[(size_t)k*n+j];
        return;
    }
    const int groups=(row_count+MR-1)/MR;
    float*ap=nullptr;
    if(posix_memalign((void**)&ap,4096,(size_t)(n/KC)*groups*KC*MR*sizeof(float)))abort();
    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]=g*MR+i<row_count?A[(size_t)(g*MR+i)*n+pc+k]:0;
    for(int jc=0;jc<n;jc+=NC)for(int pc=0;pc<n;pc+=KC){
        for(int q=0;q<NC;q+=NR)for(int k=0;k<KC;++k)
            _mm256_store_ps(bp+((q/NR*KC+k)*NR),_mm256_loadu_ps(B+(size_t)(pc+k)*n+jc+q)),
            _mm256_store_ps(bp+((q/NR*KC+k)*NR+8),_mm256_loadu_ps(B+(size_t)(pc+k)*n+jc+q+8));
        for(int ic=0;ic<row_count;ic+=MC){
            int rows=row_count-ic<MC?row_count-ic:MC;
            for(int q=0;q<NC;q+=NR)for(int r=0;r<rows;r+=MR){
                const float*a=ap+((size_t)(pc/KC)*groups+(ic+r)/MR)*KC*MR,*b=bp+q*KC;float*c=C+(size_t)(ic+r)*n+jc+q;
                if(rows-r>=6)micro<6>(a,b,c,n);
                else if(rows-r==4)micro<4>(a,b,c,n);
                else if(rows-r==2)micro<2>(a,b,c,n);
                else {for(int i=0;i<rows-r;++i)for(int j=0;j<NR;++j){float v=c[(size_t)i*n+j];for(int k=0;k<KC;++k)v+=a[k*MR+i]*b[k*NR+j];c[(size_t)i*n+j]=v;}}
            }
        }
    }
    free(ap);
}
}
void matrix_multiply(int n,const float*A,const float*B,float*C){
    for(int row=0;row<n;row+=4096){int rows=n-row<4096?n-row:4096;
        classic::multiply_rows(n,rows,A+(size_t)row*n,B,C+(size_t)row*n);
    }
}

// DUCK_PUBLIC_SURROGATE_BENCHMARK
// Public analytic matrices unrelated to production input; no random seed.
#include <stdlib.h>
#include <stdio.h>
#include <math.h>
int main(){
 constexpr int n=8192;constexpr size_t z=(size_t)n*n;
 float*A=nullptr,*B=nullptr,*C=nullptr;
 if(posix_memalign((void**)&A,4096,z*4)||posix_memalign((void**)&B,4096,z*4)||posix_memalign((void**)&C,4096,z*4))return 1;
 for(int i=0;i<n;++i)for(int j=0;j<n;++j){
  A[(size_t)i*n+j]=float(((i*23+j*7+19)%1009)-504)/504.f;
  B[(size_t)i*n+j]=float(((i*37+j*11+29)%1013)-506)/506.f;
 }
 matrix_multiply(n,A,B,C);
 double maxerr=0;
 for(int t=0;t<64;++t){int i=(t*127+31)%n,j=(t*337+53)%n;double v=0;for(int k=0;k<n;++k)v+=(double)A[(size_t)i*n+k]*B[(size_t)k*n+j];double e=fabs(v-C[(size_t)i*n+j]);if(e>maxerr)maxerr=e;if(e>0.1)return 2;}
 printf("classical_full_matrix_checked max_error=%g\n",maxerr);return 0;
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #112.845 s896 MB + 240 KBAcceptedScore: 100


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