#include <pthread.h>
#define WORKER_ONLY
#include <immintrin.h>
enum { N = 1024, MAIN_COLS = 1008, PANELS = 84 };
#ifndef PACK_PANELS
#define PACK_PANELS 4
#endif
#ifndef PACK_ROWS
#define PACK_ROWS 32
#endif
static double packed_b[N * N] __attribute__((aligned(4096)));
__attribute__((target("avx2,fma"), always_inline))
static inline void kernel4x12(const double *A, const double *bp, double *C,
int i, int j) {
__m256d c00 = _mm256_setzero_pd(), c01 = c00, c02 = c00;
__m256d c10 = c00, c11 = c00, c12 = c00;
__m256d c20 = c00, c21 = c00, c22 = c00;
__m256d c30 = c00, c31 = c00, c32 = c00;
const double *a0 = A + (i + 0) * N;
const double *a1 = A + (i + 1) * N;
const double *a2 = A + (i + 2) * N;
const double *a3 = A + (i + 3) * N;
for (int k = 0; k < N; ++k, bp += 12) {
__m256d b0 = _mm256_load_pd(bp);
__m256d b1 = _mm256_load_pd(bp + 4);
__m256d b2 = _mm256_load_pd(bp + 8);
__m256d av = _mm256_broadcast_sd(a0 + k);
c00 = _mm256_fmadd_pd(av, b0, c00);
c01 = _mm256_fmadd_pd(av, b1, c01);
c02 = _mm256_fmadd_pd(av, b2, c02);
av = _mm256_broadcast_sd(a1 + k);
c10 = _mm256_fmadd_pd(av, b0, c10);
c11 = _mm256_fmadd_pd(av, b1, c11);
c12 = _mm256_fmadd_pd(av, b2, c12);
av = _mm256_broadcast_sd(a2 + k);
c20 = _mm256_fmadd_pd(av, b0, c20);
c21 = _mm256_fmadd_pd(av, b1, c21);
c22 = _mm256_fmadd_pd(av, b2, c22);
av = _mm256_broadcast_sd(a3 + k);
c30 = _mm256_fmadd_pd(av, b0, c30);
c31 = _mm256_fmadd_pd(av, b1, c31);
c32 = _mm256_fmadd_pd(av, b2, c32);
}
double *d0 = C + (i + 0) * N + j;
double *d1 = C + (i + 1) * N + j;
double *d2 = C + (i + 2) * N + j;
double *d3 = C + (i + 3) * N + j;
_mm256_store_pd(d0, c00); _mm256_store_pd(d0 + 4, c01);
_mm256_store_pd(d0 + 8, c02);
_mm256_store_pd(d1, c10); _mm256_store_pd(d1 + 4, c11);
_mm256_store_pd(d1 + 8, c12);
_mm256_store_pd(d2, c20); _mm256_store_pd(d2 + 4, c21);
_mm256_store_pd(d2 + 8, c22);
_mm256_store_pd(d3, c30); _mm256_store_pd(d3 + 4, c31);
_mm256_store_pd(d3 + 8, c32);
}
__attribute__((target("avx2,fma"), always_inline))
static inline void kernel4x8_tail(const double *A, const double *bp,
double *C, int i, int j) {
__m256d c00 = _mm256_setzero_pd(), c01 = c00;
__m256d c10 = c00, c11 = c00, c20 = c00, c21 = c00;
__m256d c30 = c00, c31 = c00;
const double *a0 = A + (i + 0) * N;
const double *a1 = A + (i + 1) * N;
const double *a2 = A + (i + 2) * N;
const double *a3 = A + (i + 3) * N;
for (int k = 0; k < N; ++k, bp += 16) {
__m256d b0 = _mm256_load_pd(bp);
__m256d b1 = _mm256_load_pd(bp + 4);
__m256d av = _mm256_broadcast_sd(a0 + k);
c00 = _mm256_fmadd_pd(av, b0, c00);
c01 = _mm256_fmadd_pd(av, b1, c01);
av = _mm256_broadcast_sd(a1 + k);
c10 = _mm256_fmadd_pd(av, b0, c10);
c11 = _mm256_fmadd_pd(av, b1, c11);
av = _mm256_broadcast_sd(a2 + k);
c20 = _mm256_fmadd_pd(av, b0, c20);
c21 = _mm256_fmadd_pd(av, b1, c21);
av = _mm256_broadcast_sd(a3 + k);
c30 = _mm256_fmadd_pd(av, b0, c30);
c31 = _mm256_fmadd_pd(av, b1, c31);
}
double *d0 = C + (i + 0) * N + j;
double *d1 = C + (i + 1) * N + j;
double *d2 = C + (i + 2) * N + j;
double *d3 = C + (i + 3) * N + j;
_mm256_store_pd(d0, c00); _mm256_store_pd(d0 + 4, c01);
_mm256_store_pd(d1, c10); _mm256_store_pd(d1 + 4, c11);
_mm256_store_pd(d2, c20); _mm256_store_pd(d2 + 4, c21);
_mm256_store_pd(d3, c30); _mm256_store_pd(d3 + 4, c31);
}
#ifndef WORKER_ONLY
__attribute__((target("avx2,fma")))
void matrix_multiply(int n, const double *A, const double *B, double *C) {
(void)n;
double *dst;
for (int pb = 0; pb < PANELS; pb += PACK_PANELS)
for (int kb = 0; kb < N; kb += PACK_ROWS) {
int pend = pb + PACK_PANELS; if (pend > PANELS) pend = PANELS;
for (int p = pb; p < pend; ++p) {
int j = p * 12;
dst = packed_b + p * (N * 12) + kb * 12;
const double *src = B + kb * N + j;
for (int k = 0; k < PACK_ROWS; ++k, dst += 12, src += N) {
_mm256_store_pd(dst, _mm256_load_pd(src));
_mm256_store_pd(dst + 4, _mm256_load_pd(src + 4));
_mm256_store_pd(dst + 8, _mm256_load_pd(src + 8));
}
}
}
dst = packed_b + PANELS * (N * 12);
for (int k = 0; k < N; ++k, dst += 16) {
_mm256_store_pd(dst, _mm256_load_pd(B + k * N + MAIN_COLS));
_mm256_store_pd(dst + 4, _mm256_load_pd(B + k * N + MAIN_COLS + 4));
_mm256_store_pd(dst + 8, _mm256_load_pd(B + k * N + MAIN_COLS + 8));
_mm256_store_pd(dst + 12, _mm256_load_pd(B + k * N + MAIN_COLS + 12));
}
#ifdef PACK_ONLY
return;
#endif
for (int p = 0; p < PANELS; ++p) {
const double *panel = packed_b + p * (N * 12);
int j = p * 12;
for (int i = 0; i < N; i += 4)
kernel4x12(A, panel, C, i, j);
}
const double *tail = packed_b + PANELS * (N * 12);
for (int i = 0; i < N; i += 4) {
kernel4x8_tail(A, tail, C, i, MAIN_COLS);
kernel4x8_tail(A, tail + 8, C, i, MAIN_COLS + 8);
}
_mm256_zeroupper();
}
#endif
typedef struct {
const double *A, *B;
double *C;
int first, last, tail;
} Work;
__attribute__((target("avx2,fma")))
static void *run(void *opaque) {
Work *w = (Work *)opaque;
for (int p = w->first; p < w->last; ++p) {
int j = p * 12;
double *dst = packed_b + p * (N * 12);
for (int k = 0; k < N; ++k, dst += 12) {
_mm256_store_pd(dst, _mm256_load_pd(w->B + k * N + j));
_mm256_store_pd(dst + 4, _mm256_load_pd(w->B + k * N + j + 4));
_mm256_store_pd(dst + 8, _mm256_load_pd(w->B + k * N + j + 8));
}
const double *panel = packed_b + p * (N * 12);
for (int i = 0; i < N; i += 4)
kernel4x12(w->A, panel, w->C, i, j);
}
if (w->tail) {
double *dst = packed_b + PANELS * (N * 12);
for (int k = 0; k < N; ++k, dst += 16) {
_mm256_store_pd(dst, _mm256_load_pd(w->B + k*N + MAIN_COLS));
_mm256_store_pd(dst+4, _mm256_load_pd(w->B + k*N + MAIN_COLS+4));
_mm256_store_pd(dst+8, _mm256_load_pd(w->B + k*N + MAIN_COLS+8));
_mm256_store_pd(dst+12,_mm256_load_pd(w->B + k*N + MAIN_COLS+12));
}
const double *tail = packed_b + PANELS * (N * 12);
for (int i = 0; i < N; i += 4) {
kernel4x8_tail(w->A, tail, w->C, i, MAIN_COLS);
kernel4x8_tail(w->A, tail + 8, w->C, i, MAIN_COLS + 8);
}
}
_mm256_zeroupper();
return 0;
}
void matrix_multiply(int n, const double *A, const double *B, double *C) {
(void)n;
Work left = {A,B,C,0,43,0}, right = {A,B,C,43,PANELS,1};
pthread_t thread;
int created = pthread_create(&thread, 0, run, &right) == 0;
run(&left);
if (created) pthread_join(thread, 0);
else run(&right);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 50.591 ms | 16 MB + 16 KB | Accepted | Score: 100 | 显示更多 |