/* FILE-SCOPE FLAGS: judge compiles with g++ -O2 -static -nostdlib -U_FORTIFY_SOURCE and NO
* -march, so the baseline is x86-64 = SSE2 ONLY. Widen the target for EVERY function in
* the file (including all packing / C-traffic code) and lift -O2 to -O3 + unroll. */
#pragma GCC optimize("O3")
#pragma GCC target("avx2,bmi,bmi2,popcnt,lzcnt")
/* int32 GEMM exact mod 2^32, AVX2. Three-madd vpmaddwd mikro (KC=128 baked in),
* wrapped in EXACT Strassen over Z/2^32 (7/8 of the multiply work per level).
* GENERATED by gen_strassen.py -- do not edit by hand.
*/
#include <immintrin.h>
#include <string.h>
#include <stdint.h>
#include <stdlib.h>
#define MR 4
#define NCG 1
#define KC 128
#define MC 4
#define NC 128
#define NPAIR (KC / 2)
#define APSTR (NPAIR * 16 + 16)
#define NR (8 * NCG)
#define BASE 128
typedef int32_t T;
#define AVX2 __attribute__((target("avx2")))
#define LD(p) _mm256_load_si256((const __m256i *)(p))
#define LDA(p) _mm256_load_si256((const __m256i *)(p))
#define STA(p, v) _mm256_store_si256((__m256i *)(p), (v))
#define ST(p, v) _mm256_store_si256((__m256i *)(p), (v))
static int32_t Apk[(MC + 4) * APSTR + 64] __attribute__((aligned(64)));
static int32_t Bpk[16896] __attribute__((aligned(64))); /* [jg][m][B0|B1] */
#define NTSEL 0
/* MIKRO-BEGIN */
#define AVX2 __attribute__((target("avx2")))
static __m256i LA[128] __attribute__((aligned(32)));
AVX2 __attribute__((always_inline)) static inline void mikro(int npair, const int32_t *ap, const int32_t *bp,
int32_t *c, int ldc, int first) {
(void)npair;
__m256i q0, q1, q2, q3;
__asm__ volatile(
"vpxor %%xmm15, %%xmm15, %%xmm15\n\t"
"vmovdqa %%ymm15, %%ymm0\n\t"
"vmovdqa %%ymm15, %%ymm4\n\t"
"vmovdqa %%ymm15, %%ymm1\n\t"
"vmovdqa %%ymm15, %%ymm5\n\t"
"vmovdqa %%ymm15, %%ymm2\n\t"
"vmovdqa %%ymm15, %%ymm6\n\t"
"vmovdqa %%ymm15, %%ymm3\n\t"
"vmovdqa %%ymm15, %%ymm7\n\t"
"test %[cnt], %[cnt]\n\t"
"jle 2f\n\t"
"1:\n\t"
"prefetcht0 1024(%[bp])\n\t"
"vmovdqu 0(%[bp]), %%ymm8\n\t"
"vmovdqu 32(%[bp]), %%ymm9\n\t"
"vpmaddwd 0(%[ap]), %%ymm8, %%ymm10\n\t"
"vpmaddwd 4160+0(%[ap]), %%ymm8, %%ymm11\n\t"
"vpmaddwd 8320+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpmaddwd 12480+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpmaddwd 0(%[ap]), %%ymm9, %%ymm14\n\t"
"vpmaddwd 4160+0(%[ap]), %%ymm9, %%ymm15\n\t"
"vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
"vpmaddwd 8320+0(%[ap]), %%ymm9, %%ymm10\n\t"
"vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
"vpmaddwd 12480+0(%[ap]), %%ymm9, %%ymm11\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpmaddwd 32+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpmaddwd 4160+32+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
"vpmaddwd 8320+32+0(%[ap]), %%ymm8, %%ymm14\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpmaddwd 12480+32+0(%[ap]), %%ymm8, %%ymm15\n\t"
"vmovdqu 64(%[bp]), %%ymm8\n\t"
"vmovdqu 96(%[bp]), %%ymm9\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpmaddwd 64+0(%[ap]), %%ymm8, %%ymm10\n\t"
"vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
"vpmaddwd 4160+64+0(%[ap]), %%ymm8, %%ymm11\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpmaddwd 8320+64+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpmaddwd 12480+64+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
"vpmaddwd 64+0(%[ap]), %%ymm9, %%ymm14\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpmaddwd 4160+64+0(%[ap]), %%ymm9, %%ymm15\n\t"
"vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
"vpmaddwd 8320+64+0(%[ap]), %%ymm9, %%ymm10\n\t"
"vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
"vpmaddwd 12480+64+0(%[ap]), %%ymm9, %%ymm11\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpmaddwd 96+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpmaddwd 4160+96+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
"vpmaddwd 8320+96+0(%[ap]), %%ymm8, %%ymm14\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpmaddwd 12480+96+0(%[ap]), %%ymm8, %%ymm15\n\t"
"vmovdqu 128(%[bp]), %%ymm8\n\t"
"vmovdqu 160(%[bp]), %%ymm9\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpmaddwd 128+0(%[ap]), %%ymm8, %%ymm10\n\t"
"vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
"vpmaddwd 4160+128+0(%[ap]), %%ymm8, %%ymm11\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpmaddwd 8320+128+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpmaddwd 12480+128+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
"vpmaddwd 128+0(%[ap]), %%ymm9, %%ymm14\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpmaddwd 4160+128+0(%[ap]), %%ymm9, %%ymm15\n\t"
"vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
"vpmaddwd 8320+128+0(%[ap]), %%ymm9, %%ymm10\n\t"
"vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
"vpmaddwd 12480+128+0(%[ap]), %%ymm9, %%ymm11\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpmaddwd 160+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpmaddwd 4160+160+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
"vpmaddwd 8320+160+0(%[ap]), %%ymm8, %%ymm14\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpmaddwd 12480+160+0(%[ap]), %%ymm8, %%ymm15\n\t"
"vmovdqu 192(%[bp]), %%ymm8\n\t"
"vmovdqu 224(%[bp]), %%ymm9\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpmaddwd 192+0(%[ap]), %%ymm8, %%ymm10\n\t"
"vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
"vpmaddwd 4160+192+0(%[ap]), %%ymm8, %%ymm11\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpmaddwd 8320+192+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpmaddwd 12480+192+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
"vpmaddwd 192+0(%[ap]), %%ymm9, %%ymm14\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpmaddwd 4160+192+0(%[ap]), %%ymm9, %%ymm15\n\t"
"vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
"vpmaddwd 8320+192+0(%[ap]), %%ymm9, %%ymm10\n\t"
"vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
"vpmaddwd 12480+192+0(%[ap]), %%ymm9, %%ymm11\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpmaddwd 224+0(%[ap]), %%ymm8, %%ymm12\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpmaddwd 4160+224+0(%[ap]), %%ymm8, %%ymm13\n\t"
"vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
"vpmaddwd 8320+224+0(%[ap]), %%ymm8, %%ymm14\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpmaddwd 12480+224+0(%[ap]), %%ymm8, %%ymm15\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"add $256, %[ap]\n\t"
"add $256, %[bp]\n\t"
"sub $4, %[cnt]\n\t"
"jnz 1b\n\t"
"2:\n\t"
"vpslld $16, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm4, %%ymm0, %[q0]\n\t"
"vpslld $16, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm5, %%ymm1, %[q1]\n\t"
"vpslld $16, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm6, %%ymm2, %[q2]\n\t"
"vpslld $16, %%ymm7, %%ymm7\n\t"
"vpaddd %%ymm7, %%ymm3, %[q3]\n\t"
""
: [ap] "+r" (ap), [bp] "+r" (bp), [cnt] "+r" (npair),
[q0] "=x" (q0), [q1] "=x" (q1), [q2] "=x" (q2), [q3] "=x" (q3)
: [la] "r" (&LA[0])
: "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm15", "memory");
{ int32_t *cp = c + (size_t)0 * ldc + 0;
__m256i rs = q0;
if (first) _mm256_store_si256((__m256i *)cp, rs);
else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
{ int32_t *cp = c + (size_t)1 * ldc + 0;
__m256i rs = q1;
if (first) _mm256_store_si256((__m256i *)cp, rs);
else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
{ int32_t *cp = c + (size_t)2 * ldc + 0;
__m256i rs = q2;
if (first) _mm256_store_si256((__m256i *)cp, rs);
else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
{ int32_t *cp = c + (size_t)3 * ldc + 0;
__m256i rs = q3;
if (first) _mm256_store_si256((__m256i *)cp, rs);
else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
}
/* MIKRO-END */
/* ---- ld-aware packing + tiled macro-kernel (base case) ---- */
AVX2 static inline void packA(int ic, int kk, int n, int lda, const T *A) {
const __m256i m16 = _mm256_set1_epi32(0xFFFF);
const __m256i half = _mm256_set1_epi32(0x8000);
const __m256i i0 = _mm256_set1_epi32(0), i1 = _mm256_set1_epi32(1);
const __m256i i4 = _mm256_set1_epi32(4), i5 = _mm256_set1_epi32(5);
(void)n;
for (int r = 0; r < MC; r++) {
const T *src = A + (size_t)(ic + r) * lda + kk;
T *dst = Apk + (size_t)r * APSTR;
for (int m = 0; m < KC / 8; m++) {
__m256i x = LD(src + 8 * m);
__m256i al = _mm256_and_si256(x, m16);
__m256i ah = _mm256_srli_epi32(_mm256_add_epi32(x, half), 16);
__m256i pl = _mm256_packus_epi32(al, al);
__m256i ph = _mm256_packus_epi32(ah, ah);
T *d = dst + 64 * m;
ST(d + 0, _mm256_permutevar8x32_epi32(pl, i0));
ST(d + 8, _mm256_permutevar8x32_epi32(ph, i0));
ST(d + 16, _mm256_permutevar8x32_epi32(pl, i1));
ST(d + 24, _mm256_permutevar8x32_epi32(ph, i1));
ST(d + 32, _mm256_permutevar8x32_epi32(pl, i4));
ST(d + 40, _mm256_permutevar8x32_epi32(ph, i4));
ST(d + 48, _mm256_permutevar8x32_epi32(pl, i5));
ST(d + 56, _mm256_permutevar8x32_epi32(ph, i5));
}
}
}
AVX2 static inline void packB(int jc, int kk, int n, int ldb, const T *B) {
const __m256i m16 = _mm256_set1_epi32(0xFFFF);
const __m256i half = _mm256_set1_epi32(0x8000);
(void)n;
for (int m = 0; m < NPAIR; m++) {
const T *s0 = B + (size_t)(kk + 2 * m) * ldb + jc;
const T *s1 = B + (size_t)(kk + 2 * m + 1) * ldb + jc;
T *d0 = Bpk + (size_t)m * 16;
T *d1 = d0 + 8;
for (int j = 0; j < NC; j += 8) {
__m256i x0 = LD(s0 + j), x1 = LD(s1 + j);
__m256i l0 = _mm256_and_si256(x0, m16);
__m256i l1 = _mm256_and_si256(x1, m16);
__m256i h0 = _mm256_srli_epi32(_mm256_add_epi32(x0, half), 16);
__m256i h1 = _mm256_srli_epi32(_mm256_add_epi32(x1, half), 16);
ST(d0 + (size_t)(j / 8) * (NPAIR * 16 + 16), _mm256_or_si256(l0, _mm256_slli_epi32(l1, 16)));
ST(d1 + (size_t)(j / 8) * (NPAIR * 16 + 16), _mm256_or_si256(h0, _mm256_slli_epi32(h1, 16)));
}
}
}
AVX2 static inline void packA2(int ic, int kk, int n, int lda, const T *A, int sg, const T *A2) {
const __m256i m16 = _mm256_set1_epi32(0xFFFF);
const __m256i half = _mm256_set1_epi32(0x8000);
const __m256i i0 = _mm256_set1_epi32(0), i1 = _mm256_set1_epi32(1);
const __m256i i4 = _mm256_set1_epi32(4), i5 = _mm256_set1_epi32(5);
(void)n;
for (int r = 0; r < MC; r++) {
const T *src = A + (size_t)(ic + r) * lda + kk;
const T *src2 = A2 ? (A2 + (size_t)(ic + r) * lda + kk) : 0;
T *dst = Apk + (size_t)r * APSTR;
for (int m = 0; m < KC / 8; m++) {
__builtin_prefetch(src + lda + 8 * m, 0, 1);
__builtin_prefetch(src + 8 * m + 128, 0, 1);
if (src2) { __builtin_prefetch(src2 + lda + 8 * m, 0, 1);
__builtin_prefetch(src2 + 8 * m + 128, 0, 1); }
__m256i x = LD(src + 8 * m);
if (src2) { __m256i y = LD(src2 + 8 * m);
x = (sg > 0) ? _mm256_add_epi32(x, y) : _mm256_sub_epi32(x, y); }
__m256i al = _mm256_and_si256(x, m16);
__m256i ah = _mm256_srli_epi32(_mm256_add_epi32(x, half), 16);
__m256i pl = _mm256_packus_epi32(al, al);
__m256i ph = _mm256_packus_epi32(ah, ah);
T *d = dst + 64 * m;
ST(d + 0, _mm256_permutevar8x32_epi32(pl, i0));
ST(d + 8, _mm256_permutevar8x32_epi32(ph, i0));
ST(d + 16, _mm256_permutevar8x32_epi32(pl, i1));
ST(d + 24, _mm256_permutevar8x32_epi32(ph, i1));
ST(d + 32, _mm256_permutevar8x32_epi32(pl, i4));
ST(d + 40, _mm256_permutevar8x32_epi32(ph, i4));
ST(d + 48, _mm256_permutevar8x32_epi32(pl, i5));
ST(d + 56, _mm256_permutevar8x32_epi32(ph, i5));
}
}
}
AVX2 static inline void packB2(int jc, int kk, int n, int ldb, const T *B, int sg, const T *B2) {
const __m256i m16 = _mm256_set1_epi32(0xFFFF);
const __m256i half = _mm256_set1_epi32(0x8000);
(void)n;
const size_t DSTR = (size_t)(NPAIR * 16 + 16);
const int sgn = sg > 0;
for (int m = 0; m < NPAIR; m++) {
const T *s0 = B + (size_t)(kk + 2 * m) * ldb + jc;
const T *s1 = s0 + ldb;
const T *t0 = B2 ? (B2 + (size_t)(kk + 2 * m) * ldb + jc) : 0;
const T *t1 = t0 ? t0 + ldb : 0;
T *d0 = Bpk + (size_t)m * 16;
T *d1 = d0 + 8;
for (int j = 0; j < NC; j += 8, d0 += DSTR, d1 += DSTR) {
__m256i x0 = LD(s0 + j), x1 = LD(s1 + j);
if (t0) {
__m256i y0 = LD(t0 + j), y1 = LD(t1 + j);
if (sgn) { x0 = _mm256_add_epi32(x0, y0); x1 = _mm256_add_epi32(x1, y1); }
else { x0 = _mm256_sub_epi32(x0, y0); x1 = _mm256_sub_epi32(x1, y1); }
}
__m256i l0 = _mm256_and_si256(x0, m16);
__m256i l1 = _mm256_and_si256(x1, m16);
__m256i h0 = _mm256_srli_epi32(_mm256_add_epi32(x0, half), 16);
__m256i h1 = _mm256_srli_epi32(_mm256_add_epi32(x1, half), 16);
ST(d0, _mm256_or_si256(l0, _mm256_slli_epi32(l1, 16)));
ST(d1, _mm256_or_si256(h0, _mm256_slli_epi32(h1, 16)));
}
}
}
AVX2 __attribute__((always_inline)) static inline void block_mikro(int ic, int jc, int n, int ldc, T *C, int first) {
T *cp0 = C + (size_t)ic * ldc + jc;
for (int r = 0; r < MC; r += MR) {
const int32_t *ap = Apk + (size_t)r * APSTR;
T *cp = cp0 + (size_t)r * ldc;
const int32_t *bp = Bpk;
for (int jg = 0; jg < NC; jg += NR) {
mikro(NPAIR, ap, bp, cp + jg, ldc, first);
bp += (NPAIR * 16 + 16);
}
}
}
/* LEAF SPECIALISED: with BASE==KC==NC==128 every tiled2 leaf call has h==128 and
ldc==h==128, so both become compile-time constants -> the jc/kk/ic loops fully fold. */
AVX2 static void tiled2_128(const T *A1, const T *A2, int sgA, int lda,
const T *B1, const T *B2, int sgB, int ldb, T *C) {
const int n = 128;
const int ldc = 128;
for (int jc = 0; jc < n; jc += NC) {
for (int kk = 0; kk < n; kk += KC) {
packB2(jc, kk, n, ldb, B1, sgB, B2);
for (int ic = 0; ic < n; ic += MC) {
packA2(ic, kk, n, lda, A1, sgA, A2);
block_mikro(ic, jc, n, ldc, C, kk == 0);
}
}
}
}
/* correct for any n, just slow */
AVX2 static void generic(int n, const T *A, int lda, const T *B, int ldb, T *C, int ldc) {
for (int i = 0; i < n; i++)
for (int j = 0; j < n; j++) {
uint32_t s = 0;
for (int k = 0; k < n; k++)
s += (uint32_t)A[(size_t)i * lda + k] * (uint32_t)B[(size_t)k * ldb + j];
C[(size_t)i * ldc + j] = (T)s;
}
}
/* tiled base case: requires n % NC == 0, n % KC == 0, n >= MC */
AVX2 static void tiled(int n, const T *A, int lda, const T *B, int ldb, T *C, int ldc) {
/* C is fully written by the kk==0 pass (micro stores instead of RMWs) */
for (int jc = 0; jc < n; jc += NC) {
for (int kk = 0; kk < n; kk += KC) {
packB(jc, kk, n, ldb, B);
for (int ic = 0; ic < n; ic += MC) {
packA(ic, kk, n, lda, A);
block_mikro(ic, jc, n, ldc, C, kk == 0);
}
}
}
}
AVX2 static int ok_tiled(int n) {
return (n % NC) == 0 && (n % KC) == 0 && (n % MC) == 0 && (MC % MR) == 0 && (NC % NR) == 0;
}
/* ---- Strassen over Z/2^32. Exact: all ops are ring ops of Z/2^32. ---- */
#define STR_DEPTH 6
/* P = X + s*Y in ONE pass over h x h (replaces copyblk + axpy: one write instead of two) */
AVX2 static void addblk(int h, T *P, const T *X, int ldx, const T *Y, int ldy, int s) {
for (int i = 0; i < h; i++) {
const T *xp = X + (size_t)i * ldx, *yp = Y + (size_t)i * ldy;
T *pp = P + (size_t)i * h;
__builtin_prefetch(X + (size_t)(i+2)*ldx, 0, 3); __builtin_prefetch(Y + (size_t)(i+2)*ldy, 0, 3);
int j = 0;
if (s > 0) for (; j + 8 <= h; j += 8) STA(pp + j, _mm256_add_epi32(LDA(xp + j), LDA(yp + j)));
else for (; j + 8 <= h; j += 8) STA(pp + j, _mm256_sub_epi32(LDA(xp + j), LDA(yp + j)));
for (; j < h; j++) pp[j] = (T)(s > 0 ? xp[j] + yp[j] : xp[j] - yp[j]);
}
}
/* dst = s1*m1 + s2*m2 , ONE pass (M blocks are contiguous, row stride h) */
AVX2 static void comb2(int h, T *dst, int ldd, const T *m1, int s1, const T *m2, int s2) {
(void)s1;
for (int i = 0; i < h; i++) {
const T *p1 = m1 + (size_t)i * h, *p2 = m2 + (size_t)i * h;
T *dp = dst + (size_t)i * ldd;
int j = 0;
for (; j + 8 <= h; j += 8) {
__m256i v = LDA(p1 + j);
v = (s2 > 0) ? _mm256_add_epi32(v, LDA(p2 + j)) : _mm256_sub_epi32(v, LDA(p2 + j));
_mm256_stream_si256((__m256i*)(dp + j), v);
}
for (; j < h; j++) dp[j] = (T)(p1[j] + (s2 > 0 ? p2[j] : -p2[j]));
}
}
/* dst = s1*m1 + s2*m2 + s3*m3 + s4*m4 , ONE pass */
AVX2 static void comb4(int h, T *dst, int ldd, const T *m1, int s1, const T *m2, int s2,
const T *m3, int s3, const T *m4, int s4) {
(void)s1;
for (int i = 0; i < h; i++) {
const T *p1 = m1 + (size_t)i * h, *p2 = m2 + (size_t)i * h,
*p3 = m3 + (size_t)i * h, *p4 = m4 + (size_t)i * h;
T *dp = dst + (size_t)i * ldd;
int j = 0;
for (; j + 8 <= h; j += 8) {
__m256i v = LDA(p1 + j);
v = (s2 > 0) ? _mm256_add_epi32(v, LDA(p2 + j)) : _mm256_sub_epi32(v, LDA(p2 + j));
v = (s3 > 0) ? _mm256_add_epi32(v, LDA(p3 + j)) : _mm256_sub_epi32(v, LDA(p3 + j));
v = (s4 > 0) ? _mm256_add_epi32(v, LDA(p4 + j)) : _mm256_sub_epi32(v, LDA(p4 + j));
_mm256_stream_si256((__m256i*)(dp + j), v);
}
for (; j < h; j++)
dp[j] = (T)(p1[j] + (s2>0?p2[j]:-p2[j]) + (s3>0?p3[j]:-p3[j]) + (s4>0?p4[j]:-p4[j]));
}
}
#define B11 (A)
#define B12 (A + h)
#define B21 (A + (size_t)h * lda)
#define B22 (A + (size_t)h * lda + h)
#define D11 (B)
#define D12 (B + h)
#define D21 (B + (size_t)h * ldb)
#define D22 (B + (size_t)h * ldb + h)
/* ---- FUSED combine: one pass instead of four (l_sub1pct) --------------------
The shipped rec() makes FOUR passes over the seven M blocks (12 block-reads
for 4 outputs: M1..M5 are each read twice). Holding the seven row vectors in
registers lets all four quadrants be formed from 7 loads, so the combine's
issue cost per 8 outputs goes 24 -> 19 uops and its read volume 12h^2 -> 7h^2.
Same stores, same NT form, byte-exact. */
static inline void STPAIR(int nt, T *P, __m256i u, __m256i v) {
_mm256_stream_si256((__m256i *)(P), u); _mm256_stream_si256((__m256i *)(P) + 1, v);
}
AVX2 static void comb_all(int h, T *C, int ldc,
const T *M1, const T *M2, const T *M3, const T *M4,
const T *M5, const T *M6, const T *M7) {
for (int i = 0; i < h; i++) {
const T *a = M1 + (size_t)i * h, *b = M2 + (size_t)i * h, *c = M3 + (size_t)i * h,
*d = M4 + (size_t)i * h, *e = M5 + (size_t)i * h, *f = M6 + (size_t)i * h,
*g = M7 + (size_t)i * h;
T *d1 = C + (size_t)i * ldc, *d2 = d1 + h,
*d3 = d1 + (size_t)h * ldc, *d4 = d3 + h;
int j = 0;
for (; j + 16 <= h; j += 16) {
__builtin_prefetch(a + j + 256, 0, 3); __builtin_prefetch(b + j + 256, 0, 3);
__builtin_prefetch(c + j + 256, 0, 3); __builtin_prefetch(d + j + 256, 0, 3);
__builtin_prefetch(e + j + 256, 0, 3); __builtin_prefetch(f + j + 256, 0, 3);
__builtin_prefetch(g + j + 256, 0, 3);
__m256i a1 = LDA(a + j), a2 = LDA(a + j + 8),
b1 = LDA(b + j), b2 = LDA(b + j + 8),
c1 = LDA(c + j), c2 = LDA(c + j + 8),
d1v = LDA(d + j), d2v = LDA(d + j + 8),
e1 = LDA(e + j), e2 = LDA(e + j + 8),
f1 = LDA(f + j), f2 = LDA(f + j + 8),
g1 = LDA(g + j), g2 = LDA(g + j + 8);
STPAIR(NTSEL, d1 + j, _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a1, d1v), e1), g1),
_mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a2, d2v), e2), g2));
STPAIR(NTSEL, d2 + j, _mm256_add_epi32(c1, e1), _mm256_add_epi32(c2, e2));
STPAIR(NTSEL, d3 + j, _mm256_add_epi32(b1, d1v), _mm256_add_epi32(b2, d2v));
STPAIR(NTSEL, d4 + j, _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a1, c1), b1), f1),
_mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a2, c2), b2), f2));
}
for (; j + 8 <= h; j += 8) {
__m256i m1 = LDA(a + j), m2 = LDA(b + j), m3 = LDA(c + j), m4 = LDA(d + j),
m5 = LDA(e + j), m6 = LDA(f + j), m7 = LDA(g + j);
_mm256_stream_si256((__m256i *)(d1 + j),
_mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(m1, m4), m5), m7));
_mm256_stream_si256((__m256i *)(d2 + j), _mm256_add_epi32(m3, m5));
_mm256_stream_si256((__m256i *)(d3 + j), _mm256_add_epi32(m2, m4));
_mm256_stream_si256((__m256i *)(d4 + j),
_mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(m1, m3), m2), m6));
}
for (; j < h; j++) {
d1[j] = a[j] + d[j] - e[j] + g[j];
d2[j] = c[j] + e[j];
d3[j] = b[j] + d[j];
d4[j] = a[j] - b[j] + c[j] + f[j];
}
}
}
/* ---- WINOGRAD COMBINE: C11 = M1+M2 ; C12 = M1+M3+M5+M6 ; C21 = M1+M4+M6+M7 ;
C22 = M1+M5+M6+M7. Same 8 additions and same 7 block-reads as comb_all;
only the association differs (the shared M1+M6 form is computed once). */
AVX2 static void comb_all_w(int h, T *C, int ldc,
const T *M1, const T *M2, const T *M3, const T *M4,
const T *M5, const T *M6, const T *M7) {
for (int i = 0; i < h; i++) {
const T *a = M1 + (size_t)i * h, *b = M2 + (size_t)i * h, *c = M3 + (size_t)i * h,
*d = M4 + (size_t)i * h, *e = M5 + (size_t)i * h, *f = M6 + (size_t)i * h,
*g = M7 + (size_t)i * h;
T *d1 = C + (size_t)i * ldc, *d2 = d1 + h,
*d3 = d1 + (size_t)h * ldc, *d4 = d3 + h;
int j = 0;
for (; j + 16 <= h; j += 16) {
__builtin_prefetch(a + j + 256, 0, 3); __builtin_prefetch(b + j + 256, 0, 3);
__builtin_prefetch(c + j + 256, 0, 3); __builtin_prefetch(d + j + 256, 0, 3);
__builtin_prefetch(e + j + 256, 0, 3); __builtin_prefetch(f + j + 256, 0, 3);
__builtin_prefetch(g + j + 256, 0, 3);
__m256i a1 = LDA(a + j), a2 = LDA(a + j + 8),
b1 = LDA(b + j), b2 = LDA(b + j + 8),
c1 = LDA(c + j), c2 = LDA(c + j + 8),
d1v = LDA(d + j), d2v = LDA(d + j + 8),
e1 = LDA(e + j), e2 = LDA(e + j + 8),
f1 = LDA(f + j), f2 = LDA(f + j + 8),
g1 = LDA(g + j), g2 = LDA(g + j + 8);
__m256i w1 = _mm256_add_epi32(a1, f1), w2 = _mm256_add_epi32(a2, f2);
STPAIR(NTSEL, d1 + j, _mm256_add_epi32(a1, b1), _mm256_add_epi32(a2, b2));
STPAIR(NTSEL, d2 + j, _mm256_add_epi32(_mm256_add_epi32(w1, c1), e1),
_mm256_add_epi32(_mm256_add_epi32(w2, c2), e2));
STPAIR(NTSEL, d3 + j, _mm256_add_epi32(_mm256_add_epi32(w1, d1v), g1),
_mm256_add_epi32(_mm256_add_epi32(w2, d2v), g2));
STPAIR(NTSEL, d4 + j, _mm256_add_epi32(_mm256_add_epi32(w1, e1), g1),
_mm256_add_epi32(_mm256_add_epi32(w2, e2), g2));
}
for (; j + 8 <= h; j += 8) {
__m256i m1 = LDA(a + j), m2 = LDA(b + j), m3 = LDA(c + j), m4 = LDA(d + j),
m5 = LDA(e + j), m6 = LDA(f + j), m7 = LDA(g + j);
__m256i w = _mm256_add_epi32(m1, m6);
_mm256_stream_si256((__m256i *)(d1 + j), _mm256_add_epi32(m1, m2));
_mm256_stream_si256((__m256i *)(d2 + j), _mm256_add_epi32(_mm256_add_epi32(w, m3), m5));
_mm256_stream_si256((__m256i *)(d3 + j), _mm256_add_epi32(_mm256_add_epi32(w, m4), m7));
_mm256_stream_si256((__m256i *)(d4 + j), _mm256_add_epi32(_mm256_add_epi32(w, m5), m7));
}
for (; j < h; j++) {
d1[j] = a[j] + b[j];
d2[j] = a[j] + c[j] + e[j] + f[j];
d3[j] = a[j] + d[j] + f[j] + g[j];
d4[j] = a[j] + e[j] + f[j] + g[j];
}
}
}
/* ============ PER-LOP SPECIALISATION: constant signs, constant 2nd-operand presence ============ */
template<int SA>
AVX2 static inline void packA2_t(int ic, int kk, int n, int lda, const T *A, const T *A2) {
const __m256i m16 = _mm256_set1_epi32(0xFFFF);
const __m256i half = _mm256_set1_epi32(0x8000);
const __m256i i0 = _mm256_set1_epi32(0), i1 = _mm256_set1_epi32(1);
const __m256i i4 = _mm256_set1_epi32(4), i5 = _mm256_set1_epi32(5);
(void)n;
for (int r = 0; r < MC; r++) {
const T *src = A + (size_t)(ic + r) * lda + kk;
const T *src2 = SA ? (A2 + (size_t)(ic + r) * lda + kk) : (const T *)0;
T *dst = Apk + (size_t)r * APSTR;
for (int m = 0; m < KC / 8; m++) {
__builtin_prefetch(src + (size_t)lda + 8 * m, 0, 1);
__builtin_prefetch(src + (size_t)2 * lda + 8 * m, 0, 1);
__m256i x = LD(src + 8 * m);
if (SA) { __builtin_prefetch(src2 + (size_t)lda + 8 * m, 0, 1); __m256i y = LD(src2 + 8 * m);
x = (SA > 0) ? _mm256_add_epi32(x, y) : _mm256_sub_epi32(x, y); }
__m256i al = _mm256_and_si256(x, m16);
__m256i ah = _mm256_srli_epi32(_mm256_add_epi32(x, half), 16);
__m256i pl = _mm256_packus_epi32(al, al);
__m256i ph = _mm256_packus_epi32(ah, ah);
T *d = dst + 64 * m;
ST(d + 0, _mm256_permutevar8x32_epi32(pl, i0));
ST(d + 8, _mm256_permutevar8x32_epi32(ph, i0));
ST(d + 16, _mm256_permutevar8x32_epi32(pl, i1));
ST(d + 24, _mm256_permutevar8x32_epi32(ph, i1));
ST(d + 32, _mm256_permutevar8x32_epi32(pl, i4));
ST(d + 40, _mm256_permutevar8x32_epi32(ph, i4));
ST(d + 48, _mm256_permutevar8x32_epi32(pl, i5));
ST(d + 56, _mm256_permutevar8x32_epi32(ph, i5));
}
}
}
template<int SB>
AVX2 static inline void packB2_t(int jc, int kk, int n, int ldb, const T *B, const T *B2) {
const __m256i m16 = _mm256_set1_epi32(0xFFFF);
const __m256i half = _mm256_set1_epi32(0x8000);
(void)n;
const size_t DSTR = (size_t)(NPAIR * 16 + 16);
for (int m = 0; m < NPAIR; m++) {
const T *s0 = B + (size_t)(kk + 2 * m) * ldb + jc;
const T *s1 = s0 + ldb;
const T *t0 = SB ? (B2 + (size_t)(kk + 2 * m) * ldb + jc) : (const T *)0;
const T *t1 = t0 ? t0 + ldb : (const T *)0;
T *d0 = Bpk + (size_t)m * 16;
T *d1 = d0 + 8;
for (int j = 0; j < NC; j += 8, d0 += DSTR, d1 += DSTR) {
__m256i x0 = LD(s0 + j), x1 = LD(s1 + j);
if (SB) {
__m256i y0 = LD(t0 + j), y1 = LD(t1 + j);
if (SB > 0) { x0 = _mm256_add_epi32(x0, y0); x1 = _mm256_add_epi32(x1, y1); }
else { x0 = _mm256_sub_epi32(x0, y0); x1 = _mm256_sub_epi32(x1, y1); }
}
__m256i l0 = _mm256_and_si256(x0, m16);
__m256i l1 = _mm256_and_si256(x1, m16);
__m256i h0 = _mm256_srli_epi32(_mm256_add_epi32(x0, half), 16);
__m256i h1 = _mm256_srli_epi32(_mm256_add_epi32(x1, half), 16);
ST(d0, _mm256_or_si256(l0, _mm256_slli_epi32(l1, 16)));
ST(d1, _mm256_or_si256(h0, _mm256_slli_epi32(h1, 16)));
}
}
}
template<int SA, int SB>
AVX2 static void tiled2_128_t(const T *A1, const T *A2, int lda,
const T *B1, const T *B2, int ldb, T *C) {
const int n = 128;
const int ldc = 128;
for (int jc = 0; jc < n; jc += NC) {
for (int kk = 0; kk < n; kk += KC) {
packB2_t<SB>(jc, kk, n, ldb, B1, B2);
for (int ic = 0; ic < n; ic += MC) {
packA2_t<SA>(ic, kk, n, lda, A1, A2);
block_mikro(ic, jc, n, ldc, C, kk == 0);
}
}
}
}
AVX2 static void rec(int n, const T *A, int lda, const T *B, int ldb, T *C, int ldc,
T *scr, int depth) {
if (depth >= STR_DEPTH || (n & 1) || n <= BASE || !ok_tiled(n)) {
if (ok_tiled(n)) tiled(n, A, lda, B, ldb, C, ldc);
else generic(n, A, lda, B, ldb, C, ldc);
return;
}
const int h = n / 2;
const size_t h2 = (size_t)h * h;
T *S = scr; /* 7 result blocks M1..M7 then 2 scratch S blocks */
T *M1 = S, *M2 = M1 + h2 + 4608, *M3 = M2 + h2 + 4608, *M4 = M3 + h2 + 4608,
*M5 = M4 + h2 + 4608, *M6 = M5 + h2 + 4608, *M7 = M6 + h2 + 4608;
T *P = M7 + h2 + 4608; /* operand scratch block 1 (h x h, ldp = h) */
T *Q = P + h2 + 4608; /* operand scratch block 2 */
T *sub = Q + h2 + 4608; /* recursion arena */
/* Each of the 7 products uses an operand that is A1 (+/- A2) on the A side and
B1 (+/- B2) on the B side. When the children are LEAVES (h <= BASE) the summed
operand is consumed only by the packers, so we fuse the sum INTO the packer and
never materialise P/Q: that removes 2 of every 4 accesses per summed operand
element (addblk's write + the packer's read of it, replaced by a second load of
the original block, which is exactly what addblk did anyway). */
if (h <= BASE) {
#define LOP(a1, a2, sgA, b1, b2, sgB, Mi) \
do { \
if (h <= BASE) { \
tiled2_128((a1), (a2), (sgA), lda, \
(b1), (b2), (sgB), ldb, (Mi)); \
} else { \
const T *pa = (a1); int la = lda; \
const T *pb = (b1); int lb = ldb; \
if (a2) { addblk(h, P, (a1), lda, (a2), lda, (sgA)); pa = P; la = h; } \
if (b2) { addblk(h, Q, (b1), ldb, (b2), ldb, (sgB)); pb = Q; lb = h; } \
rec(h, pa, la, pb, lb, (Mi), h, sub, depth + 1); \
} \
} while (0)
do { if (h <= BASE) { tiled2_128_t<1,1>((B11), (B22), lda, (D11), (D22), ldb, (M1)); } else { const T *pa = (B11); int la = lda; const T *pb = (D11); int lb = ldb; if (1) { addblk(h, P, (B11), lda, (B22), lda, 1); pa = P; la = h; } if (1) { addblk(h, Q, (D11), ldb, (D22), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M1), h, sub, depth + 1); } } while (0);
do { if (h <= BASE) { tiled2_128_t<1,0>((B21), (B22), lda, (D11), (0), ldb, (M2)); } else { const T *pa = (B21); int la = lda; const T *pb = (D11); int lb = ldb; if (1) { addblk(h, P, (B21), lda, (B22), lda, 1); pa = P; la = h; } if (0) { addblk(h, Q, (D11), ldb, (D11), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M2), h, sub, depth + 1); } } while (0);
do { if (h <= BASE) { tiled2_128_t<0,-1>((B11), (0), lda, (D12), (D22), ldb, (M3)); } else { const T *pa = (B11); int la = lda; const T *pb = (D12); int lb = ldb; if (0) { addblk(h, P, (B11), lda, (B11), lda, 1); pa = P; la = h; } if (1) { addblk(h, Q, (D12), ldb, (D22), ldb, -1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M3), h, sub, depth + 1); } } while (0);
do { if (h <= BASE) { tiled2_128_t<0,-1>((B22), (0), lda, (D21), (D11), ldb, (M4)); } else { const T *pa = (B22); int la = lda; const T *pb = (D21); int lb = ldb; if (0) { addblk(h, P, (B22), lda, (B22), lda, 1); pa = P; la = h; } if (1) { addblk(h, Q, (D21), ldb, (D11), ldb, -1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M4), h, sub, depth + 1); } } while (0);
do { if (h <= BASE) { tiled2_128_t<1,0>((B11), (B12), lda, (D22), (0), ldb, (M5)); } else { const T *pa = (B11); int la = lda; const T *pb = (D22); int lb = ldb; if (1) { addblk(h, P, (B11), lda, (B12), lda, 1); pa = P; la = h; } if (0) { addblk(h, Q, (D22), ldb, (D22), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M5), h, sub, depth + 1); } } while (0);
do { if (h <= BASE) { tiled2_128_t<-1,1>((B21), (B11), lda, (D11), (D12), ldb, (M6)); } else { const T *pa = (B21); int la = lda; const T *pb = (D11); int lb = ldb; if (1) { addblk(h, P, (B21), lda, (B11), lda, -1); pa = P; la = h; } if (1) { addblk(h, Q, (D11), ldb, (D12), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M6), h, sub, depth + 1); } } while (0);
do { if (h <= BASE) { tiled2_128_t<-1,1>((B12), (B22), lda, (D21), (D22), ldb, (M7)); } else { const T *pa = (B12); int la = lda; const T *pb = (D21); int lb = ldb; if (1) { addblk(h, P, (B12), lda, (B22), lda, -1); pa = P; la = h; } if (1) { addblk(h, Q, (D21), ldb, (D22), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M7), h, sub, depth + 1); } } while (0);
#undef LOP
/* C11 = M1 + M4 - M5 + M7 ; C12 = M3 + M5 ; C21 = M2 + M4 ; C22 = M1 - M2 + M3 + M6 */
comb_all(h, C, ldc, M1, M2, M3, M4, M5, M6, M7);
} else {
/* ---- STRASSEN-WINOGRAD: 8 materialised operand passes per node (classical: 10) ----
The chains run IN PLACE and the child calls are ordered so that only one A-side
and one B-side form are live at a time -> the SAME two scratch blocks P and Q
suffice, and the arena layout is unchanged. The leaf path (h <= BASE) above is
the shipped classical bytes, untouched. */
addblk(h, P, (B21), lda, (B22), lda, 1); /* S1 = A21 + A22 */
addblk(h, Q, (D12), ldb, (D11), ldb, -1); /* T1 = B12 - B11 */
rec(h, P, h, Q, h, (M5), h, sub, depth + 1); /* M5 = S1 * T1 */
addblk(h, P, P, h, (B11), lda, -1); /* S2 = S1 - A11 (ip) */
addblk(h, Q, (D22), ldb, Q, h, -1); /* T2 = B22 - T1 (ip) */
rec(h, P, h, Q, h, (M6), h, sub, depth + 1); /* M6 = S2 * T2 */
addblk(h, Q, (D21), ldb, Q, h, -1); /* T4 = B21 - T2 (ip) */
rec(h, (B22), lda, Q, h, (M4), h, sub, depth + 1); /* M4 = A22 * T4 */
addblk(h, P, (B12), lda, P, h, -1); /* S4 = A12 - S2 (ip) */
rec(h, P, h, (D22), ldb, (M3), h, sub, depth + 1); /* M3 = S4 * B22 */
addblk(h, P, (B11), lda, (B21), lda, -1); /* S3 = A11 - A21 */
addblk(h, Q, (D22), ldb, (D12), ldb, -1); /* T3 = B22 - B12 */
rec(h, P, h, Q, h, (M7), h, sub, depth + 1); /* M7 = S3 * T3 */
rec(h, (B11), lda, (D11), ldb, (M1), h, sub, depth + 1); /* M1 = A11 * B11 */
rec(h, (B12), lda, (D21), ldb, (M2), h, sub, depth + 1); /* M2 = A12 * B21 */
comb_all_w(h, C, ldc, M1, M2, M3, M4, M5, M6, M7);
}
}
AVX2 void matrix_multiply(int n, const T *A, const T *B, T *C) {
unsigned long long al = (unsigned long long)(size_t)A | (unsigned long long)(size_t)B | (unsigned long long)(size_t)C;
if (n <= BASE || !ok_tiled(n) || (al & 31)) {
if (ok_tiled(n) && !(al & 31)) tiled(n, A, n, B, n, C, n);
else generic(n, A, n, B, n, C, n);
return;
}
/* depth-first arena: 9 h^2 blocks at each level along the path */
size_t need = 0, h = (size_t)n;
while (h > (size_t)BASE) { h /= 2; need += 9 * (h * h + 4608); }
static T *scr = 0; static size_t cap = 0;
if (need > cap) { if (scr) free(scr); cap = need;
void *pv = 0; if (posix_memalign(&pv, 4096, cap * sizeof(T))) return; scr = (T *)pv; }
rec(n, A, n, B, n, C, n, scr, 0);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 1.778 s | 255 MB + 896 KB | Accepted | Score: 100 | 显示更多 |