#include <immintrin.h>
#include <stddef.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#pragma GCC target("avx2,fma")
// Tunable via -DMC=... etc. Defaults.
#ifndef MC
#define MC 512
#endif
#ifndef NC
#define NC 128
#endif
#ifndef KC
#define KC 128
#endif
#define MR 4
#define NR 8
static double A_pack[MC * KC] __attribute__((aligned(4096)));
static double B_pack[KC * NC] __attribute__((aligned(4096)));
// A_pack: [(p*KC + k)*MR + r] = A[(ic + p*MR + r)*n + (pc + k)]
// B_pack: [(q*KC + k)*NR + c] = B[(pc + k)*n + (jc + q*NR + c)]
static inline void microkernel(const double *a, const double *b, double *c, int lda) {
__m256d c00 = _mm256_loadu_pd(c + 0*lda);
__m256d c01 = _mm256_loadu_pd(c + 0*lda + 4);
__m256d c10 = _mm256_loadu_pd(c + 1*lda);
__m256d c11 = _mm256_loadu_pd(c + 1*lda + 4);
__m256d c20 = _mm256_loadu_pd(c + 2*lda);
__m256d c21 = _mm256_loadu_pd(c + 2*lda + 4);
__m256d c30 = _mm256_loadu_pd(c + 3*lda);
__m256d c31 = _mm256_loadu_pd(c + 3*lda + 4);
const double *ap = a;
const double *bp = b;
#pragma GCC unroll 4
for (int k = 0; k < KC; k++) {
__m256d b0 = _mm256_loadu_pd(bp);
__m256d b1 = _mm256_loadu_pd(bp + 4);
__m256d a0 = _mm256_broadcast_sd(ap + 0);
c00 = _mm256_fmadd_pd(a0, b0, c00);
c01 = _mm256_fmadd_pd(a0, b1, c01);
__m256d a1 = _mm256_broadcast_sd(ap + 1);
c10 = _mm256_fmadd_pd(a1, b0, c10);
c11 = _mm256_fmadd_pd(a1, b1, c11);
__m256d a2 = _mm256_broadcast_sd(ap + 2);
c20 = _mm256_fmadd_pd(a2, b0, c20);
c21 = _mm256_fmadd_pd(a2, b1, c21);
__m256d a3 = _mm256_broadcast_sd(ap + 3);
c30 = _mm256_fmadd_pd(a3, b0, c30);
c31 = _mm256_fmadd_pd(a3, b1, c31);
ap += MR;
bp += NR;
}
_mm256_storeu_pd(c + 0*lda, c00);
_mm256_storeu_pd(c + 0*lda + 4, c01);
_mm256_storeu_pd(c + 1*lda, c10);
_mm256_storeu_pd(c + 1*lda + 4, c11);
_mm256_storeu_pd(c + 2*lda, c20);
_mm256_storeu_pd(c + 2*lda + 4, c21);
_mm256_storeu_pd(c + 3*lda, c30);
_mm256_storeu_pd(c + 3*lda + 4, c31);
}
void matrix_multiply(int n, const double *A, const double *B, double *C) {
// Loop order: ic (m), pc (k), jc (n) -- A_pack reused across jc.
for (int ic = 0; ic < n; ic += MC) {
for (int pc = 0; pc < n; pc += KC) {
// pack A: MC x KC
for (int p = 0; p < MC / MR; p++) {
for (int k = 0; k < KC; k++) {
const double *Acol = A + (ic + p * MR) * n + (pc + k);
double *ap = A_pack + (p * KC + k) * MR;
ap[0] = Acol[0];
ap[1] = Acol[n];
ap[2] = Acol[2*n];
ap[3] = Acol[3*n];
}
}
for (int jc = 0; jc < n; jc += NC) {
// pack B: KC x NC
for (int k = 0; k < KC; k++) {
const double *Brow = B + (pc + k) * n + jc;
double *bp = B_pack + k * NR;
for (int q = 0; q < NC / NR; q++) {
__m256d v0 = _mm256_loadu_pd(Brow + q * NR);
__m256d v1 = _mm256_loadu_pd(Brow + q * NR + 4);
_mm256_storeu_pd(bp + q * KC * NR, v0);
_mm256_storeu_pd(bp + q * KC * NR + 4, v1);
}
}
// compute
for (int p = 0; p < MC / MR; p++) {
double *cp = C + (ic + p * MR) * n + jc;
const double *ap = A_pack + p * KC * MR;
for (int q = 0; q < NC / NR; q++) {
const double *bp = B_pack + q * KC * NR;
microkernel(ap, bp, cp + q * NR, n);
}
}
}
}
}
}
// ---- bench harness (for customtest) ----
#ifdef BENCH
static double Ab[1024*1024] __attribute__((aligned(4096)));
static double Bb[1024*1024] __attribute__((aligned(4096)));
static double Cb[1024*1024] __attribute__((aligned(4096)));
int main() {
const int n = 1024;
// Fast fill (cheap, no division) so the measurement is dominated by GEMM.
for (int i = 0; i < n*n; i++) {
Ab[i] = (double)(i & 1023) * (1.0/1024.0) - 0.5;
Bb[i] = (double)((i * 7) & 1023) * (1.0/1024.0) - 0.5;
}
matrix_multiply(n, Ab, Bb, Cb);
volatile double acc = Cb[0] + Cb[n*n-1] + Cb[12345];
printf("%g\n", (double)acc);
return 0;
}
#endif
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 44.793 ms | 8 MB + 652 KB | Accepted | Score: 100 | 显示更多 |