// mmmd1k: C = A*B, n = 1024, row-major doubles.
// Blocked GEMM, AVX2+FMA 6x8 register microkernel (4x8 for tail rows).
#pragma GCC target("avx2,fma")
#include <immintrin.h>
#include <cstring>
#define MR 4
#define MR6 6
#define NR 8
#define MC 120
#define KC 256
#define NC 1024
static double Apack[(size_t)MC * KC + 64] __attribute__((aligned(64)));
static double Bpack[(size_t)(KC + 4) * (NC + NR) + 64] __attribute__((aligned(64)));
// Pack 6 rows of an A block, k-major: Apack[(ig/6)*kc*6 + k*6 + i]
static void packA6(const double *A, int n, int ic, int pc, int mc, int kc) {
double *dst = Apack;
for (int ig = 0; ig < mc; ig += MR6) {
const double *r0 = A + (size_t)(ic + ig) * n + pc;
const double *r1 = r0 + n, *r2 = r1 + n, *r3 = r2 + n, *r4 = r3 + n, *r5 = r4 + n;
for (int k = 0; k < kc; k++) {
dst[0] = r0[k]; dst[1] = r1[k]; dst[2] = r2[k];
dst[3] = r3[k]; dst[4] = r4[k]; dst[5] = r5[k];
dst += MR6;
}
}
}
// Pack B into NR-wide panels: panel p holds kc rows of NR contiguous doubles.
static void packB(const double *B, int n, int pc, int jc, int kc, int nc) {
double *dst = Bpack;
for (int jr = 0; jr < nc; jr += NR) {
const double *src = B + (size_t)pc * n + jc + jr;
for (int k = 0; k < kc; k++) {
_mm256_store_pd(dst, _mm256_loadu_pd(src + (size_t)k * n));
_mm256_store_pd(dst + 4, _mm256_loadu_pd(src + (size_t)k * n + 4));
dst += NR;
}
}
}
#define STROW(R, A, B) t = _mm256_loadu_pd(cp); _mm256_storeu_pd(cp, _mm256_add_pd(t, A)); \
t = _mm256_loadu_pd(cp + 4); _mm256_storeu_pd(cp + 4, _mm256_add_pd(t, B)); cp += ldc;
static void kernel6x8(const double *Ap, const double *Bp, int kc, double *C, int ldc) {
__m256d c0=_mm256_setzero_pd(),c1=_mm256_setzero_pd(),c2=_mm256_setzero_pd();
__m256d c3=_mm256_setzero_pd(),c4=_mm256_setzero_pd(),c5=_mm256_setzero_pd();
__m256d c6=_mm256_setzero_pd(),c7=_mm256_setzero_pd(),c8=_mm256_setzero_pd();
__m256d c9=_mm256_setzero_pd(),c10=_mm256_setzero_pd(),c11=_mm256_setzero_pd();
const double *ap = Ap, *bp = Bp;
#define K6STEP() { \
__m256d b0 = _mm256_load_pd(bp); \
__m256d b1 = _mm256_load_pd(bp + 4); \
__m256d a; \
a = _mm256_broadcast_sd(ap + 0); c0 = _mm256_fmadd_pd(a, b0, c0); c6 = _mm256_fmadd_pd(a, b1, c6); \
a = _mm256_broadcast_sd(ap + 1); c1 = _mm256_fmadd_pd(a, b0, c1); c7 = _mm256_fmadd_pd(a, b1, c7); \
a = _mm256_broadcast_sd(ap + 2); c2 = _mm256_fmadd_pd(a, b0, c2); c8 = _mm256_fmadd_pd(a, b1, c8); \
a = _mm256_broadcast_sd(ap + 3); c3 = _mm256_fmadd_pd(a, b0, c3); c9 = _mm256_fmadd_pd(a, b1, c9); \
a = _mm256_broadcast_sd(ap + 4); c4 = _mm256_fmadd_pd(a, b0, c4); c10 = _mm256_fmadd_pd(a, b1, c10); \
a = _mm256_broadcast_sd(ap + 5); c5 = _mm256_fmadd_pd(a, b0, c5); c11 = _mm256_fmadd_pd(a, b1, c11); \
ap += MR6; bp += NR; \
}
int k = 0;
for (; k + 4 <= kc; k += 4) {
_mm_prefetch((const char *)(ap + 384), _MM_HINT_T0);
_mm_prefetch((const char *)(bp + 96), _MM_HINT_T0);
K6STEP()
K6STEP()
K6STEP()
K6STEP()
}
for (; k < kc; k++) K6STEP()
double *cp = C; __m256d t;
STROW(0, c0, c6) STROW(1, c1, c7) STROW(2, c2, c8)
STROW(3, c3, c9) STROW(4, c4, c10) STROW(5, c5, c11)
}
static void kernel4x8(const double *Ap, const double *Bp, int kc, int nc,
double *C, int ldc) {
__m256d c0=_mm256_setzero_pd(),c1=_mm256_setzero_pd(),c2=_mm256_setzero_pd(),c3=_mm256_setzero_pd();
__m256d c4=_mm256_setzero_pd(),c5=_mm256_setzero_pd(),c6=_mm256_setzero_pd(),c7=_mm256_setzero_pd();
const double *ap = Ap, *bp = Bp;
for (int k = 0; k < kc; k++) {
__m256d b0 = _mm256_load_pd(bp), b1 = _mm256_load_pd(bp + 4), a;
a = _mm256_broadcast_sd(ap + 0); c0 = _mm256_fmadd_pd(a, b0, c0); c4 = _mm256_fmadd_pd(a, b1, c4);
a = _mm256_broadcast_sd(ap + 1); c1 = _mm256_fmadd_pd(a, b0, c1); c5 = _mm256_fmadd_pd(a, b1, c5);
a = _mm256_broadcast_sd(ap + 2); c2 = _mm256_fmadd_pd(a, b0, c2); c6 = _mm256_fmadd_pd(a, b1, c6);
a = _mm256_broadcast_sd(ap + 3); c3 = _mm256_fmadd_pd(a, b0, c3); c7 = _mm256_fmadd_pd(a, b1, c7);
ap += MR; bp += NR;
}
double *cp = C; __m256d t;
STROW(0, c0, c4) STROW(1, c1, c5) STROW(2, c2, c6) STROW(3, c3, c7)
}
static void packA4(const double *A, int n, int ic, int pc, int mc, int kc) {
double *dst = Apack;
for (int ig = 0; ig < mc; ig += MR) {
const double *r0 = A + (size_t)(ic + ig) * n + pc;
const double *r1 = r0 + n, *r2 = r1 + n, *r3 = r2 + n;
for (int k = 0; k < kc; k++) {
dst[0] = r0[k]; dst[1] = r1[k]; dst[2] = r2[k]; dst[3] = r3[k];
dst += MR;
}
}
}
static void edge_rows(const double *A, const double *B, double *C, int n, int i0, int i1, int j0, int j1) {
for (int i = i0; i < i1; i++) {
for (int j = j0; j < j1; j++) {
const double *ar = A + (size_t)i * n;
double s = 0;
for (int k = 0; k < n; k++) s += ar[k] * B[(size_t)k * n + j];
C[(size_t)i * n + j] = s;
}
}
}
void matrix_multiply(int n, const double *A, const double *B, double *C) {
memset(C, 0, sizeof(double) * (size_t)n * n);
int n6 = n - (n % MR6); // rows handled by the 6-row kernel
int n4 = n - (n % NR); // cols handled by the vector kernels
int ntail = n - n6;
int n4r = n6 + (ntail - (ntail % MR)); // rows handled by the vector kernels
for (int jc = 0; jc < n4; jc += NC) {
int nc = n4 - jc; if (nc > NC) nc = NC;
for (int pc = 0; pc < n; pc += KC) {
int kc = n - pc; if (kc > KC) kc = KC;
packB(B, n, pc, jc, kc, nc);
for (int ic = 0; ic < n6; ic += MC) {
int mc = n6 - ic; if (mc > MC) mc = MC;
packA6(A, n, ic, pc, mc, kc);
for (int jr = 0; jr < nc; jr += NR) {
const double *bp = Bpack + (size_t)(jr / NR) * kc * NR;
double *cb = C + (size_t)ic * n + jc + jr;
for (int ir = 0; ir < mc; ir += MR6)
kernel6x8(Apack + (size_t)(ir / MR6) * kc * MR6, bp, kc,
cb + (size_t)ir * n, n);
}
}
if (n6 < n4r) { // rows [n6, n4r) with the 4-row kernel
int mc = n4r - n6;
packA4(A, n, n6, pc, mc, kc);
for (int jr = 0; jr < nc; jr += NR) {
const double *bp = Bpack + (size_t)(jr / NR) * kc * NR;
double *cb = C + (size_t)n6 * n + jc + jr;
for (int ir = 0; ir < mc; ir += MR)
kernel4x8(Apack + (size_t)(ir / MR) * kc * MR, bp, kc, nc,
cb + (size_t)ir * n, n);
}
}
}
}
if (n4 < n) edge_rows(A, B, C, n, 0, n, n4, n);
if (n4r < n) edge_rows(A, B, C, n, n4r, n, 0, n4);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 48.014 ms | 10 MB + 252 KB | Accepted | Score: 100 | 显示更多 |