#include <immintrin.h>
enum { N = 1024 };
__attribute__((target("avx2,fma"), always_inline))
static inline void kernel6x8(const double *A, const double *B, double *C,
int i, int j) {
__m256d c00 = _mm256_setzero_pd(), c01 = _mm256_setzero_pd();
__m256d c10 = c00, c11 = c00, c20 = c00, c21 = c00;
__m256d c30 = c00, c31 = c00, c40 = c00, c41 = c00;
__m256d c50 = c00, c51 = 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;
const double *a4 = A + (i + 4) * N;
const double *a5 = A + (i + 5) * N;
const double *bp = B + j;
for (int k = 0; k < N; ++k, bp += N) {
__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);
av = _mm256_broadcast_sd(a4 + k);
c40 = _mm256_fmadd_pd(av, b0, c40);
c41 = _mm256_fmadd_pd(av, b1, c41);
av = _mm256_broadcast_sd(a5 + k);
c50 = _mm256_fmadd_pd(av, b0, c50);
c51 = _mm256_fmadd_pd(av, b1, c51);
}
_mm256_store_pd(C + (i + 0) * N + j, c00);
_mm256_store_pd(C + (i + 0) * N + j + 4, c01);
_mm256_store_pd(C + (i + 1) * N + j, c10);
_mm256_store_pd(C + (i + 1) * N + j + 4, c11);
_mm256_store_pd(C + (i + 2) * N + j, c20);
_mm256_store_pd(C + (i + 2) * N + j + 4, c21);
_mm256_store_pd(C + (i + 3) * N + j, c30);
_mm256_store_pd(C + (i + 3) * N + j + 4, c31);
_mm256_store_pd(C + (i + 4) * N + j, c40);
_mm256_store_pd(C + (i + 4) * N + j + 4, c41);
_mm256_store_pd(C + (i + 5) * N + j, c50);
_mm256_store_pd(C + (i + 5) * N + j + 4, c51);
}
__attribute__((target("avx2,fma"), always_inline))
static inline void kernel4x8(const double *A, const double *B, double *C,
int i, int j) {
__m256d c00 = _mm256_setzero_pd(), c01 = _mm256_setzero_pd();
__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;
const double *bp = B + j;
for (int k = 0; k < N; ++k, bp += N) {
__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);
}
_mm256_store_pd(C + (i + 0) * N + j, c00);
_mm256_store_pd(C + (i + 0) * N + j + 4, c01);
_mm256_store_pd(C + (i + 1) * N + j, c10);
_mm256_store_pd(C + (i + 1) * N + j + 4, c11);
_mm256_store_pd(C + (i + 2) * N + j, c20);
_mm256_store_pd(C + (i + 2) * N + j + 4, c21);
_mm256_store_pd(C + (i + 3) * N + j, c30);
_mm256_store_pd(C + (i + 3) * N + j + 4, c31);
}
__attribute__((target("avx2,fma")))
void matrix_multiply(int n, const double *A, const double *B, double *C) {
(void)n;
int i;
for (i = 0; i + 6 <= N; i += 6)
for (int j = 0; j < N; j += 8)
kernel6x8(A, B, C, i, j);
for (int j = 0; j < N; j += 8)
kernel4x8(A, B, C, i, j);
_mm256_zeroupper();
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 298.798 ms | 8 MB + 8 KB | Accepted | Score: 100 | 显示更多 |