#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
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 176.929 ms | 20 MB + 104 KB | Accepted | Score: 100 | 显示更多 |