提交记录 86355


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 mmmd2k. 测测你的双精度矩阵乘法-2k Accepted 100 399.429 ms 33576 KB C++17 3.61 KB
提交时间 评测时间
2026-09-24 23:26:08 2026-09-24 23:26:11
#include <string.h>
#include <immintrin.h>

#ifndef KC
#define KC 256
#endif
#ifndef NC
#define NC 384
#endif
#ifndef MC
#define MC 16
#endif
#define MR 4
#define NR 8

static double Abuf[MC*KC+64];
static double Bbuf[NC*KC+64];

__attribute__((target("avx2,fma")))
static inline void mkernel4x8(int kc, const double*ap, const double*bp, double*c, int ldc){
  __m256d c0=_mm256_setzero_pd(),c1=_mm256_setzero_pd(),c2=_mm256_setzero_pd(),c3=_mm256_setzero_pd();
  __m256d c4=_mm256_setzero_pd(),c5=_mm256_setzero_pd(),c6=_mm256_setzero_pd(),c7=_mm256_setzero_pd();
  int k=0;
  for(;k+4<=kc;k+=4){
    __m256d b0,b1,a0,a1,a2,a3;
#define STEP(OFF) \
    b0=_mm256_loadu_pd(bp+(k+OFF)*NR); b1=_mm256_loadu_pd(bp+(k+OFF)*NR+4); \
    a0=_mm256_broadcast_sd(ap+k+OFF); a1=_mm256_broadcast_sd(ap+KC+k+OFF); \
    a2=_mm256_broadcast_sd(ap+2*KC+k+OFF); a3=_mm256_broadcast_sd(ap+3*KC+k+OFF); \
    c0=_mm256_fmadd_pd(a0,b0,c0); c1=_mm256_fmadd_pd(a0,b1,c1); \
    c2=_mm256_fmadd_pd(a1,b0,c2); c3=_mm256_fmadd_pd(a1,b1,c3); \
    c4=_mm256_fmadd_pd(a2,b0,c4); c5=_mm256_fmadd_pd(a2,b1,c5); \
    c6=_mm256_fmadd_pd(a3,b0,c6); c7=_mm256_fmadd_pd(a3,b1,c7);
    STEP(0) STEP(1) STEP(2) STEP(3)
#undef STEP
  }
  for(;k<kc;k++){
    __m256d b0=_mm256_loadu_pd(bp+k*NR), b1=_mm256_loadu_pd(bp+k*NR+4);
    __m256d a0=_mm256_broadcast_sd(ap+k), a1=_mm256_broadcast_sd(ap+KC+k);
    __m256d a2=_mm256_broadcast_sd(ap+2*KC+k), a3=_mm256_broadcast_sd(ap+3*KC+k);
    c0=_mm256_fmadd_pd(a0,b0,c0); c1=_mm256_fmadd_pd(a0,b1,c1);
    c2=_mm256_fmadd_pd(a1,b0,c2); c3=_mm256_fmadd_pd(a1,b1,c3);
    c4=_mm256_fmadd_pd(a2,b0,c4); c5=_mm256_fmadd_pd(a2,b1,c5);
    c6=_mm256_fmadd_pd(a3,b0,c6); c7=_mm256_fmadd_pd(a3,b1,c7);
  }
  _mm256_storeu_pd(c+0*ldc, _mm256_add_pd(_mm256_loadu_pd(c+0*ldc), c0));
  _mm256_storeu_pd(c+0*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+0*ldc+4), c1));
  _mm256_storeu_pd(c+1*ldc, _mm256_add_pd(_mm256_loadu_pd(c+1*ldc), c2));
  _mm256_storeu_pd(c+1*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+1*ldc+4), c3));
  _mm256_storeu_pd(c+2*ldc, _mm256_add_pd(_mm256_loadu_pd(c+2*ldc), c4));
  _mm256_storeu_pd(c+2*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+2*ldc+4), c5));
  _mm256_storeu_pd(c+3*ldc, _mm256_add_pd(_mm256_loadu_pd(c+3*ldc), c6));
  _mm256_storeu_pd(c+3*ldc+4, _mm256_add_pd(_mm256_loadu_pd(c+3*ldc+4), c7));
}

__attribute__((target("avx2,fma")))
void matrix_multiply(int n, const double *A, const double *B, double *C){
  for(int jc=0; jc<n; jc+=NC){
    int nb=n-jc; if(nb>NC) nb=NC;
    int nrb=(nb+NR-1)/NR;
    for(int kk=0; kk<n; kk+=KC){
      int kb=n-kk; if(kb>KC) kb=KC;
      for(int jb=0; jb<nrb; jb++){
        int j0=jc+jb*NR;
        double* bp = Bbuf + (size_t)jb*KC*NR;
        for(int k=0;k<kb;k++){
          const double* src = B + (size_t)(kk+k)*n + j0;
          double* dst = bp + k*NR;
          dst[0]=src[0];dst[1]=src[1];dst[2]=src[2];dst[3]=src[3];
          dst[4]=src[4];dst[5]=src[5];dst[6]=src[6];dst[7]=src[7];
        }
      }
      for(int ii=0; ii<n; ii+=MC){
        int mb=n-ii; if(mb>MC) mb=MC;
        for(int i=0;i<mb;i++)
          memcpy(Abuf+(size_t)i*KC, A+(size_t)(ii+i)*n+kk, (size_t)kb*sizeof(double));
        int im=mb-((mb/MR)*MR);
        for(int jb=0; jb<nrb; jb++){
          double* bp = Bbuf + (size_t)jb*KC*NR;
          double* cp = C + (size_t)ii*n + jc + jb*NR;
          int i=0;
          for(; i<mb-im; i+=MR) mkernel4x8(kb, Abuf+(size_t)i*KC, bp, cp+(size_t)i*n, n);
          for(; i<mb; i++){
            for(int k=0;k<kb;k++){
              double a=Abuf[(size_t)i*KC+k];
              double* c=cp+(size_t)i*n;
              for(int j=0;j<NR;j++) c[j]+=a*bp[k*NR+j];
            }
          }
        }
      }
    }
  }
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #1399.429 ms32 MB + 808 KBAcceptedScore: 100


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