#include <immintrin.h>
#include <stddef.h>
// mmmd1k: C = A*B, n=1024, row-major, double precision.
// Blocked GEMM: MC=NC=KC=256 (divide 1024), microkernel 4x8 AVX2+FMA.
#define MC 256
#define KC 256
#define NC 256
#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;
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) {
const int M = n, N = n, K = n;
const int npM = M / MR; // 256
const int npN = N / NR; // 128
for (int jc = 0; jc < N; jc += NC) {
for (int pc = 0; pc < K; pc += KC) {
// 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);
}
}
for (int ic = 0; ic < M; ic += MC) {
// 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 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);
}
}
}
}
}
}
| Compilation | N/A | N/A | Compile Error | Score: N/A | 显示更多 |