// mmms* -- C = A*B, short, exact mod 2^16, Strassen on top of a packed base-case GEMM.
// The identity used is exact in Z/2^16, so recursion is valid for arbitrary short data.
#include <immintrin.h>
#include <stddef.h>
#include <stdlib.h>
#pragma GCC push_options
#pragma GCC target("avx2")
typedef short TE;
#ifndef CUT
#define CUT 1024
#endif
#ifndef MR
#define MR 6
#endif
#ifndef NR
#define NR 32
#endif
#ifndef KC
#define KC 64
#endif
#ifndef NC
#define NC 1024
#endif
#define LDV(p) _mm256_loadu_si256((const __m256i *)(p))
#define STV(p, v) _mm256_storeu_si256((__m256i *)(p), (v))
static TE Apanel[MR * KC * 16 + 64] __attribute__((aligned(64)));
static TE Bpanel[(size_t)NC * KC + 128] __attribute__((aligned(64)));
static inline __m256i b16(const TE *p) {
__m256i r; __asm__("vpbroadcastw %1, %0" : "=x"(r) : "m"(*p)); return r;
}
// ---------------- micro kernel: MR x NR, A panel pre-replicated 16x ----------------
#define KSTEP(off, C0, C1) \
"vpmullw " #off "(%[a]), %%ymm14, %%ymm13\n\t" \
"vpaddw %%ymm13, %" C0 ", %" C0 "\n\t" \
"vpmullw " #off "(%[a]), %%ymm15, %%ymm13\n\t" \
"vpaddw %%ymm13, %" C1 ", %" C1 "\n\t"
#define MICRO_BODY \
"1:\n\t" \
"vmovdqa 0(%[b]), %%ymm14\n\t" \
"vmovdqa 32(%[b]), %%ymm15\n\t" \
KSTEP(0, "0", "1") KSTEP(32, "2", "3") KSTEP(64, "4", "5") \
KSTEP(96, "6", "7") KSTEP(128, "8", "9") KSTEP(160, "10", "11") \
"addq $192, %[a]\n\t" \
"addq $64, %[b]\n\t" \
"decq %[k]\n\t" \
"jnz 1b\n\t"
#define MICRO_CLOB "ymm12", "ymm13", "ymm14", "ymm15", "cc", "memory"
template <int OP> // 0 = store, 1 = add, 2 = sub
static inline void micro(long k, const TE *ap, const TE *bp, TE *C, int ldc, int rows) {
__m256i c00 = _mm256_setzero_si256(), c01 = c00, c10 = c00, c11 = c00, c20 = c00, c21 = c00;
__m256i c30 = c00, c31 = c00, c40 = c00, c41 = c00, c50 = c00, c51 = c00;
__asm__ volatile(
MICRO_BODY
: "+x"(c00), "+x"(c01), "+x"(c10), "+x"(c11), "+x"(c20), "+x"(c21),
"+x"(c30), "+x"(c31), "+x"(c40), "+x"(c41), "+x"(c50), "+x"(c51),
[a] "+r"(ap), [b] "+r"(bp), [k] "+r"(k)
:
: MICRO_CLOB);
#define ST(r, v0, v1) \
if (rows > (r)) { \
TE *p0 = C + (size_t)(r) * ldc; \
if (OP == 0) { STV(p0, v0); STV(p0 + 16, v1); } \
else if (OP == 1) { STV(p0, _mm256_add_epi16(LDV(p0), v0)); STV(p0 + 16, _mm256_add_epi16(LDV(p0 + 16), v1)); } \
else { STV(p0, _mm256_sub_epi16(LDV(p0), v0)); STV(p0 + 16, _mm256_sub_epi16(LDV(p0 + 16), v1)); } }
ST(0, c00, c01) ST(1, c10, c11) ST(2, c20, c21)
ST(3, c30, c31) ST(4, c40, c41) ST(5, c50, c51)
#undef ST
}
// ---------------- base-case GEMM: C = A*B (lda/ldb arbitrary), n<=NC ----------------
static void leaf(TE *C, int ldc, const TE *A, int lda, const TE *B, int ldb, int n) {
const int nbc = n - (n % NR);
for (int jc = 0; jc < nbc; jc += NC) {
const int nc = (nbc - jc < NC) ? (nbc - jc) : NC;
for (int pc = 0; pc < n; pc += KC) {
const int kc = (n - pc < KC) ? (n - pc) : KC;
for (int jr = 0, jb = 0; jr < nc; jr += NR, jb++) {
TE *q = Bpanel + (size_t)jb * kc * NR;
for (int k = 0; k < kc; k++) {
const TE *src = B + (size_t)(pc + k) * ldb + jc + jr;
for (int j = 0; j < NR; j++) q[j] = src[j];
q += NR;
}
}
const int firstpc = (pc == 0);
for (int ic = 0; ic < n; ic += MR) {
const int rows = (n - ic < MR) ? (n - ic) : MR;
TE *p = Apanel;
for (int k = 0; k < kc; k++) {
for (int r = 0; r < rows; r++) { STV(p, b16(A + (size_t)(ic + r) * lda + pc + k)); p += 16; }
for (int r = rows; r < MR; r++) { STV(p, _mm256_setzero_si256()); p += 16; }
}
const TE *b0 = Bpanel;
TE *cp = C + (size_t)ic * ldc + jc;
for (int jr = 0; jr < nc; jr += NR) {
if (firstpc) micro<0>(kc, Apanel, b0, cp + jr, ldc, rows);
else micro<1>(kc, Apanel, b0, cp + jr, ldc, rows);
b0 += (size_t)kc * NR;
}
}
}
}
if (nbc < n) { // ragged columns (not reachable for our power-of-two sizes)
for (int i = 0; i < n; i++)
for (int j = nbc; j < n; j++) {
unsigned acc = 0;
for (int k = 0; k < n; k++)
acc += (unsigned)(unsigned short)A[(size_t)i * lda + k] * (unsigned)(unsigned short)B[(size_t)k * ldb + j];
C[(size_t)i * ldc + j] = (TE)acc;
}
}
}
// ---------------- Strassen driver ----------------
static TE *g_pool;
static size_t g_off;
// form the 5 A-sums and 5 B-sums of an h-split block
static void mk5(TE *AS, TE *BS, const TE *A, int lda, const TE *B, int ldb, int h) {
const size_t q = (size_t)h * h;
const TE *a11 = A, *a12 = A + h, *a21 = A + (size_t)h * lda, *a22 = a21 + h;
const TE *b11 = B, *b12 = B + h, *b21 = B + (size_t)h * ldb, *b22 = b21 + h;
for (int i = 0; i < h; i++) {
const TE *r11 = a11 + (size_t)i * lda, *r12 = a12 + (size_t)i * lda;
const TE *r21 = a21 + (size_t)i * lda, *r22 = a22 + (size_t)i * lda;
const TE *s11 = b11 + (size_t)i * ldb, *s12 = b12 + (size_t)i * ldb;
const TE *s21 = b21 + (size_t)i * ldb, *s22 = b22 + (size_t)i * ldb;
TE *o0 = AS + (size_t)i * h, *o1 = o0 + q, *o2 = o1 + q, *o3 = o2 + q, *o4 = o3 + q;
TE *p0 = BS + (size_t)i * h, *p1 = p0 + q, *p2 = p1 + q, *p3 = p2 + q, *p4 = p3 + q;
int j = 0;
for (; j + 16 <= h; j += 16) {
__m256i v11 = LDV(r11 + j), v12 = LDV(r12 + j), v21 = LDV(r21 + j), v22 = LDV(r22 + j);
__m256i w11 = LDV(s11 + j), w12 = LDV(s12 + j), w21 = LDV(s21 + j), w22 = LDV(s22 + j);
STV(o0 + j, _mm256_add_epi16(v11, v22)); // A11+A22
STV(o1 + j, _mm256_add_epi16(v21, v22)); // A21+A22
STV(o2 + j, _mm256_add_epi16(v11, v12)); // A11+A12
STV(o3 + j, _mm256_sub_epi16(v21, v11)); // A21-A11
STV(o4 + j, _mm256_sub_epi16(v12, v22)); // A12-A22
STV(p0 + j, _mm256_add_epi16(w11, w22)); // B11+B22
STV(p1 + j, _mm256_sub_epi16(w12, w22)); // B12-B22
STV(p2 + j, _mm256_sub_epi16(w21, w11)); // B21-B11
STV(p3 + j, _mm256_add_epi16(w11, w12)); // B11+B12
STV(p4 + j, _mm256_add_epi16(w21, w22)); // B21+B22
}
for (; j < h; j++) {
TE x11 = r11[j], x12 = r12[j], x21 = r21[j], x22 = r22[j];
TE y11 = s11[j], y12 = s12[j], y21 = s21[j], y22 = s22[j];
o0[j] = x11 + x22; o1[j] = x21 + x22; o2[j] = x11 + x12;
o3[j] = x21 - x11; o4[j] = x12 - x22;
p0[j] = y11 + y22; p1[j] = y12 - y22; p2[j] = y21 - y11;
p3[j] = y11 + y12; p4[j] = y21 + y22;
}
}
}
static void combine(TE *C, int ldc, const TE *M1, const TE *M2, const TE *M3,
const TE *M4, const TE *M5, const TE *M6, const TE *M7, int h) {
const size_t q = (size_t)h * h;
for (int i = 0; i < h; i++) {
const TE *m1 = M1 + (size_t)i * h, *m2 = M2 + (size_t)i * h, *m3 = M3 + (size_t)i * h;
const TE *m4 = M4 + (size_t)i * h, *m5 = M5 + (size_t)i * h, *m6 = M6 + (size_t)i * h;
const TE *m7 = M7 + (size_t)i * h;
TE *c11 = C + (size_t)i * ldc, *c12 = c11 + h;
TE *c21 = c11 + (size_t)h * ldc, *c22 = c21 + h;
int j = 0;
for (; j + 16 <= h; j += 16) {
__m256i a1 = LDV(m1 + j), a2 = LDV(m2 + j), a3 = LDV(m3 + j), a4 = LDV(m4 + j);
__m256i a5 = LDV(m5 + j), a6 = LDV(m6 + j), a7 = LDV(m7 + j);
// C11 = M1+M4-M5+M7 C12 = M3+M5 C21 = M2+M4 C22 = M1-M2+M3+M6
STV(c11 + j, _mm256_add_epi16(_mm256_sub_epi16(_mm256_add_epi16(a1, a4), a5), a7));
STV(c12 + j, _mm256_add_epi16(a3, a5));
STV(c21 + j, _mm256_add_epi16(a2, a4));
STV(c22 + j, _mm256_add_epi16(_mm256_sub_epi16(_mm256_add_epi16(a1, a3), a2), a6));
}
for (; j < h; j++) {
c11[j] = m1[j] + m4[j] - m5[j] + m7[j];
c12[j] = m3[j] + m5[j];
c21[j] = m2[j] + m4[j];
c22[j] = m1[j] - m2[j] + m3[j] + m6[j];
}
}
(void)q;
}
static void mm(TE *C, int ldc, const TE *A, int lda, const TE *B, int ldb, int n) {
if (n <= CUT) { leaf(C, ldc, A, lda, B, ldb, n); return; }
const int h = n >> 1;
const size_t q = (size_t)h * h;
TE *AS = g_pool + g_off;
TE *BS = AS + 5 * q;
TE *M = BS + 5 * q;
g_off += 17 * q;
mk5(AS, BS, A, lda, B, ldb, h);
const TE *A11 = A, *A22 = A + (size_t)h * lda + h;
const TE *B11 = B, *B22 = B + (size_t)h * ldb + h;
mm(M + 0 * q, h, AS + 0 * q, h, BS + 0 * q, h, h); // M1 = (A11+A22)(B11+B22)
mm(M + 1 * q, h, AS + 1 * q, h, B11, ldb, h); // M2 = (A21+A22)B11
mm(M + 2 * q, h, A11, lda, BS + 1 * q, h, h); // M3 = A11(B12-B22)
mm(M + 3 * q, h, A22, lda, BS + 2 * q, h, h); // M4 = A22(B21-B11)
mm(M + 4 * q, h, AS + 2 * q, h, B22, ldb, h); // M5 = (A11+A12)B22
mm(M + 5 * q, h, AS + 3 * q, h, BS + 3 * q, h, h); // M6 = (A21-A11)(B11+B12)
mm(M + 6 * q, h, AS + 4 * q, h, BS + 4 * q, h, h); // M7 = (A12-A22)(B21+B22)
combine(C, ldc, M, M + q, M + 2 * q, M + 3 * q, M + 4 * q, M + 5 * q, M + 6 * q, h);
g_off -= 17 * q;
}
void matrix_multiply(int n, const short *A, const short *B, short *C) {
size_t need = 64;
for (int m = n; m > CUT; m >>= 1) need += 17 * (size_t)(m >> 1) * (m >> 1);
if (need > 64) {
void *p = 0;
if (posix_memalign(&p, 4096, need * sizeof(TE))) return;
g_pool = (TE *)p;
g_off = 0;
mm(C, n, A, n, B, n, n);
free(p);
} else {
leaf(C, n, A, n, B, n, n);
}
}
#pragma GCC pop_options
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 954.976 ms | 202 MB + 156 KB | Accepted | Score: 100 | 显示更多 |