#pragma GCC optimize("O3,unroll-loops")
#pragma GCC target("avx2")
#include <immintrin.h>
#include <stdint.h>
#include <string.h>
#ifndef BASE
#define BASE 64
#endif
using U = uint16_t;
static constexpr int MAXN = 4096;
alignas(4096) static U pa[MAXN * MAXN], pb[MAXN * MAXN];
alignas(4096) static U pc[MAXN * MAXN], ws[MAXN * MAXN];
static inline __m256i ld(const U *p) {
return _mm256_load_si256((const __m256i *)p);
}
static inline void st(U *p, __m256i x) {
_mm256_store_si256((__m256i *)p, x);
}
// The recursive quadrant layout makes every Strassen addition contiguous.
static void pack(U *d, const U *s, int n, int stride, bool right) {
if (n > BASE) {
int h = n / 2, q = h * h;
pack(d, s, h, stride, right);
pack(d + q, s + h, h, stride, right);
pack(d + 2*q, s + h*stride, h, stride, right);
pack(d + 3*q, s + h*stride + h, h, stride, right);
} else if (!right) {
for(int i=0;i<n;i+=4) for(int k=0;k<n;k+=2)
for(int r=0;r<4;++r) {
memcpy(d,s+(i+r)*stride+k,4); d+=2;
}
} else {
for(int j=0;j<n;) {
int w=n-j>=24?24:n-j;
for(int k=0;k<n;k+=2) {
for(int t=0;t<w;t+=8) {
__m128i x=_mm_load_si128((const __m128i*)(s+k*stride+j+t));
__m128i y=_mm_load_si128((const __m128i*)(s+(k+1)*stride+j+t));
__m128i lo=_mm_unpacklo_epi16(x,y),hi=_mm_unpackhi_epi16(x,y);
_mm_store_si128((__m128i*)d,lo);
_mm_store_si128((__m128i*)(d+8),hi);
d+=16;
}
}
j+=w;
}
}
}
static void unpack(U *d, const U *s, int n, int stride) {
if (n > BASE) {
int h = n/2, q = h*h;
unpack(d,s,h,stride);
unpack(d+h,s+q,h,stride);
unpack(d+h*stride,s+2*q,h,stride);
unpack(d+h*stride+h,s+3*q,h,stride);
} else {
for (int i=0;i<n;++i) memcpy(d+i*stride,s+i*n,n*sizeof(U));
}
}
static inline void micro24(const U *a,const U *b,U *c,int n) {
const U *end=a+4*n; U *c3=c+3*n; intptr_t stride=2*n;
asm volatile(
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
"vpxor %%ymm2, %%ymm2, %%ymm2\n\t"
"vpxor %%ymm3, %%ymm3, %%ymm3\n\t"
"vpxor %%ymm4, %%ymm4, %%ymm4\n\t"
"vpxor %%ymm5, %%ymm5, %%ymm5\n\t"
"vpxor %%ymm6, %%ymm6, %%ymm6\n\t"
"vpxor %%ymm7, %%ymm7, %%ymm7\n\t"
"vpxor %%ymm8, %%ymm8, %%ymm8\n\t"
"vpxor %%ymm9, %%ymm9, %%ymm9\n\t"
"vpxor %%ymm10, %%ymm10, %%ymm10\n\t"
"vpxor %%ymm11, %%ymm11, %%ymm11\n\t"
"1:\n\t"
"vmovdqa 0(%[b]), %%ymm12\n\t"
"vmovdqa 32(%[b]), %%ymm13\n\t"
"vpbroadcastd 0(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 4(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 8(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
"vpbroadcastd 12(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
"vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
"vmovdqa 96(%[b]), %%ymm12\n\t"
"vmovdqa 128(%[b]), %%ymm13\n\t"
"vpbroadcastd 16(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 20(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 24(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
"vpbroadcastd 28(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
"vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
"vmovdqa 192(%[b]), %%ymm12\n\t"
"vmovdqa 224(%[b]), %%ymm13\n\t"
"vpbroadcastd 32(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 36(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 40(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
"vpbroadcastd 44(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
"vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
"vmovdqa 288(%[b]), %%ymm12\n\t"
"vmovdqa 320(%[b]), %%ymm13\n\t"
"vpbroadcastd 48(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 52(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 56(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
"vpbroadcastd 60(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
"vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
"add $64, %[a]\n\t"
"add $384, %[b]\n\t"
"cmp %[end], %[a]\n\t"
"jne 1b\n\t"
"vpcmpeqd %%ymm15, %%ymm15, %%ymm15\n\t"
"vpsrld $16, %%ymm15, %%ymm15\n\t"
"vpand %%ymm15, %%ymm0, %%ymm0\n\t"
"vpand %%ymm15, %%ymm1, %%ymm1\n\t"
"vpand %%ymm15, %%ymm2, %%ymm2\n\t"
"vpackusdw %%ymm1, %%ymm0, %%ymm0\n\t"
"vpermq $216, %%ymm0, %%ymm0\n\t"
"vmovdqu %%ymm0, 0(%[c])\n\t"
"vpackusdw %%ymm2, %%ymm2, %%ymm2\n\t"
"vpermq $216, %%ymm2, %%ymm2\n\t"
"vmovdqu %%xmm2, 32(%[c])\n\t"
"vpand %%ymm15, %%ymm3, %%ymm3\n\t"
"vpand %%ymm15, %%ymm4, %%ymm4\n\t"
"vpand %%ymm15, %%ymm5, %%ymm5\n\t"
"vpackusdw %%ymm4, %%ymm3, %%ymm3\n\t"
"vpermq $216, %%ymm3, %%ymm3\n\t"
"vmovdqu %%ymm3, 0(%[c],%[s],1)\n\t"
"vpackusdw %%ymm5, %%ymm5, %%ymm5\n\t"
"vpermq $216, %%ymm5, %%ymm5\n\t"
"vmovdqu %%xmm5, 32(%[c],%[s],1)\n\t"
"vpand %%ymm15, %%ymm6, %%ymm6\n\t"
"vpand %%ymm15, %%ymm7, %%ymm7\n\t"
"vpand %%ymm15, %%ymm8, %%ymm8\n\t"
"vpackusdw %%ymm7, %%ymm6, %%ymm6\n\t"
"vpermq $216, %%ymm6, %%ymm6\n\t"
"vmovdqu %%ymm6, 0(%[c],%[s],2)\n\t"
"vpackusdw %%ymm8, %%ymm8, %%ymm8\n\t"
"vpermq $216, %%ymm8, %%ymm8\n\t"
"vmovdqu %%xmm8, 32(%[c],%[s],2)\n\t"
"vpand %%ymm15, %%ymm9, %%ymm9\n\t"
"vpand %%ymm15, %%ymm10, %%ymm10\n\t"
"vpand %%ymm15, %%ymm11, %%ymm11\n\t"
"vpackusdw %%ymm10, %%ymm9, %%ymm9\n\t"
"vpermq $216, %%ymm9, %%ymm9\n\t"
"vmovdqu %%ymm9, 0(%[c3])\n\t"
"vpackusdw %%ymm11, %%ymm11, %%ymm11\n\t"
"vpermq $216, %%ymm11, %%ymm11\n\t"
"vmovdqu %%xmm11, 32(%[c3])\n\t"
: [a] "+&r"(a), [b] "+&r"(b)
: [c] "r"(c), [c3] "r"(c3), [s] "r"(stride), [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5",
"ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm11", "ymm12",
"ymm13", "ymm14", "ymm15");
}
static inline void micro16(const U *a,const U *b,U *c,int n) {
const U *end=a+4*n; U *c3=c+3*n; intptr_t stride=2*n;
asm volatile(
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
"vpxor %%ymm2, %%ymm2, %%ymm2\n\t"
"vpxor %%ymm3, %%ymm3, %%ymm3\n\t"
"vpxor %%ymm4, %%ymm4, %%ymm4\n\t"
"vpxor %%ymm5, %%ymm5, %%ymm5\n\t"
"vpxor %%ymm6, %%ymm6, %%ymm6\n\t"
"vpxor %%ymm7, %%ymm7, %%ymm7\n\t"
"1:\n\t"
"vmovdqa 0(%[b]), %%ymm12\n\t"
"vmovdqa 32(%[b]), %%ymm13\n\t"
"vpbroadcastd 0(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 4(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 8(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 12(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 64(%[b]), %%ymm12\n\t"
"vmovdqa 96(%[b]), %%ymm13\n\t"
"vpbroadcastd 16(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 20(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 24(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 28(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 128(%[b]), %%ymm12\n\t"
"vmovdqa 160(%[b]), %%ymm13\n\t"
"vpbroadcastd 32(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 36(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 40(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 44(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 192(%[b]), %%ymm12\n\t"
"vmovdqa 224(%[b]), %%ymm13\n\t"
"vpbroadcastd 48(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 52(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 56(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 60(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"add $64, %[a]\n\t"
"add $256, %[b]\n\t"
"cmp %[end], %[a]\n\t"
"jne 1b\n\t"
"vpcmpeqd %%ymm15, %%ymm15, %%ymm15\n\t"
"vpsrld $16, %%ymm15, %%ymm15\n\t"
"vpand %%ymm15, %%ymm0, %%ymm0\n\t"
"vpand %%ymm15, %%ymm1, %%ymm1\n\t"
"vpackusdw %%ymm1, %%ymm0, %%ymm0\n\t"
"vpermq $216, %%ymm0, %%ymm0\n\t"
"vmovdqu %%ymm0, 0(%[c])\n\t"
"vpand %%ymm15, %%ymm2, %%ymm2\n\t"
"vpand %%ymm15, %%ymm3, %%ymm3\n\t"
"vpackusdw %%ymm3, %%ymm2, %%ymm2\n\t"
"vpermq $216, %%ymm2, %%ymm2\n\t"
"vmovdqu %%ymm2, 0(%[c],%[s],1)\n\t"
"vpand %%ymm15, %%ymm4, %%ymm4\n\t"
"vpand %%ymm15, %%ymm5, %%ymm5\n\t"
"vpackusdw %%ymm5, %%ymm4, %%ymm4\n\t"
"vpermq $216, %%ymm4, %%ymm4\n\t"
"vmovdqu %%ymm4, 0(%[c],%[s],2)\n\t"
"vpand %%ymm15, %%ymm6, %%ymm6\n\t"
"vpand %%ymm15, %%ymm7, %%ymm7\n\t"
"vpackusdw %%ymm7, %%ymm6, %%ymm6\n\t"
"vpermq $216, %%ymm6, %%ymm6\n\t"
"vmovdqu %%ymm6, 0(%[c3])\n\t"
: [a] "+&r"(a), [b] "+&r"(b)
: [c] "r"(c), [c3] "r"(c3), [s] "r"(stride), [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5",
"ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm11", "ymm12",
"ymm13", "ymm14", "ymm15");
}
static inline void micro8(const U *a,const U *b,U *c,int n) {
const U *end=a+4*n; U *c3=c+3*n; intptr_t stride=2*n;
asm volatile(
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
"vpxor %%ymm2, %%ymm2, %%ymm2\n\t"
"vpxor %%ymm3, %%ymm3, %%ymm3\n\t"
"1:\n\t"
"vmovdqa 0(%[b]), %%ymm12\n\t"
"vpbroadcastd 0(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpbroadcastd 4(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 8(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 12(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vmovdqa 32(%[b]), %%ymm12\n\t"
"vpbroadcastd 16(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpbroadcastd 20(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 24(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 28(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vmovdqa 64(%[b]), %%ymm12\n\t"
"vpbroadcastd 32(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpbroadcastd 36(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 40(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 44(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"vmovdqa 96(%[b]), %%ymm12\n\t"
"vpbroadcastd 48(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
"vpbroadcastd 52(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 56(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
"vpbroadcastd 60(%[a]), %%ymm14\n\t"
"vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
"add $64, %[a]\n\t"
"add $128, %[b]\n\t"
"cmp %[end], %[a]\n\t"
"jne 1b\n\t"
"vpcmpeqd %%ymm15, %%ymm15, %%ymm15\n\t"
"vpsrld $16, %%ymm15, %%ymm15\n\t"
"vpand %%ymm15, %%ymm0, %%ymm0\n\t"
"vpackusdw %%ymm0, %%ymm0, %%ymm0\n\t"
"vpermq $216, %%ymm0, %%ymm0\n\t"
"vmovdqu %%xmm0, 0(%[c])\n\t"
"vpand %%ymm15, %%ymm1, %%ymm1\n\t"
"vpackusdw %%ymm1, %%ymm1, %%ymm1\n\t"
"vpermq $216, %%ymm1, %%ymm1\n\t"
"vmovdqu %%xmm1, 0(%[c],%[s],1)\n\t"
"vpand %%ymm15, %%ymm2, %%ymm2\n\t"
"vpackusdw %%ymm2, %%ymm2, %%ymm2\n\t"
"vpermq $216, %%ymm2, %%ymm2\n\t"
"vmovdqu %%xmm2, 0(%[c],%[s],2)\n\t"
"vpand %%ymm15, %%ymm3, %%ymm3\n\t"
"vpackusdw %%ymm3, %%ymm3, %%ymm3\n\t"
"vpermq $216, %%ymm3, %%ymm3\n\t"
"vmovdqu %%xmm3, 0(%[c3])\n\t"
: [a] "+&r"(a), [b] "+&r"(b)
: [c] "r"(c), [c3] "r"(c3), [s] "r"(stride), [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5",
"ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm11", "ymm12",
"ymm13", "ymm14", "ymm15");
}
static void base(const U *a,const U *b,U *c,int n) {
int j=0;
for(;j+24<=n;j+=24)
for(int i=0;i<n;i+=4) micro24(a+i*n,b+j*n,c+i*n+j,n);
if(n-j==16)
for(int i=0;i<n;i+=4) micro16(a+i*n,b+j*n,c+i*n+j,n);
if(n-j==8)
for(int i=0;i<n;i+=4) micro8(a+i*n,b+j*n,c+i*n+j,n);
}
template<bool Sub>
static inline void combine(U *d,const U *a,const U *b,int len) {
for(int i=0;i<len;i+=16) {
__m256i x=ld(a+i),y=ld(b+i);
st(d+i,Sub?_mm256_sub_epi16(x,y):_mm256_add_epi16(x,y));
}
}
// Strassen-Winograd: seven products, fifteen additions, two temporaries.
// All operations are in Z/(2^16); truncation never changes the answer.
static void mul(const U *a,const U *b,U *c,int n,U *work) {
if(n<=BASE) { base(a,b,c,n); return; }
int h=n/2,q=h*h;
const U *a11=a,*a12=a+q,*a21=a+2*q,*a22=a+3*q;
const U *b11=b,*b12=b+q,*b21=b+2*q,*b22=b+3*q;
U *c11=c,*c12=c+q,*c21=c+2*q,*c22=c+3*q;
U *x=work,*y=work+q,*next=work+2*q;
combine<true>(x,a11,a21,q);
combine<true>(y,b22,b12,q);
mul(x,y,c21,h,next); // P7
combine<false>(x,a21,a22,q);
combine<true>(y,b12,b11,q);
mul(x,y,c22,h,next); // P5
combine<true>(x,x,a11,q);
combine<true>(y,b22,y,q);
mul(x,y,c12,h,next); // P6
combine<true>(x,a12,x,q);
mul(x,b22,c11,h,next); // P3
combine<true>(y,y,b21,q);
mul(a22,y,x,h,next); // P4
mul(a11,b11,y,h,next); // P1
for(int i=0;i<q;i+=16) {
__m256i u2=_mm256_add_epi16(ld(y+i),ld(c12+i));
__m256i u3=_mm256_add_epi16(u2,ld(c21+i));
__m256i p5=ld(c22+i);
st(c12+i,_mm256_add_epi16(_mm256_add_epi16(u2,p5),ld(c11+i)));
st(c21+i,_mm256_sub_epi16(u3,ld(x+i)));
st(c22+i,_mm256_add_epi16(u3,p5));
}
mul(a12,b21,c11,h,next); // P2
combine<false>(c11,c11,y,q);
}
void matrix_multiply(int n,const short *A,const short *B,short *C) {
pack(pa,(const U *)A,n,n,false);
pack(pb,(const U *)B,n,n,true);
mul(pa,pb,pc,n,ws);
unpack((U *)C,pc,n,n);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 610.219 ms | 149 MB + 348 KB | Accepted | Score: 100 | 显示更多 |