提交记录 86916


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 mmmf2k. 测测你的单精度矩阵乘法-2k Accepted 100 176.929 ms 20584 KB C++17 14.50 KB
提交时间 评测时间
2026-09-25 01:24:03 2026-09-25 01:24:05
#include <immintrin.h>
#pragma GCC push_options
#pragma GCC target("avx2,fma")
#define MR 6
#define NR 16
#define KC 512
#define MC 48
#define NC 2048
static float Apack[MC*KC+256] __attribute__((aligned(64)));
static float Bpack[NC*KC+256] __attribute__((aligned(64)));
static inline void micro(int kc, const float *Ap, const float *Bp, float *C, int ldc) {
  __m256 c0_0,c0_1,c1_0,c1_1,c2_0,c2_1,c3_0,c3_1,c4_0,c4_1,c5_0,c5_1;
  c0_0 = c0_1 = c1_0 = c1_1 = c2_0 = c2_1 = c3_0 = c3_1 = c4_0 = c4_1 = c5_0 = c5_1 = _mm256_setzero_ps();
  for (int k = 0; k < kc; k++) {
    __m256 b0 = _mm256_loadu_ps(Bp + (size_t)k*NR + 0);
    __m256 b1 = _mm256_loadu_ps(Bp + (size_t)k*NR + 8);
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 0);
      c0_0 = _mm256_fmadd_ps(t, b0, c0_0);
      c0_1 = _mm256_fmadd_ps(t, b1, c0_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 1);
      c1_0 = _mm256_fmadd_ps(t, b0, c1_0);
      c1_1 = _mm256_fmadd_ps(t, b1, c1_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 2);
      c2_0 = _mm256_fmadd_ps(t, b0, c2_0);
      c2_1 = _mm256_fmadd_ps(t, b1, c2_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 3);
      c3_0 = _mm256_fmadd_ps(t, b0, c3_0);
      c3_1 = _mm256_fmadd_ps(t, b1, c3_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 4);
      c4_0 = _mm256_fmadd_ps(t, b0, c4_0);
      c4_1 = _mm256_fmadd_ps(t, b1, c4_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 5);
      c5_0 = _mm256_fmadd_ps(t, b0, c5_0);
      c5_1 = _mm256_fmadd_ps(t, b1, c5_1);
    }
  }
  _mm256_storeu_ps(C + (size_t)0*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 0), c0_0));
  _mm256_storeu_ps(C + (size_t)0*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 8), c0_1));
  _mm256_storeu_ps(C + (size_t)1*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 0), c1_0));
  _mm256_storeu_ps(C + (size_t)1*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 8), c1_1));
  _mm256_storeu_ps(C + (size_t)2*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 0), c2_0));
  _mm256_storeu_ps(C + (size_t)2*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 8), c2_1));
  _mm256_storeu_ps(C + (size_t)3*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)3*ldc + 0), c3_0));
  _mm256_storeu_ps(C + (size_t)3*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)3*ldc + 8), c3_1));
  _mm256_storeu_ps(C + (size_t)4*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)4*ldc + 0), c4_0));
  _mm256_storeu_ps(C + (size_t)4*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)4*ldc + 8), c4_1));
  _mm256_storeu_ps(C + (size_t)5*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)5*ldc + 0), c5_0));
  _mm256_storeu_ps(C + (size_t)5*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)5*ldc + 8), c5_1));
}
static inline void micro_st(int kc, const float *Ap, const float *Bp, float *C, int ldc) {
  __m256 c0_0,c0_1,c1_0,c1_1,c2_0,c2_1,c3_0,c3_1,c4_0,c4_1,c5_0,c5_1;
  c0_0 = c0_1 = c1_0 = c1_1 = c2_0 = c2_1 = c3_0 = c3_1 = c4_0 = c4_1 = c5_0 = c5_1 = _mm256_setzero_ps();
  for (int k = 0; k < kc; k++) {
    __m256 b0 = _mm256_loadu_ps(Bp + (size_t)k*NR + 0);
    __m256 b1 = _mm256_loadu_ps(Bp + (size_t)k*NR + 8);
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 0);
      c0_0 = _mm256_fmadd_ps(t, b0, c0_0);
      c0_1 = _mm256_fmadd_ps(t, b1, c0_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 1);
      c1_0 = _mm256_fmadd_ps(t, b0, c1_0);
      c1_1 = _mm256_fmadd_ps(t, b1, c1_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 2);
      c2_0 = _mm256_fmadd_ps(t, b0, c2_0);
      c2_1 = _mm256_fmadd_ps(t, b1, c2_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 3);
      c3_0 = _mm256_fmadd_ps(t, b0, c3_0);
      c3_1 = _mm256_fmadd_ps(t, b1, c3_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 4);
      c4_0 = _mm256_fmadd_ps(t, b0, c4_0);
      c4_1 = _mm256_fmadd_ps(t, b1, c4_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 5);
      c5_0 = _mm256_fmadd_ps(t, b0, c5_0);
      c5_1 = _mm256_fmadd_ps(t, b1, c5_1);
    }
  }
  _mm256_storeu_ps(C + (size_t)0*ldc + 0, c0_0);
  _mm256_storeu_ps(C + (size_t)0*ldc + 8, c0_1);
  _mm256_storeu_ps(C + (size_t)1*ldc + 0, c1_0);
  _mm256_storeu_ps(C + (size_t)1*ldc + 8, c1_1);
  _mm256_storeu_ps(C + (size_t)2*ldc + 0, c2_0);
  _mm256_storeu_ps(C + (size_t)2*ldc + 8, c2_1);
  _mm256_storeu_ps(C + (size_t)3*ldc + 0, c3_0);
  _mm256_storeu_ps(C + (size_t)3*ldc + 8, c3_1);
  _mm256_storeu_ps(C + (size_t)4*ldc + 0, c4_0);
  _mm256_storeu_ps(C + (size_t)4*ldc + 8, c4_1);
  _mm256_storeu_ps(C + (size_t)5*ldc + 0, c5_0);
  _mm256_storeu_ps(C + (size_t)5*ldc + 8, c5_1);
}
static inline void micro_p(int mr, int kc, const float *Ap, const float *Bp, float *C, int ldc) {
  __m256 c0_0,c0_1,c1_0,c1_1,c2_0,c2_1,c3_0,c3_1,c4_0,c4_1,c5_0,c5_1;
  c0_0 = c0_1 = c1_0 = c1_1 = c2_0 = c2_1 = c3_0 = c3_1 = c4_0 = c4_1 = c5_0 = c5_1 = _mm256_setzero_ps();
  for (int k = 0; k < kc; k++) {
    __m256 b0 = _mm256_loadu_ps(Bp + (size_t)k*NR + 0);
    __m256 b1 = _mm256_loadu_ps(Bp + (size_t)k*NR + 8);
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 0);
      c0_0 = _mm256_fmadd_ps(t, b0, c0_0);
      c0_1 = _mm256_fmadd_ps(t, b1, c0_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 1);
      c1_0 = _mm256_fmadd_ps(t, b0, c1_0);
      c1_1 = _mm256_fmadd_ps(t, b1, c1_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 2);
      c2_0 = _mm256_fmadd_ps(t, b0, c2_0);
      c2_1 = _mm256_fmadd_ps(t, b1, c2_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 3);
      c3_0 = _mm256_fmadd_ps(t, b0, c3_0);
      c3_1 = _mm256_fmadd_ps(t, b1, c3_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 4);
      c4_0 = _mm256_fmadd_ps(t, b0, c4_0);
      c4_1 = _mm256_fmadd_ps(t, b1, c4_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 5);
      c5_0 = _mm256_fmadd_ps(t, b0, c5_0);
      c5_1 = _mm256_fmadd_ps(t, b1, c5_1);
    }
  }
  switch (mr) {
  case 1:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 0), c0_0));
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 8), c0_1));
    break;
  case 2:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 0), c0_0));
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 8), c0_1));
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 0), c1_0));
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 8), c1_1));
    break;
  case 3:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 0), c0_0));
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 8), c0_1));
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 0), c1_0));
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 8), c1_1));
    _mm256_storeu_ps(C + (size_t)2*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 0), c2_0));
    _mm256_storeu_ps(C + (size_t)2*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 8), c2_1));
    break;
  case 4:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 0), c0_0));
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 8), c0_1));
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 0), c1_0));
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 8), c1_1));
    _mm256_storeu_ps(C + (size_t)2*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 0), c2_0));
    _mm256_storeu_ps(C + (size_t)2*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 8), c2_1));
    _mm256_storeu_ps(C + (size_t)3*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)3*ldc + 0), c3_0));
    _mm256_storeu_ps(C + (size_t)3*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)3*ldc + 8), c3_1));
    break;
  case 5:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 0), c0_0));
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)0*ldc + 8), c0_1));
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 0), c1_0));
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)1*ldc + 8), c1_1));
    _mm256_storeu_ps(C + (size_t)2*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 0), c2_0));
    _mm256_storeu_ps(C + (size_t)2*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)2*ldc + 8), c2_1));
    _mm256_storeu_ps(C + (size_t)3*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)3*ldc + 0), c3_0));
    _mm256_storeu_ps(C + (size_t)3*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)3*ldc + 8), c3_1));
    _mm256_storeu_ps(C + (size_t)4*ldc + 0, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)4*ldc + 0), c4_0));
    _mm256_storeu_ps(C + (size_t)4*ldc + 8, _mm256_add_ps(_mm256_loadu_ps(C + (size_t)4*ldc + 8), c4_1));
    break;
  default: break;
  }
}
static inline void micro_pst(int mr, int kc, const float *Ap, const float *Bp, float *C, int ldc) {
  __m256 c0_0,c0_1,c1_0,c1_1,c2_0,c2_1,c3_0,c3_1,c4_0,c4_1,c5_0,c5_1;
  c0_0 = c0_1 = c1_0 = c1_1 = c2_0 = c2_1 = c3_0 = c3_1 = c4_0 = c4_1 = c5_0 = c5_1 = _mm256_setzero_ps();
  for (int k = 0; k < kc; k++) {
    __m256 b0 = _mm256_loadu_ps(Bp + (size_t)k*NR + 0);
    __m256 b1 = _mm256_loadu_ps(Bp + (size_t)k*NR + 8);
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 0);
      c0_0 = _mm256_fmadd_ps(t, b0, c0_0);
      c0_1 = _mm256_fmadd_ps(t, b1, c0_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 1);
      c1_0 = _mm256_fmadd_ps(t, b0, c1_0);
      c1_1 = _mm256_fmadd_ps(t, b1, c1_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 2);
      c2_0 = _mm256_fmadd_ps(t, b0, c2_0);
      c2_1 = _mm256_fmadd_ps(t, b1, c2_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 3);
      c3_0 = _mm256_fmadd_ps(t, b0, c3_0);
      c3_1 = _mm256_fmadd_ps(t, b1, c3_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 4);
      c4_0 = _mm256_fmadd_ps(t, b0, c4_0);
      c4_1 = _mm256_fmadd_ps(t, b1, c4_1);
    }
    { __m256 t = _mm256_broadcast_ss(Ap + (size_t)k*MR + 5);
      c5_0 = _mm256_fmadd_ps(t, b0, c5_0);
      c5_1 = _mm256_fmadd_ps(t, b1, c5_1);
    }
  }
  switch (mr) {
  case 1:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, c0_0);
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, c0_1);
    break;
  case 2:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, c0_0);
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, c0_1);
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, c1_0);
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, c1_1);
    break;
  case 3:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, c0_0);
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, c0_1);
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, c1_0);
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, c1_1);
    _mm256_storeu_ps(C + (size_t)2*ldc + 0, c2_0);
    _mm256_storeu_ps(C + (size_t)2*ldc + 8, c2_1);
    break;
  case 4:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, c0_0);
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, c0_1);
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, c1_0);
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, c1_1);
    _mm256_storeu_ps(C + (size_t)2*ldc + 0, c2_0);
    _mm256_storeu_ps(C + (size_t)2*ldc + 8, c2_1);
    _mm256_storeu_ps(C + (size_t)3*ldc + 0, c3_0);
    _mm256_storeu_ps(C + (size_t)3*ldc + 8, c3_1);
    break;
  case 5:
    _mm256_storeu_ps(C + (size_t)0*ldc + 0, c0_0);
    _mm256_storeu_ps(C + (size_t)0*ldc + 8, c0_1);
    _mm256_storeu_ps(C + (size_t)1*ldc + 0, c1_0);
    _mm256_storeu_ps(C + (size_t)1*ldc + 8, c1_1);
    _mm256_storeu_ps(C + (size_t)2*ldc + 0, c2_0);
    _mm256_storeu_ps(C + (size_t)2*ldc + 8, c2_1);
    _mm256_storeu_ps(C + (size_t)3*ldc + 0, c3_0);
    _mm256_storeu_ps(C + (size_t)3*ldc + 8, c3_1);
    _mm256_storeu_ps(C + (size_t)4*ldc + 0, c4_0);
    _mm256_storeu_ps(C + (size_t)4*ldc + 8, c4_1);
    break;
  default: break;
  }
}

static inline void packB(int jc, int pc, int nc, int kc, int n, const float *B) {
  float *p = Bpack;
  for (int jr = 0; jr < nc; jr += NR) {
    int nr = (nc - jr < NR) ? (nc - jr) : NR;
    const float *src = B + (size_t)pc * n + jc + jr;
    if (nr == NR) {
      for (int k = 0; k < kc; k++) {
        _mm256_storeu_ps(p, _mm256_loadu_ps(src));
        _mm256_storeu_ps(p + 8, _mm256_loadu_ps(src + 8));
        p += NR; src += n;
      }
    } else {
      for (int k = 0; k < kc; k++) {
        for (int r = 0; r < nr; r++) p[r] = src[r];
        for (int r = nr; r < NR; r++) p[r] = 0.f;
        p += NR; src += n;
      }
    }
  }
}

static inline void packA(int ic, int pc, int mc, int kc, int n, const float *A) {
  float *p = Apack;
  for (int ir = 0; ir < mc; ir += MR) {
    int mr = (mc - ir < MR) ? (mc - ir) : MR;
    for (int k = 0; k < kc; k++) {
      const float *src = A + (size_t)(ic + ir) * n + pc + k;
      for (int r = 0; r < mr; r++) *p++ = src[(size_t)r * n];
      for (int r = mr; r < MR; r++) *p++ = 0.f;
    }
  }
}
void matrix_multiply(int n, const float *A, const float *B, float *C) {
  int nbc = n - (n % NR);
  for (int jc = 0; jc < nbc; jc += NC) {
    int nc = (nbc - jc < NC) ? (nbc - jc) : NC;
    for (int pc = 0; pc < n; pc += KC) {
      int kc = (n - pc < KC) ? (n - pc) : KC;
      packB(jc, pc, nc, kc, n, B);
      for (int ic = 0; ic < n; ic += MC) {
        int mc = (n - ic < MC) ? (n - ic) : MC;
        packA(ic, pc, mc, kc, n, A);
        int im = mc % MR;
        int mcf = mc - im;
        { const float *bp = Bpack;
        for (int jr = 0; jr < nc; jr += NR) {
          const float *a = Apack;
          for (int ir = 0; ir < mcf; ir += MR) {
            if (pc == 0) micro_st(kc, a, bp, C + (size_t)(ic+ir)*n + jc + jr, n);
            else micro(kc, a, bp, C + (size_t)(ic+ir)*n + jc + jr, n);
            a += (size_t)kc * MR;
          }
          if (im) {
            if (pc == 0) micro_pst(im, kc, a, bp, C + (size_t)(ic+mcf)*n + jc + jr, n);
            else micro_p(im, kc, a, bp, C + (size_t)(ic+mcf)*n + jc + jr, n);
          }
          bp += (size_t)kc * NR;
        }
        }
      }
    }
  }
  if (nbc != n) {
  for (int i = 0; i < n; i++)
    for (int j = nbc; j < n; j++) { float acc=0; for (int k=0;k<n;k++) acc += A[(size_t)i*n+k]*B[(size_t)k*n+j]; C[(size_t)i*n+j]=acc; }
  }
}
#pragma GCC pop_options

CompilationN/AN/ACompile OKScore: N/A

Testcase #1176.929 ms20 MB + 104 KBAcceptedScore: 100


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