#pragma GCC optimize("O3","unroll-loops")
#pragma GCC target("avx2")
/* mmml: C = A*B mod 2^64, long long, FUNCTION-style C++ linkage.
5-ALU-uop mikro (vpmuludq + vpaddq + vpmulld + vpaddd per k x row x 4 cols, no p5 uop)
plus Strassen with fused operand combination and dual-destination stores
(no extra memory). MR=6 KC=512 MC=48 NC=512 CUTL=512 U=1 BSW=0 EXP=0 */
#include <immintrin.h>
/* Single-instruction 256-bit unaligned access, NO alignment precondition.
g++-9 -O2 with no -march expands the unaligned 256-bit intrinsics into 2-3 uops, one
of them a p5-only vinsertf128/vextractf128; these emit the single instruction the
intrinsic is named for. AT&T order is `vmovdqu src,dst`: the STORE is %0,%1 (ymm
first, memory second). A REVERSED store operand silently emits a LOAD instead --
that is how this was first written and it produced wrong output with no diagnostic.
volatile + "memory" is REQUIRED: a non-volatile asm with only an "m" input is pure
and gets eliminated. */
#define LDQ(p) ({ __m256i _ldq_v; __asm__("vmovdqu %1,%0" : "=x"(_ldq_v) : "m"(*(const __m256i *)(p)) : "memory"); _ldq_v; })
#define STQ(p, v) __asm__ volatile("vmovdqu %0,%1" :: "x"(v), "m"(*(__m256i *)(p)) : "memory")
#include <stddef.h>
/* lane l2k12_rv PFW: PREFETCHW on the panel destination, issued ahead of the
stores. The panel lines are written once, fully, and the write-allocate RFO
(packA 3.96M + packB 8.70M lines) is the phase's critical path; issuing the
RFO early lets the store hit L1 instead of waiting for the line. NOTE:
__builtin_prefetch(p,1,N) is NOT a write prefetch (it emits prefetcht0) --
this is the raw instruction. */
#define PFW(p) __asm__ volatile("prefetchw %0" :: "m"(*(char *)(p)))
#pragma GCC optimize("O3","unroll-loops","O2","unroll-loops",)
#pragma GCC push_options
#pragma GCC target("avx2,bmi2")
typedef long long TE;
#define MR 6
#define KC 64
#define MC 12
#define NC 512
#define UU 1
#define CUTL 64
/* A-panel row stride PADDED by 8 uint64 (64 B). With the natural stride KC*8 = 4096 B
every one of the 6 rows of a micro-tile maps to the SAME L1 set (4096/64 = 64 sets
apart, 64 sets in the cache) -- a textbook conflict. +8 uint64 makes the stride
4160 B = 65 lines, so consecutive rows land in consecutive sets. */
#define KCP (KC + 8)
static TE Apanel[(size_t)MC * KCP + 64] __attribute__((aligned(4096)));
static TE Bpanel[(size_t)2 * (NC / 4 + 1) * KC * 8 + 64] __attribute__((aligned(4096)));
static __m256i AC[12] __attribute__((aligned(64)));
/* op: 0 = store, 1 = add, 2 = sub; c2 may be 0 (single destination) */
static inline void mikro(long cnt, const long long *ap, const long long *bp,
long long *c1, int ldc1, int op1,
long long *c2, int ldc2, int op2, int rows) {
long idx = -(cnt << 3);
long long *apE = (long long *)ap + cnt;
long long *bpE = (long long *)bp + (cnt << 2);
const long st1 = (long)ldc1 * 8, st2 = (long)ldc2 * 8;
if (rows == 6) {
__asm__ volatile(
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vmovdqa %%ymm0, %%ymm1\n\t"
"vmovdqa %%ymm0, %%ymm2\n\t"
"vmovdqa %%ymm0, %%ymm3\n\t"
"vmovdqa %%ymm0, %%ymm4\n\t"
"vmovdqa %%ymm0, %%ymm5\n\t"
"vmovdqa %%ymm0, %%ymm6\n\t"
"vmovdqa %%ymm0, %%ymm7\n\t"
"vmovdqa %%ymm0, %%ymm8\n\t"
"vmovdqa %%ymm0, %%ymm9\n\t"
"vmovdqa %%ymm0, %%ymm10\n\t"
"vmovdqa %%ymm0, %%ymm11\n\t"
"test %[idx], %[idx]\n\t"
"jz 9f\n\t"
"vmovdqu (%[bp],%[idx],4), %%ymm12\n\t"
"vpshufd $0xB1, %%ymm12, %%ymm13\n\t"
"vbroadcastsd 0(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm0\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm1\n\t"
"vbroadcastsd 576(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm2\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm3\n\t"
"vbroadcastsd 1152(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm4\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm5\n\t"
"vbroadcastsd 1728(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm6\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm7\n\t"
"vbroadcastsd 2304(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm8\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm9\n\t"
"vbroadcastsd 2880(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm10\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm11\n\t"
"vbroadcastsd 8(%[ap],%[idx],1), %%ymm14\n\t"
"vmovdqu 32(%[bp],%[idx],4), %%ymm12\n\t"
"add $8, %[idx]\n\t"
"jz 2f\n\t"
"1:\n\t"
"vpshufd $0xB1, %%ymm12, %%ymm13\n\t"
"vbroadcastsd 0(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm1, %%ymm1\n\t"
"vbroadcastsd 576(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm3, %%ymm3\n\t"
"vbroadcastsd 1152(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm2, %%ymm2\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm5, %%ymm5\n\t"
"vbroadcastsd 1728(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm7, %%ymm7\n\t"
"vbroadcastsd 2304(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm9, %%ymm9\n\t"
"vbroadcastsd 2880(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm8, %%ymm8\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm11, %%ymm11\n\t"
"vbroadcastsd 8(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm10, %%ymm10\n\t"
"vmovdqu 32(%[bp],%[idx],4), %%ymm12\n\t"
"add $8, %[idx]\n\t"
"jnz 1b\n\t"
"jmp 2f\n\t"
"9:\n\t"
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vmovdqa %%ymm0, %%ymm1\n\t"
"vmovdqa %%ymm0, %%ymm2\n\t"
"vmovdqa %%ymm0, %%ymm3\n\t"
"vmovdqa %%ymm0, %%ymm4\n\t"
"vmovdqa %%ymm0, %%ymm5\n\t"
"vmovdqa %%ymm0, %%ymm6\n\t"
"vmovdqa %%ymm0, %%ymm7\n\t"
"vmovdqa %%ymm0, %%ymm8\n\t"
"vmovdqa %%ymm0, %%ymm9\n\t"
"vmovdqa %%ymm0, %%ymm10\n\t"
"vmovdqa %%ymm0, %%ymm11\n\t"
"2:\n\t"
/* ---- E2 epilogue, 6 rows, ONE fixup shared by both destinations ---- */
"vpsrlq $32, %%ymm1, %%ymm12\n\t"
"vpaddd %%ymm1, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm0, %%ymm12, %%ymm0\n\t"
"vpsrlq $32, %%ymm3, %%ymm12\n\t"
"vpaddd %%ymm3, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm2, %%ymm12, %%ymm2\n\t"
"vpsrlq $32, %%ymm5, %%ymm12\n\t"
"vpaddd %%ymm5, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm4, %%ymm12, %%ymm4\n\t"
"vpsrlq $32, %%ymm7, %%ymm12\n\t"
"vpaddd %%ymm7, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm6, %%ymm12, %%ymm6\n\t"
"vpsrlq $32, %%ymm9, %%ymm12\n\t"
"vpaddd %%ymm9, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm8, %%ymm12, %%ymm8\n\t"
"vpsrlq $32, %%ymm11, %%ymm12\n\t"
"vpaddd %%ymm11, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm10, %%ymm12, %%ymm10\n\t"
"movq %[c1], %%r8\n\t"
"movq %[st1], %%r9\n\t"
"cmpl $1, %[op1]\n\t"
"je .LF6%=_1b\n\t"
"jg .LF6%=_1c\n\t"
"vmovdqu %%ymm0, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm2, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm4, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm6, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm8, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm10, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"jmp .LF6%=_1z\n\t"
".LF6%=_1b:\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm8, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm10, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"jmp .LF6%=_1z\n\t"
".LF6%=_1c:\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm8, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm10, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
".LF6%=_1z:\n\t"
"testq %[c2], %[c2]\n\t"
"jz .LF6%=_ez\n\t"
"movq %[c2], %%r10\n\t"
"movq %[st2], %%r11\n\t"
"cmpl $1, %[op2]\n\t"
"je .LF6%=_2b\n\t"
"jg .LF6%=_2c\n\t"
"vmovdqu %%ymm0, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm2, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm4, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm6, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm8, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm10, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"jmp .LF6%=_ez\n\t"
".LF6%=_2b:\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm8, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm10, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"jmp .LF6%=_ez\n\t"
".LF6%=_2c:\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm8, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm10, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
".LF6%=_ez:\n\t"
""
: [ap] "+r"(apE), [bp] "+r"(bpE), [idx] "+r"(idx)
: [c1] "r"(c1), [st1] "r"(st1), [op1] "r"(op1),
[c2] "r"(c2), [st2] "r"(st2), [op2] "r"(op2)
: "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7",
"ymm8", "ymm9", "ymm10", "ymm11", "ymm12", "ymm13", "ymm14", "ymm15",
"r8", "r9", "r10", "r11", "cc", "memory");
} else
if (rows == 4) {
__asm__ volatile(
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vmovdqa %%ymm0, %%ymm1\n\t"
"vmovdqa %%ymm0, %%ymm2\n\t"
"vmovdqa %%ymm0, %%ymm3\n\t"
"vmovdqa %%ymm0, %%ymm4\n\t"
"vmovdqa %%ymm0, %%ymm5\n\t"
"vmovdqa %%ymm0, %%ymm6\n\t"
"vmovdqa %%ymm0, %%ymm7\n\t"
"test %[idx], %[idx]\n\t"
"jz 9f\n\t"
"vmovdqu (%[bp],%[idx],4), %%ymm12\n\t"
"vpshufd $0xB1, %%ymm12, %%ymm13\n\t"
"vbroadcastsd 0(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm0\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm1\n\t"
"vbroadcastsd 576(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm2\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm3\n\t"
"vbroadcastsd 1152(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm4\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm5\n\t"
"vbroadcastsd 1728(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm6\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm7\n\t"
"vbroadcastsd 8(%[ap],%[idx],1), %%ymm14\n\t"
"vmovdqu 32(%[bp],%[idx],4), %%ymm12\n\t"
"add $8, %[idx]\n\t"
"jz 2f\n\t"
"1:\n\t"
"vpshufd $0xB1, %%ymm12, %%ymm13\n\t"
"vbroadcastsd 0(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm1, %%ymm1\n\t"
"vbroadcastsd 576(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm3, %%ymm3\n\t"
"vbroadcastsd 1152(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm2, %%ymm2\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm5, %%ymm5\n\t"
"vbroadcastsd 1728(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm7, %%ymm7\n\t"
"vbroadcastsd 8(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm6, %%ymm6\n\t"
"vmovdqu 32(%[bp],%[idx],4), %%ymm12\n\t"
"add $8, %[idx]\n\t"
"jnz 1b\n\t"
"jmp 2f\n\t"
"9:\n\t"
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vmovdqa %%ymm0, %%ymm1\n\t"
"vmovdqa %%ymm0, %%ymm2\n\t"
"vmovdqa %%ymm0, %%ymm3\n\t"
"vmovdqa %%ymm0, %%ymm4\n\t"
"vmovdqa %%ymm0, %%ymm5\n\t"
"vmovdqa %%ymm0, %%ymm6\n\t"
"vmovdqa %%ymm0, %%ymm7\n\t"
"2:\n\t"
/* ---- E2 epilogue, 4 rows, ONE fixup shared by both destinations ---- */
"vpsrlq $32, %%ymm1, %%ymm12\n\t"
"vpaddd %%ymm1, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm0, %%ymm12, %%ymm0\n\t"
"vpsrlq $32, %%ymm3, %%ymm12\n\t"
"vpaddd %%ymm3, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm2, %%ymm12, %%ymm2\n\t"
"vpsrlq $32, %%ymm5, %%ymm12\n\t"
"vpaddd %%ymm5, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm4, %%ymm12, %%ymm4\n\t"
"vpsrlq $32, %%ymm7, %%ymm12\n\t"
"vpaddd %%ymm7, %%ymm12, %%ymm12\n\t"
"vpsllq $32, %%ymm12, %%ymm12\n\t"
"vpaddq %%ymm6, %%ymm12, %%ymm6\n\t"
"movq %[c1], %%r8\n\t"
"movq %[st1], %%r9\n\t"
"cmpl $1, %[op1]\n\t"
"je .LF4%=_1b\n\t"
"jg .LF4%=_1c\n\t"
"vmovdqu %%ymm0, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm2, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm4, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu %%ymm6, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"jmp .LF4%=_1z\n\t"
".LF4%=_1b:\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpaddq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"jmp .LF4%=_1z\n\t"
".LF4%=_1c:\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
"vmovdqu (%%r8), %%ymm13\n\t"
"vpsubq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r8)\n\t"
"addq %%r9, %%r8\n\t"
".LF4%=_1z:\n\t"
"testq %[c2], %[c2]\n\t"
"jz .LF4%=_ez\n\t"
"movq %[c2], %%r10\n\t"
"movq %[st2], %%r11\n\t"
"cmpl $1, %[op2]\n\t"
"je .LF4%=_2b\n\t"
"jg .LF4%=_2c\n\t"
"vmovdqu %%ymm0, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm2, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm4, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu %%ymm6, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"jmp .LF4%=_ez\n\t"
".LF4%=_2b:\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpaddq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"jmp .LF4%=_ez\n\t"
".LF4%=_2c:\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm0, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm2, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm4, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
"vmovdqu (%%r10), %%ymm13\n\t"
"vpsubq %%ymm6, %%ymm13, %%ymm13\n\t"
"vmovdqu %%ymm13, (%%r10)\n\t"
"addq %%r11, %%r10\n\t"
".LF4%=_ez:\n\t"
""
: [ap] "+r"(apE), [bp] "+r"(bpE), [idx] "+r"(idx)
: [c1] "r"(c1), [st1] "r"(st1), [op1] "r"(op1),
[c2] "r"(c2), [st2] "r"(st2), [op2] "r"(op2)
: "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7",
"ymm8", "ymm9", "ymm10", "ymm11", "ymm12", "ymm13", "ymm14", "ymm15",
"r8", "r9", "r10", "r11", "cc", "memory");
} else
if (rows == 2) {
__asm__ volatile(
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vmovdqa %%ymm0, %%ymm1\n\t"
"vmovdqa %%ymm0, %%ymm2\n\t"
"vmovdqa %%ymm0, %%ymm3\n\t"
"test %[idx], %[idx]\n\t"
"jz 2f\n\t"
"vmovdqu (%[bp],%[idx],4), %%ymm12\n\t"
"1:\n\t"
"vpshufd $0xB1, %%ymm12, %%ymm13\n\t"
"vbroadcastsd 0(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm1, %%ymm1\n\t"
"vbroadcastsd 576(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm3, %%ymm3\n\t"
"vbroadcastsd 8(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm2, %%ymm2\n\t"
"vmovdqu 32(%[bp],%[idx],4), %%ymm12\n\t"
"add $8, %[idx]\n\t"
"jnz 1b\n\t"
"2:\n\t"
"vmovdqa %%ymm0, 0(%[ac])\n\t"
"vmovdqa %%ymm1, 32(%[ac])\n\t"
"vmovdqa %%ymm2, 64(%[ac])\n\t"
"vmovdqa %%ymm3, 96(%[ac])\n\t"
""
: [ap] "+r"(apE), [bp] "+r"(bpE), [idx] "+r"(idx)
: [ac] "r"(&AC[0])
: "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7",
"ymm8", "ymm9", "ymm10", "ymm11", "ymm12", "ymm13", "ymm14", "ymm15",
"cc", "memory");
} else
{
__asm__ volatile(
"vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
"vmovdqa %%ymm0, %%ymm1\n\t"
"vmovdqa %%ymm0, %%ymm2\n\t"
"vmovdqa %%ymm0, %%ymm3\n\t"
"vmovdqa %%ymm0, %%ymm4\n\t"
"vmovdqa %%ymm0, %%ymm5\n\t"
"vmovdqa %%ymm0, %%ymm6\n\t"
"vmovdqa %%ymm0, %%ymm7\n\t"
"vmovdqa %%ymm0, %%ymm8\n\t"
"vmovdqa %%ymm0, %%ymm9\n\t"
"vmovdqa %%ymm0, %%ymm10\n\t"
"vmovdqa %%ymm0, %%ymm11\n\t"
"test %[idx], %[idx]\n\t"
"jz 2f\n\t"
"vmovdqu (%[bp],%[idx],4), %%ymm12\n\t"
"1:\n\t"
"vpshufd $0xB1, %%ymm12, %%ymm13\n\t"
"vbroadcastsd 0(%[ap],%[idx],1), %%ymm14\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm1, %%ymm1\n\t"
"vbroadcastsd 576(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm0, %%ymm0\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm3, %%ymm3\n\t"
"vbroadcastsd 1152(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm2, %%ymm2\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm5, %%ymm5\n\t"
"vbroadcastsd 1728(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm4, %%ymm4\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm7, %%ymm7\n\t"
"vbroadcastsd 2304(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm6, %%ymm6\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm9, %%ymm9\n\t"
"vbroadcastsd 2880(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm8, %%ymm8\n\t"
"vpmuludq %%ymm14, %%ymm12, %%ymm15\n\t"
"vpmulld %%ymm13, %%ymm14, %%ymm14\n\t"
"vpaddd %%ymm14, %%ymm11, %%ymm11\n\t"
"vbroadcastsd 8(%[ap],%[idx],1), %%ymm14\n\t"
"vpaddq %%ymm15, %%ymm10, %%ymm10\n\t"
"vmovdqu 32(%[bp],%[idx],4), %%ymm12\n\t"
"add $8, %[idx]\n\t"
"jnz 1b\n\t"
"2:\n\t"
"vmovdqa %%ymm0, 0(%[ac])\n\t"
"vmovdqa %%ymm1, 32(%[ac])\n\t"
"vmovdqa %%ymm2, 64(%[ac])\n\t"
"vmovdqa %%ymm3, 96(%[ac])\n\t"
"vmovdqa %%ymm4, 128(%[ac])\n\t"
"vmovdqa %%ymm5, 160(%[ac])\n\t"
"vmovdqa %%ymm6, 192(%[ac])\n\t"
"vmovdqa %%ymm7, 224(%[ac])\n\t"
"vmovdqa %%ymm8, 256(%[ac])\n\t"
"vmovdqa %%ymm9, 288(%[ac])\n\t"
"vmovdqa %%ymm10, 320(%[ac])\n\t"
"vmovdqa %%ymm11, 352(%[ac])\n\t"
""
: [ap] "+r"(apE), [bp] "+r"(bpE), [idx] "+r"(idx)
: [ac] "r"(&AC[0])
: "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7",
"ymm8", "ymm9", "ymm10", "ymm11", "ymm12", "ymm13", "ymm14", "ymm15",
"cc", "memory");
}
if (rows == 6 || rows == 4) return;
if (rows > 0) {
__m256i lo = AC[0], hi = AC[1];
__m256i y = _mm256_add_epi32(hi, _mm256_srli_epi64(hi, 32));
__m256i res = _mm256_add_epi64(lo, _mm256_slli_epi64(y, 32));
long long *cp = c1 + (size_t)0 * ldc1;
if (op1 == 0) STQ((__m256i *)cp, res);
else if (op1 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res));
if (c2) { cp = c2 + (size_t)0 * ldc2;
if (op2 == 0) STQ((__m256i *)cp, res);
else if (op2 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res)); } }
if (rows > 1) {
__m256i lo = AC[2], hi = AC[3];
__m256i y = _mm256_add_epi32(hi, _mm256_srli_epi64(hi, 32));
__m256i res = _mm256_add_epi64(lo, _mm256_slli_epi64(y, 32));
long long *cp = c1 + (size_t)1 * ldc1;
if (op1 == 0) STQ((__m256i *)cp, res);
else if (op1 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res));
if (c2) { cp = c2 + (size_t)1 * ldc2;
if (op2 == 0) STQ((__m256i *)cp, res);
else if (op2 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res)); } }
if (rows > 2) {
__m256i lo = AC[4], hi = AC[5];
__m256i y = _mm256_add_epi32(hi, _mm256_srli_epi64(hi, 32));
__m256i res = _mm256_add_epi64(lo, _mm256_slli_epi64(y, 32));
long long *cp = c1 + (size_t)2 * ldc1;
if (op1 == 0) STQ((__m256i *)cp, res);
else if (op1 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res));
if (c2) { cp = c2 + (size_t)2 * ldc2;
if (op2 == 0) STQ((__m256i *)cp, res);
else if (op2 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res)); } }
if (rows > 3) {
__m256i lo = AC[6], hi = AC[7];
__m256i y = _mm256_add_epi32(hi, _mm256_srli_epi64(hi, 32));
__m256i res = _mm256_add_epi64(lo, _mm256_slli_epi64(y, 32));
long long *cp = c1 + (size_t)3 * ldc1;
if (op1 == 0) STQ((__m256i *)cp, res);
else if (op1 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res));
if (c2) { cp = c2 + (size_t)3 * ldc2;
if (op2 == 0) STQ((__m256i *)cp, res);
else if (op2 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res)); } }
if (rows > 4) {
__m256i lo = AC[8], hi = AC[9];
__m256i y = _mm256_add_epi32(hi, _mm256_srli_epi64(hi, 32));
__m256i res = _mm256_add_epi64(lo, _mm256_slli_epi64(y, 32));
long long *cp = c1 + (size_t)4 * ldc1;
if (op1 == 0) STQ((__m256i *)cp, res);
else if (op1 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res));
if (c2) { cp = c2 + (size_t)4 * ldc2;
if (op2 == 0) STQ((__m256i *)cp, res);
else if (op2 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res)); } }
if (rows > 5) {
__m256i lo = AC[10], hi = AC[11];
__m256i y = _mm256_add_epi32(hi, _mm256_srli_epi64(hi, 32));
__m256i res = _mm256_add_epi64(lo, _mm256_slli_epi64(y, 32));
long long *cp = c1 + (size_t)5 * ldc1;
if (op1 == 0) STQ((__m256i *)cp, res);
else if (op1 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res));
if (c2) { cp = c2 + (size_t)5 * ldc2;
if (op2 == 0) STQ((__m256i *)cp, res);
else if (op2 == 1) STQ((__m256i *)cp, _mm256_add_epi64(LDQ((const __m256i *)cp), res));
else STQ((__m256i *)cp, _mm256_sub_epi64(LDQ((const __m256i *)cp), res)); } }
}
struct Opd { /* value(i,j) = sum_t s[t]*p[(r[t]+i)*ld + c[t]+j] (nt <= 4)
z=1: the ld x ld block is stored QUADRANT-MAJOR -- four dense
(ld/2) x (ld/2) quadrants at (2*qi+qj)*(ld/2)^2. A z=1 descriptor
is only ever consumed by oq(), which returns a PLAIN AFFINE
descriptor of the dense quadrant. */
const long long *p; int ld; int z; int nt; int r[8], c[8], s[8];
};
static inline Opd oq(const Opd *o, int qi, int qj, int h) {
Opd d = *o;
if (o->z) { /* quadrant-major parent: the child IS one dense block */
const int q = 2 * qi + qj; /* child base = r*ld + q*h*h, ld = 2h => r' = 2r + q*h */
for (int t = 0; t < o->nt; t++) { d.r[t] = 2 * o->r[t] + q * h; d.c[t] = 0; }
d.ld = h; d.z = 0;
} else {
for (int t = 0; t < o->nt; t++) { d.r[t] += qi * h; d.c[t] += qj * h; }
}
return d;
}
static inline Opd oadd(const Opd *a, const Opd *b, int sgn) {
Opd d = *a;
for (int t = 0; t < b->nt; t++) { d.r[d.nt] = b->r[t]; d.c[d.nt] = b->c[t]; d.s[d.nt] = b->s[t] < 0 ? -sgn : sgn; d.nt++; }
return d;
}
static inline Opd o1(const long long *p, int ld) {
Opd d; d.p = p; d.ld = ld; d.z = 0; d.nt = 1; d.r[0] = 0; d.c[0] = 0; d.s[0] = 1; return d;
}
static inline Opd o1z(const long long *p, int n) { /* quadrant-major n x n block */
Opd d; d.p = p; d.ld = n; d.z = 1; d.nt = 1; d.r[0] = 0; d.c[0] = 0; d.s[0] = 1; return d;
}
/* A panel: rows [ic, ic+mc), k in [pc, pc+kc) of the operand */
static void packA(const Opd *o, int ic, int mc, int pc, int kc) {
const int nt = o->nt;
const long long *P = o->p; const int ld = o->ld;
for (int r = 0; r < MC; r++) {
long long *p = Apanel + (size_t)r * KCP;
int k = 0;
if (r < mc) {
const long long *s0 = P + (size_t)(o->r[0] + ic + r) * ld + o->c[0] + pc;
if (nt == 1) {
for (; k + 16 <= kc; k += 16) {
__m256i a0=LDQ((const __m256i *)(s0+k)), a1=LDQ((const __m256i *)(s0+k+4));
__m256i a2=LDQ((const __m256i *)(s0+k+8)), a3=LDQ((const __m256i *)(s0+k+12));
PFW((char *)p + 1024); PFW((char *)p + 1088);
PFW((char *)p + 512);
STQ((__m256i *)p, a0); STQ((__m256i *)(p+4), a1);
STQ((__m256i *)(p+8), a2); STQ((__m256i *)(p+12), a3);
p += 16;
}
for (; k + 4 <= kc; k += 4) { STQ((__m256i *)p, LDQ((const __m256i *)(s0 + k))); p += 4; }
for (; k < kc; k++) *p++ = s0[k];
} else if (nt == 2) {
const long long *s1 = P + (size_t)(o->r[1] + ic + r) * ld + o->c[1] + pc;
if (o->s[1] > 0) {
for (; k + 16 <= kc; k += 16) {
__m256i a0=LDQ((const __m256i *)(s0+k)), a1=LDQ((const __m256i *)(s0+k+4));
__m256i a2=LDQ((const __m256i *)(s0+k+8)), a3=LDQ((const __m256i *)(s0+k+12));
__m256i b0=LDQ((const __m256i *)(s1+k)), b1=LDQ((const __m256i *)(s1+k+4));
__m256i b2=LDQ((const __m256i *)(s1+k+8)), b3=LDQ((const __m256i *)(s1+k+12));
STQ((__m256i *)p, _mm256_add_epi64(a0,b0)); STQ((__m256i *)(p+4), _mm256_add_epi64(a1,b1));
STQ((__m256i *)(p+8), _mm256_add_epi64(a2,b2)); STQ((__m256i *)(p+12), _mm256_add_epi64(a3,b3));
p += 16;
}
for (; k + 4 <= kc; k += 4) {
__m256i a = LDQ((const __m256i *)(s0+k)), b = LDQ((const __m256i *)(s1+k));
STQ((__m256i *)p, _mm256_add_epi64(a,b)); p += 4;
}
for (; k < kc; k++) *p++ = s0[k] + s1[k];
} else {
for (; k + 16 <= kc; k += 16) {
__m256i a0=LDQ((const __m256i *)(s0+k)), a1=LDQ((const __m256i *)(s0+k+4));
__m256i a2=LDQ((const __m256i *)(s0+k+8)), a3=LDQ((const __m256i *)(s0+k+12));
__m256i b0=LDQ((const __m256i *)(s1+k)), b1=LDQ((const __m256i *)(s1+k+4));
__m256i b2=LDQ((const __m256i *)(s1+k+8)), b3=LDQ((const __m256i *)(s1+k+12));
STQ((__m256i *)p, _mm256_sub_epi64(a0,b0)); STQ((__m256i *)(p+4), _mm256_sub_epi64(a1,b1));
STQ((__m256i *)(p+8), _mm256_sub_epi64(a2,b2)); STQ((__m256i *)(p+12), _mm256_sub_epi64(a3,b3));
p += 16;
}
for (; k + 4 <= kc; k += 4) {
__m256i a = LDQ((const __m256i *)(s0+k)), b = LDQ((const __m256i *)(s1+k));
STQ((__m256i *)p, _mm256_sub_epi64(a,b)); p += 4;
}
for (; k < kc; k++) *p++ = s0[k] - s1[k];
}
} else {
for (; k + 4 <= kc; k += 4) {
__m256i v = _mm256_setzero_si256();
for (int t = 0; t < nt; t++) {
__m256i w = LDQ((const __m256i *)(P + (size_t)(o->r[t] + ic + r) * ld + o->c[t] + pc + k));
v = (o->s[t] < 0) ? _mm256_sub_epi64(v, w) : _mm256_add_epi64(v, w);
}
STQ((__m256i *)p, v);
p += 4;
}
for (; k < kc; k++) {
long long v = 0;
for (int t = 0; t < nt; t++)
v += o->s[t] < 0 ? -P[(size_t)(o->r[t] + ic + r) * ld + o->c[t] + pc + k]
: P[(size_t)(o->r[t] + ic + r) * ld + o->c[t] + pc + k];
*p++ = (v);
}
}
}
p += KC - k; /* k >= kc is never read: mikro runs kc steps */
}
}
static void packB(const Opd *o, int jc, int nc, int pc, int kc) {
const int nt = o->nt;
const long long *P = o->p; const int ld = o->ld;
const int jbn = (nc + 3) & ~3; /* mikro reads only jg < nc, step 4 */
int jb = 0;
if (nt <= 2) {
for (; jb + 8 <= nc; jb += 8) {
long long *q0 = Bpanel + (size_t)(jb / 4) * KC * 4;
long long *q1 = q0 + (size_t)KC * 4;
const long long *sq = P + (size_t)o->r[0] * ld + o->c[0] + jc + jb;
const long long *sq1 = (nt >= 2) ? P + (size_t)o->r[1] * ld + o->c[1] + jc + jb : 0;
const long long *ap0 = sq + (size_t)pc * ld;
const long long *ap1 = (nt >= 2) ? sq1 + (size_t)pc * ld : 0;
int k = 0;
if (nt == 1) {
for (; k + 2 <= kc; k += 2, ap0 += 2 * ld) {
if ((k & 7) == 0) {
char *d0 = (char *)(q0 + (size_t)k * 4) + 2048, *d1 = (char *)(q1 + (size_t)k * 4) + 2048;
for (int tt = 0; tt < 4; tt++) { PFW(d0 + 64 * tt); PFW(d1 + 64 * tt); }
}
const long long *a0 = ap0, *a1 = a0 + ld;
__m256i v0 = LDQ((const __m256i *)a0), u0 = LDQ((const __m256i *)(a0 + 4));
__m256i v1 = LDQ((const __m256i *)a1), u1 = LDQ((const __m256i *)(a1 + 4));
_mm256_store_si256((__m256i *)(q0 + (size_t)k * 4), v0);
_mm256_store_si256((__m256i *)(q1 + (size_t)k * 4), u0);
_mm256_store_si256((__m256i *)(q0 + (size_t)(k + 1) * 4), v1);
_mm256_store_si256((__m256i *)(q1 + (size_t)(k + 1) * 4), u1);
}
for (; k < kc; k++) {
const long long *a = sq + (size_t)(pc + k) * ld;
_mm256_store_si256((__m256i *)(q0 + (size_t)k * 4), LDQ((const __m256i *)a));
_mm256_store_si256((__m256i *)(q1 + (size_t)k * 4), LDQ((const __m256i *)(a + 4)));
}
} else if (o->s[1] > 0) {
for (; k + 2 <= kc; k += 2, ap0 += 2 * ld, ap1 += 2 * ld) {
if ((k & 7) == 0) {
char *d0 = (char *)(q0 + (size_t)k * 4) + 2048, *d1 = (char *)(q1 + (size_t)k * 4) + 2048;
for (int tt = 0; tt < 4; tt++) { PFW(d0 + 64 * tt); PFW(d1 + 64 * tt); }
}
const long long *a0 = ap0, *a1 = a0 + ld;
const long long *b0 = ap1, *b1 = b0 + ld;
__m256i v0 = _mm256_add_epi64(LDQ((const __m256i *)a0), LDQ((const __m256i *)b0));
__m256i u0 = _mm256_add_epi64(LDQ((const __m256i *)(a0 + 4)), LDQ((const __m256i *)(b0 + 4)));
__m256i v1 = _mm256_add_epi64(LDQ((const __m256i *)a1), LDQ((const __m256i *)b1));
__m256i u1 = _mm256_add_epi64(LDQ((const __m256i *)(a1 + 4)), LDQ((const __m256i *)(b1 + 4)));
_mm256_store_si256((__m256i *)(q0 + (size_t)k * 4), v0);
_mm256_store_si256((__m256i *)(q1 + (size_t)k * 4), u0);
_mm256_store_si256((__m256i *)(q0 + (size_t)(k + 1) * 4), v1);
_mm256_store_si256((__m256i *)(q1 + (size_t)(k + 1) * 4), u1);
}
for (; k < kc; k++) {
const long long *a = sq + (size_t)(pc + k) * ld;
const long long *b = sq1 + (size_t)(pc + k) * ld;
_mm256_store_si256((__m256i *)(q0 + (size_t)k * 4), _mm256_add_epi64(LDQ((const __m256i *)a), LDQ((const __m256i *)b)));
_mm256_store_si256((__m256i *)(q1 + (size_t)k * 4), _mm256_add_epi64(LDQ((const __m256i *)(a + 4)), LDQ((const __m256i *)(b + 4))));
}
} else {
for (; k + 2 <= kc; k += 2, ap0 += 2 * ld, ap1 += 2 * ld) {
if ((k & 7) == 0) {
char *d0 = (char *)(q0 + (size_t)k * 4) + 2048, *d1 = (char *)(q1 + (size_t)k * 4) + 2048;
for (int tt = 0; tt < 4; tt++) { PFW(d0 + 64 * tt); PFW(d1 + 64 * tt); }
}
const long long *a0 = ap0, *a1 = a0 + ld;
const long long *b0 = ap1, *b1 = b0 + ld;
__m256i v0 = _mm256_sub_epi64(LDQ((const __m256i *)a0), LDQ((const __m256i *)b0));
__m256i u0 = _mm256_sub_epi64(LDQ((const __m256i *)(a0 + 4)), LDQ((const __m256i *)(b0 + 4)));
__m256i v1 = _mm256_sub_epi64(LDQ((const __m256i *)a1), LDQ((const __m256i *)b1));
__m256i u1 = _mm256_sub_epi64(LDQ((const __m256i *)(a1 + 4)), LDQ((const __m256i *)(b1 + 4)));
_mm256_store_si256((__m256i *)(q0 + (size_t)k * 4), v0);
_mm256_store_si256((__m256i *)(q1 + (size_t)k * 4), u0);
_mm256_store_si256((__m256i *)(q0 + (size_t)(k + 1) * 4), v1);
_mm256_store_si256((__m256i *)(q1 + (size_t)(k + 1) * 4), u1);
}
for (; k < kc; k++) {
const long long *a = sq + (size_t)(pc + k) * ld;
const long long *b = sq1 + (size_t)(pc + k) * ld;
_mm256_store_si256((__m256i *)(q0 + (size_t)k * 4), _mm256_sub_epi64(LDQ((const __m256i *)a), LDQ((const __m256i *)b)));
_mm256_store_si256((__m256i *)(q1 + (size_t)k * 4), _mm256_sub_epi64(LDQ((const __m256i *)(a + 4)), LDQ((const __m256i *)(b + 4))));
}
}
}
}
for (; jb < jbn; jb += 4) {
long long *q = Bpanel + (size_t)(jb / 4) * KC * 4;
int full = (jb + 4 <= nc);
int lim = nc - jb; if (lim > 4) lim = 4; if (lim < 0) lim = 0;
const long long *sq = P + (size_t)o->r[0] * ld + o->c[0] + jc + jb;
const long long *sq1 = (nt >= 2) ? P + (size_t)o->r[1] * ld + o->c[1] + jc + jb : 0;
if (full && nt <= 2) {
for (int k = 0; k < kc; k++) {
__m256i v = LDQ((const __m256i *)(sq + (size_t)(pc + k) * ld));
if (nt == 2) {
__m256i w = LDQ((const __m256i *)(sq1 + (size_t)(pc + k) * ld));
v = (o->s[1] > 0) ? _mm256_add_epi64(v, w) : _mm256_sub_epi64(v, w);
}
_mm256_store_si256((__m256i *)(q + (size_t)k * 4), v);
}
} else if (full) {
const long long *bt[8];
for (int tt = 0; tt < nt; tt++) bt[tt] = P + (size_t)(o->r[tt] + pc) * ld + o->c[tt] + jc + jb;
for (int k = 0; k < kc; k++) {
__m256i v = _mm256_setzero_si256();
for (int tt = 0; tt < nt; tt++) {
__m256i w = LDQ((const __m256i *)(bt[tt] + (size_t)k * ld));
v = (o->s[tt] < 0) ? _mm256_sub_epi64(v, w) : _mm256_add_epi64(v, w);
}
_mm256_store_si256((__m256i *)(q + (size_t)k * 4), v);
}
} else {
for (int k = 0; k < KC; k++) {
long long t[4] = {0, 0, 0, 0};
if (k < kc) {
int ncol = full ? 4 : lim;
for (int c = 0; c < ncol; c++) {
long long v = 0;
for (int tt = 0; tt < nt; tt++)
v += o->s[tt] < 0 ? -P[(size_t)(o->r[tt] + pc + k) * ld + o->c[tt] + jc + jb + c]
: P[(size_t)(o->r[tt] + pc + k) * ld + o->c[tt] + jc + jb + c];
t[c] = v;
}
}
__m256i v = LDQ((const __m256i *)t);
_mm256_store_si256((__m256i *)(q + (size_t)k * 4), v);
}
}
}
}
static inline int effop(int base, int term) { /* term: 1 = +M, 2 = -M ; result: 0 store 1 add 2 sub */
if (base == 0) return term == 1 ? 0 : 2;
if (base == 1) return term == 2 ? 2 : 1;
return term == 2 ? 1 : 2;
}
/* C1 op1 (+- second destination C2 op2, either may be absent) = Ao * Bo, block size n */
static void leaf(int n, const Opd *Ao, const Opd *Bo,
long long *C1, int ldc1, int op1, long long *C2, int ldc2, int op2) {
int nb4 = n - (n & 3);
for (int jc = 0; jc < nb4; jc += NC) {
int nc = (nb4 - jc < NC) ? (nb4 - jc) : NC;
for (int pc = 0; pc < n; pc += KC) {
int kc = (n - pc < KC) ? (n - pc) : KC;
packB(Bo, jc, nc, pc, kc);
int e1 = (pc == 0) ? op1 : (op1 == 2 ? 2 : 1);
int e2 = (pc == 0) ? op2 : (op2 == 2 ? 2 : 1);
for (int ic = 0; ic < n; ic += MC) {
int mc = (n - ic < MC) ? (n - ic) : MC;
packA(Ao, ic, mc, pc, kc);
for (int jgb = 0; jgb < nc; jgb += 128) {
int jge = (nc - jgb < 128) ? nc : jgb + 128;
for (int ir = 0; ir < mc; ir += MR) {
int rows = (mc - ir < MR) ? (mc - ir) : MR;
const long long *ap = Apanel + (size_t)ir * KCP * 1;
for (int jg = jgb; jg < jge; jg += 4) {
const long long *bp = Bpanel + (size_t)(jg / 4) * KC * 4;
mikro(kc / UU, ap, bp, C1 + (size_t)(ic + ir) * ldc1 + jc + jg, ldc1, e1,
C2 ? C2 + (size_t)(ic + ir) * ldc2 + jc + jg : 0, ldc2, e2, rows);
}
}
}
}
}
}
for (int j = nb4; j < n; j++)
for (int i = 0; i < n; i++) {
unsigned long long acc = 0;
for (int k = 0; k < n; k++) {
long long av = 0, bv = 0;
for (int t = 0; t < Ao->nt; t++)
av += Ao->s[t] < 0 ? -Ao->p[(size_t)(Ao->r[t] + i) * Ao->ld + Ao->c[t] + k]
: Ao->p[(size_t)(Ao->r[t] + i) * Ao->ld + Ao->c[t] + k];
for (int t = 0; t < Bo->nt; t++)
bv += Bo->s[t] < 0 ? -Bo->p[(size_t)(Bo->r[t] + k) * Bo->ld + Bo->c[t] + j]
: Bo->p[(size_t)(Bo->r[t] + k) * Bo->ld + Bo->c[t] + j];
acc += (unsigned long long)av * (unsigned long long)bv;
}
long long *cp = C1 + (size_t)i * ldc1 + j;
if (op1 == 0) *cp = (long long)acc;
else if (op1 == 1) *cp += (long long)acc;
else *cp -= (long long)acc;
if (C2) {
long long *cq = C2 + (size_t)i * ldc2 + j;
if (op2 == 0) *cq = (long long)acc;
else if (op2 == 1) *cq += (long long)acc;
else *cq -= (long long)acc;
}
}
}
/* recursive Strassen: C1 op1 op= Ao*Bo and C2 op2 op= Ao*Bo */
/* Scratch: one h*h product block per recursion level, carved from the tail of the
buffer. Only the levels that actually recurse consume space, so a walk of the
recursion computes exactly the bytes matrix_multiply() hands us below. */
static long long g_scratch[6u * 1024u * 1024u] __attribute__((aligned(4096)));
/* dst op= sign * T (op: 0 store, 1 add, 2 sub) */
static void cmb(long long *dst, int ldd, const long long *T, int ldt, int h, int sign, int op) {
for (int i = 0; i < h; i++) {
const long long *t = T + (size_t)i * ldt;
long long *c = dst + (size_t)i * ldd;
int j = 0;
if (op == 0) {
for (; j + 16 <= h; j += 16) {
__m256i v0=LDQ((const __m256i*)(t+j)), v1=LDQ((const __m256i*)(t+j+4));
__m256i v2=LDQ((const __m256i*)(t+j+8)), v3=LDQ((const __m256i*)(t+j+12));
if (sign < 0) { v0=_mm256_sub_epi64(_mm256_setzero_si256(),v0);
v1=_mm256_sub_epi64(_mm256_setzero_si256(),v1);
v2=_mm256_sub_epi64(_mm256_setzero_si256(),v2);
v3=_mm256_sub_epi64(_mm256_setzero_si256(),v3); }
STQ((__m256i*)(c+j),v0); STQ((__m256i*)(c+j+4),v1);
STQ((__m256i*)(c+j+8),v2); STQ((__m256i*)(c+j+12),v3);
}
for (; j + 4 <= h; j += 4) {
__m256i v = LDQ((const __m256i *)(t + j));
if (sign < 0) v = _mm256_sub_epi64(_mm256_setzero_si256(), v);
STQ((__m256i *)(c + j), v);
}
for (; j < h; j++) c[j] = sign < 0 ? -t[j] : t[j];
} else {
for (; j + 16 <= h; j += 16) {
__m256i v0=LDQ((const __m256i*)(t+j)), v1=LDQ((const __m256i*)(t+j+4));
__m256i v2=LDQ((const __m256i*)(t+j+8)), v3=LDQ((const __m256i*)(t+j+12));
__m256i w0=LDQ((const __m256i*)(c+j)), w1=LDQ((const __m256i*)(c+j+4));
__m256i w2=LDQ((const __m256i*)(c+j+8)), w3=LDQ((const __m256i*)(c+j+12));
if ((op == 1) != (sign < 0)) { STQ((__m256i*)(c+j),_mm256_add_epi64(w0,v0));
STQ((__m256i*)(c+j+4),_mm256_add_epi64(w1,v1));
STQ((__m256i*)(c+j+8),_mm256_add_epi64(w2,v2));
STQ((__m256i*)(c+j+12),_mm256_add_epi64(w3,v3)); }
else { STQ((__m256i*)(c+j),_mm256_sub_epi64(w0,v0));
STQ((__m256i*)(c+j+4),_mm256_sub_epi64(w1,v1));
STQ((__m256i*)(c+j+8),_mm256_sub_epi64(w2,v2));
STQ((__m256i*)(c+j+12),_mm256_sub_epi64(w3,v3)); }
}
for (; j + 4 <= h; j += 4) {
__m256i v = LDQ((const __m256i *)(t + j));
__m256i w = LDQ((const __m256i *)(c + j));
if ((op == 1) != (sign < 0)) STQ((__m256i *)(c + j), _mm256_add_epi64(w, v));
else STQ((__m256i *)(c + j), _mm256_sub_epi64(w, v));
}
for (; j < h; j++) { if ((op == 1) != (sign < 0)) c[j] += t[j]; else c[j] -= t[j]; }
}
}
}
/* fused two-destination combine: reads T ONCE for both destinations.
Semantics identical to cmb(d1,s1,o1) then cmb(d2,s2,o2) on the same T. */
static void cmb2(long long *d1,int ld1,int op1,int s1,
long long *d2,int ld2,int op2,int s2,
const long long *T,int ldt,int h) {
for (int i = 0; i < h; i++) {
const long long *t = T + (size_t)i * ldt;
long long *c1 = d1 + (size_t)i * ld1, *c2 = d2 + (size_t)i * ld2;
int j = 0;
for (; j + 16 <= h; j += 16) {
__m256i v0=LDQ((const __m256i*)(t+j)), v1=LDQ((const __m256i*)(t+j+4));
__m256i v2=LDQ((const __m256i*)(t+j+8)), v3=LDQ((const __m256i*)(t+j+12));
__m256i a0=LDQ((const __m256i*)(c1+j)), a1=LDQ((const __m256i*)(c1+j+4));
__m256i a2=LDQ((const __m256i*)(c1+j+8)), a3=LDQ((const __m256i*)(c1+j+12));
__m256i b0=LDQ((const __m256i*)(c2+j)), b1=LDQ((const __m256i*)(c2+j+4));
__m256i b2=LDQ((const __m256i*)(c2+j+8)), b3=LDQ((const __m256i*)(c2+j+12));
if ((op1 == 1) != (s1 < 0)) { STQ((__m256i*)(c1+j),_mm256_add_epi64(a0,v0));
STQ((__m256i*)(c1+j+4),_mm256_add_epi64(a1,v1)); STQ((__m256i*)(c1+j+8),_mm256_add_epi64(a2,v2));
STQ((__m256i*)(c1+j+12),_mm256_add_epi64(a3,v3)); }
else { STQ((__m256i*)(c1+j),_mm256_sub_epi64(a0,v0));
STQ((__m256i*)(c1+j+4),_mm256_sub_epi64(a1,v1)); STQ((__m256i*)(c1+j+8),_mm256_sub_epi64(a2,v2));
STQ((__m256i*)(c1+j+12),_mm256_sub_epi64(a3,v3)); }
if ((op2 == 1) != (s2 < 0)) { STQ((__m256i*)(c2+j),_mm256_add_epi64(b0,v0));
STQ((__m256i*)(c2+j+4),_mm256_add_epi64(b1,v1)); STQ((__m256i*)(c2+j+8),_mm256_add_epi64(b2,v2));
STQ((__m256i*)(c2+j+12),_mm256_add_epi64(b3,v3)); }
else { STQ((__m256i*)(c2+j),_mm256_sub_epi64(b0,v0));
STQ((__m256i*)(c2+j+4),_mm256_sub_epi64(b1,v1)); STQ((__m256i*)(c2+j+8),_mm256_sub_epi64(b2,v2));
STQ((__m256i*)(c2+j+12),_mm256_sub_epi64(b3,v3)); }
}
for (; j + 4 <= h; j += 4) {
__m256i v = LDQ((const __m256i *)(t + j));
__m256i a = LDQ((const __m256i *)(c1 + j));
STQ((__m256i *)(c1 + j), ((op1 == 1) != (s1 < 0)) ? _mm256_add_epi64(a,v)
: _mm256_sub_epi64(a,v));
__m256i b = LDQ((const __m256i *)(c2 + j));
STQ((__m256i *)(c2 + j), ((op2 == 1) != (s2 < 0)) ? _mm256_add_epi64(b,v)
: _mm256_sub_epi64(b,v));
}
for (; j < h; j++) {
long long v = t[j];
if ((op1 == 1) != (s1 < 0)) c1[j] += v; else c1[j] -= v;
if ((op2 == 1) != (s2 < 0)) c2[j] += v; else c2[j] -= v;
}
}
}
/* single-destination recursive Strassen: d op= Ao*Bo, with T a scratch block of h*h */
static long long matbufA[3][(size_t)64*CUTL*CUTL + 64] __attribute__((aligned(4096))), matbufB[3][(size_t)64*CUTL*CUTL + 64] __attribute__((aligned(4096)));
/* ---- M1: matop with the nt/sign dispatch HOISTED out of the j-loop and the
per-row zi multiplies replaced by pointer walks. Semantics identical. ---- */
static inline void mat1(long long *w0, long long *dp, int q, int s0) {
int j = 0;
if (s0 > 0) {
for (; j + 16 <= q; j += 16) {
__m256i v0=LDQ((const __m256i*)(w0+j)), v1=LDQ((const __m256i*)(w0+j+4));
__m256i v2=LDQ((const __m256i*)(w0+j+8)), v3=LDQ((const __m256i*)(w0+j+12));
PFW((char *)(dp+j) + 2048); PFW((char *)(dp+j) + 2112);
STQ((__m256i*)(dp+j), v0); STQ((__m256i*)(dp+j+4), v1);
STQ((__m256i*)(dp+j+8), v2); STQ((__m256i*)(dp+j+12), v3);
}
for (; j + 4 <= q; j += 4) STQ((__m256i*)(dp+j), LDQ((const __m256i*)(w0+j)));
for (; j < q; j++) dp[j] = w0[j];
} else {
__m256i z = _mm256_setzero_si256();
for (; j + 16 <= q; j += 16) {
__m256i v0=LDQ((const __m256i*)(w0+j)), v1=LDQ((const __m256i*)(w0+j+4));
__m256i v2=LDQ((const __m256i*)(w0+j+8)), v3=LDQ((const __m256i*)(w0+j+12));
STQ((__m256i*)(dp+j), _mm256_sub_epi64(z,v0)); STQ((__m256i*)(dp+j+4), _mm256_sub_epi64(z,v1));
STQ((__m256i*)(dp+j+8), _mm256_sub_epi64(z,v2)); STQ((__m256i*)(dp+j+12), _mm256_sub_epi64(z,v3));
}
for (; j + 4 <= q; j += 4) { __m256i v=LDQ((const __m256i*)(w0+j)); STQ((__m256i*)(dp+j), _mm256_sub_epi64(z, v)); }
for (; j < q; j++) dp[j] = -w0[j];
}
}
static inline void mat2(long long *w0, long long *w1, long long *dp, int q, int s0, int s1) {
int j = 0;
__m256i z = _mm256_setzero_si256();
switch ((s0 > 0 ? 2 : 0) | (s1 > 0 ? 1 : 0)) {
case 3:
for (; j + 16 <= q; j += 16) {
__m256i v0=LDQ((const __m256i*)(w0+j)), v1=LDQ((const __m256i*)(w0+j+4));
__m256i v2=LDQ((const __m256i*)(w0+j+8)), v3=LDQ((const __m256i*)(w0+j+12));
__m256i u0=LDQ((const __m256i*)(w1+j)), u1=LDQ((const __m256i*)(w1+j+4));
__m256i u2=LDQ((const __m256i*)(w1+j+8)), u3=LDQ((const __m256i*)(w1+j+12));
PFW((char *)(dp+j) + 2048); PFW((char *)(dp+j) + 2112);
STQ((__m256i*)(dp+j), _mm256_add_epi64(v0,u0)); STQ((__m256i*)(dp+j+4), _mm256_add_epi64(v1,u1));
STQ((__m256i*)(dp+j+8), _mm256_add_epi64(v2,u2)); STQ((__m256i*)(dp+j+12), _mm256_add_epi64(v3,u3));
}
for (; j + 4 <= q; j += 4) { __m256i v=LDQ((const __m256i*)(w0+j)), u=LDQ((const __m256i*)(w1+j));
STQ((__m256i*)(dp+j), _mm256_add_epi64(v,u)); }
for (; j < q; j++) dp[j] = w0[j] + w1[j];
return;
case 2:
for (; j + 16 <= q; j += 16) {
__m256i v0=LDQ((const __m256i*)(w0+j)), v1=LDQ((const __m256i*)(w0+j+4));
__m256i v2=LDQ((const __m256i*)(w0+j+8)), v3=LDQ((const __m256i*)(w0+j+12));
__m256i u0=LDQ((const __m256i*)(w1+j)), u1=LDQ((const __m256i*)(w1+j+4));
__m256i u2=LDQ((const __m256i*)(w1+j+8)), u3=LDQ((const __m256i*)(w1+j+12));
PFW((char *)(dp+j) + 2048); PFW((char *)(dp+j) + 2112);
STQ((__m256i*)(dp+j), _mm256_sub_epi64(v0,u0)); STQ((__m256i*)(dp+j+4), _mm256_sub_epi64(v1,u1));
STQ((__m256i*)(dp+j+8), _mm256_sub_epi64(v2,u2)); STQ((__m256i*)(dp+j+12), _mm256_sub_epi64(v3,u3));
}
for (; j + 4 <= q; j += 4) { __m256i v=LDQ((const __m256i*)(w0+j)), u=LDQ((const __m256i*)(w1+j));
STQ((__m256i*)(dp+j), _mm256_sub_epi64(v,u)); }
for (; j < q; j++) dp[j] = w0[j] - w1[j];
return;
case 1:
for (; j + 16 <= q; j += 16) {
__m256i v0=LDQ((const __m256i*)(w0+j)), v1=LDQ((const __m256i*)(w0+j+4));
__m256i v2=LDQ((const __m256i*)(w0+j+8)), v3=LDQ((const __m256i*)(w0+j+12));
__m256i u0=LDQ((const __m256i*)(w1+j)), u1=LDQ((const __m256i*)(w1+j+4));
__m256i u2=LDQ((const __m256i*)(w1+j+8)), u3=LDQ((const __m256i*)(w1+j+12));
PFW((char *)(dp+j) + 2048); PFW((char *)(dp+j) + 2112);
STQ((__m256i*)(dp+j), _mm256_sub_epi64(u0,v0)); STQ((__m256i*)(dp+j+4), _mm256_sub_epi64(u1,v1));
STQ((__m256i*)(dp+j+8), _mm256_sub_epi64(u2,v2)); STQ((__m256i*)(dp+j+12), _mm256_sub_epi64(u3,v3));
}
for (; j + 4 <= q; j += 4) { __m256i v=LDQ((const __m256i*)(w0+j)), u=LDQ((const __m256i*)(w1+j));
STQ((__m256i*)(dp+j), _mm256_sub_epi64(u,v)); }
for (; j < q; j++) dp[j] = w1[j] - w0[j];
return;
default:
for (; j + 16 <= q; j += 16) {
__m256i v0=LDQ((const __m256i*)(w0+j)), v1=LDQ((const __m256i*)(w0+j+4));
__m256i v2=LDQ((const __m256i*)(w0+j+8)), v3=LDQ((const __m256i*)(w0+j+12));
__m256i u0=LDQ((const __m256i*)(w1+j)), u1=LDQ((const __m256i*)(w1+j+4));
__m256i u2=LDQ((const __m256i*)(w1+j+8)), u3=LDQ((const __m256i*)(w1+j+12));
STQ((__m256i*)(dp+j), _mm256_sub_epi64(z,_mm256_add_epi64(v0,u0)));
STQ((__m256i*)(dp+j+4), _mm256_sub_epi64(z,_mm256_add_epi64(v1,u1)));
STQ((__m256i*)(dp+j+8), _mm256_sub_epi64(z,_mm256_add_epi64(v2,u2)));
STQ((__m256i*)(dp+j+12), _mm256_sub_epi64(z,_mm256_add_epi64(v3,u3)));
}
for (; j + 4 <= q; j += 4) { __m256i v=LDQ((const __m256i*)(w0+j)), u=LDQ((const __m256i*)(w1+j));
STQ((__m256i*)(dp+j), _mm256_sub_epi64(z,_mm256_add_epi64(v,u))); }
for (; j < q; j++) dp[j] = -w0[j] - w1[j];
return;
}
}
static void matop(const Opd *o, int n, long long *dst) {
const int nt = o->nt; const long long *P = o->p; const int ld = o->ld;
int r0=0,r1=0,r2=0,r3=0, s0=1,s1=1,s2=1,s3=1;
long long *b0=(long long *)P,*b1=b0,*b2=b0,*b3=b0;
if (nt > 0) { r0=o->r[0]; b0=(long long *)P + o->c[0]; s0=o->s[0]; }
if (nt > 1) { r1=o->r[1]; b1=(long long *)P + o->c[1]; s1=o->s[1]; }
if (nt > 2) { r2=o->r[2]; b2=(long long *)P + o->c[2]; s2=o->s[2]; }
if (nt > 3) { r3=o->r[3]; b3=(long long *)P + o->c[3]; s3=o->s[3]; }
const int qh = n >> 1;
long long *w0 = b0 + (size_t)r0 * ld, *w1 = b1 + (size_t)r1 * ld;
long long *w2 = b2 + (size_t)r2 * ld, *w3 = b3 + (size_t)r3 * ld;
const size_t qhh = (size_t)qh * qh;
long long *dpA = dst, *dpB = dst + qhh;
if (nt == 1) {
for (int i = 0; i < n; i++, w0 += ld, dpA += qh, dpB += qh) {
if (i == qh) { dpA = dst + 2*qhh; dpB = dpA + qhh; }
mat1(w0, dpA, qh, s0);
mat1(w0 + qh, dpB, qh, s0);
}
return;
}
if (nt == 2) {
for (int i = 0; i < n; i++, w0 += ld, w1 += ld, dpA += qh, dpB += qh) {
if (i == qh) { dpA = dst + 2*qhh; dpB = dpA + qhh; }
mat2(w0, w1, dpA, qh, s0, s1);
mat2(w0 + qh, w1 + qh, dpB, qh, s0, s1);
}
return;
}
for (int i = 0; i < n; i++, w0 += ld, w1 += ld, w2 += ld, w3 += ld, dpA += qh, dpB += qh) {
if (i == qh) { dpA = dst + 2*qhh; dpB = dpA + qhh; }
for (int qq = 0; qq < 2; qq++) {
long long *dp = qq ? dpB : dpA;
const int jb = qq * qh;
int j = 0;
for (; j + 16 <= qh; j += 16) {
__m256i v0=_mm256_setzero_si256(),v1=v0,v2=v0,v3=v0;
{ const long long *sp = w0 + jb + j;
__m256i a0=LDQ((const __m256i *)sp), a1=LDQ((const __m256i *)(sp+4));
__m256i a2=LDQ((const __m256i *)(sp+8)), a3=LDQ((const __m256i *)(sp+12));
if (s0 < 0) { v0=_mm256_sub_epi64(v0,a0); v1=_mm256_sub_epi64(v1,a1); v2=_mm256_sub_epi64(v2,a2); v3=_mm256_sub_epi64(v3,a3); }
else { v0=_mm256_add_epi64(v0,a0); v1=_mm256_add_epi64(v1,a1); v2=_mm256_add_epi64(v2,a2); v3=_mm256_add_epi64(v3,a3); } }
{ const long long *sp = w1 + jb + j;
__m256i a0=LDQ((const __m256i *)sp), a1=LDQ((const __m256i *)(sp+4));
__m256i a2=LDQ((const __m256i *)(sp+8)), a3=LDQ((const __m256i *)(sp+12));
if (s1 < 0) { v0=_mm256_sub_epi64(v0,a0); v1=_mm256_sub_epi64(v1,a1); v2=_mm256_sub_epi64(v2,a2); v3=_mm256_sub_epi64(v3,a3); }
else { v0=_mm256_add_epi64(v0,a0); v1=_mm256_add_epi64(v1,a1); v2=_mm256_add_epi64(v2,a2); v3=_mm256_add_epi64(v3,a3); } }
if (nt > 2) { const long long *sp = w2 + jb + j;
__m256i a0=LDQ((const __m256i *)sp), a1=LDQ((const __m256i *)(sp+4));
__m256i a2=LDQ((const __m256i *)(sp+8)), a3=LDQ((const __m256i *)(sp+12));
if (s2 < 0) { v0=_mm256_sub_epi64(v0,a0); v1=_mm256_sub_epi64(v1,a1); v2=_mm256_sub_epi64(v2,a2); v3=_mm256_sub_epi64(v3,a3); }
else { v0=_mm256_add_epi64(v0,a0); v1=_mm256_add_epi64(v1,a1); v2=_mm256_add_epi64(v2,a2); v3=_mm256_add_epi64(v3,a3); } }
if (nt > 3) { const long long *sp = w3 + jb + j;
__m256i a0=LDQ((const __m256i *)sp), a1=LDQ((const __m256i *)(sp+4));
__m256i a2=LDQ((const __m256i *)(sp+8)), a3=LDQ((const __m256i *)(sp+12));
if (s3 < 0) { v0=_mm256_sub_epi64(v0,a0); v1=_mm256_sub_epi64(v1,a1); v2=_mm256_sub_epi64(v2,a2); v3=_mm256_sub_epi64(v3,a3); }
else { v0=_mm256_add_epi64(v0,a0); v1=_mm256_add_epi64(v1,a1); v2=_mm256_add_epi64(v2,a2); v3=_mm256_add_epi64(v3,a3); } }
STQ((__m256i *)(dp+j), v0); STQ((__m256i *)(dp+j+4), v1);
STQ((__m256i *)(dp+j+8), v2); STQ((__m256i *)(dp+j+12), v3);
}
for (; j + 4 <= qh; j += 4) {
__m256i v = _mm256_setzero_si256();
{ __m256i w = LDQ((const __m256i *)(w0 + jb + j));
v = (s0 < 0) ? _mm256_sub_epi64(v, w) : _mm256_add_epi64(v, w); }
{ __m256i w = LDQ((const __m256i *)(w1 + jb + j));
v = (s1 < 0) ? _mm256_sub_epi64(v, w) : _mm256_add_epi64(v, w); }
if (nt > 2) { __m256i w = LDQ((const __m256i *)(w2 + jb + j));
v = (s2 < 0) ? _mm256_sub_epi64(v, w) : _mm256_add_epi64(v, w); }
if (nt > 3) { __m256i w = LDQ((const __m256i *)(w3 + jb + j));
v = (s3 < 0) ? _mm256_sub_epi64(v, w) : _mm256_add_epi64(v, w); }
STQ((__m256i *)(dp + j), v);
}
for (; j < qh; j++) {
long long v = (s0<0?-w0[jb+j]:w0[jb+j]) + (s1<0?-w1[jb+j]:w1[jb+j]);
if (nt > 2) v += (s2<0?-w2[jb+j]:w2[jb+j]);
if (nt > 3) v += (s3<0?-w3[jb+j]:w3[jb+j]);
dp[j] = v;
}
}
}
}
static void mmxR(int n, const Opd *Ao0, const Opd *Bo0, long long *d, int ld, int op,
long long *T) {
const Opd *Ao = Ao0, *Bo = Bo0; Opd A1, B1;
if (n <= CUTL || (n & 1) || Ao->nt > 4 || Bo->nt > 4) {
leaf(n, Ao, Bo, d, ld, op, 0, 0, 0); return; }
if (n <= 8 * CUTL && (Ao->nt > 1 || Bo->nt > 1)) {
int Lm = (n > 4 * CUTL) ? 0 : (n > 2 * CUTL) ? 1 : 2;
if (Ao->nt > 1) { matop(Ao, n, matbufA[Lm]); A1 = o1z(matbufA[Lm], n); Ao = &A1; }
if (Bo->nt > 1) { matop(Bo, n, matbufB[Lm]); B1 = o1z(matbufB[Lm], n); Bo = &B1; }
}
if (n <= CUTL || (n & 1) || Ao->nt > 4 || Bo->nt > 4) {
leaf(n, Ao, Bo, d, ld, op, 0, 0, 0); return; }
const int h = n >> 1;
Opd A11 = oq(Ao, 0, 0, h), A12 = oq(Ao, 0, 1, h), A21 = oq(Ao, 1, 0, h), A22 = oq(Ao, 1, 1, h);
Opd B11 = oq(Bo, 0, 0, h), B12 = oq(Bo, 0, 1, h), B21 = oq(Bo, 1, 0, h), B22 = oq(Bo, 1, 1, h);
long long *Q[4];
Q[0] = d; Q[1] = d + h; Q[2] = d + (size_t)h * ld; Q[3] = Q[2] + h;
long long *T2 = T + (size_t)h * h;
int wr[4] = {0, 0, 0, 0};
const int LF = (h <= CUTL); /* children ARE leaves: one pass into both quadrants */
#define MK(qi, si, qj, sj, Ma, Mb) do { \
int oa = wr[qi] ? effop(1, (si)) : effop(op, (si)); \
int ob = wr[qj] ? effop(1, (sj)) : effop(op, (sj)); \
wr[qi] = 1; wr[qj] = 1; \
if (LF) { \
leaf(h, &(Ma), &(Mb), Q[qi], ld, oa, Q[qj], ld, ob); \
} else if (oa == 0) { \
mmxR(h, &(Ma), &(Mb), Q[qi], ld, 0, T); \
cmb(Q[qj], ld, Q[qi], ld, h, (sj), ob); \
} else { \
mmxR(h, &(Ma), &(Mb), T, h, 0, T2); \
cmb2(Q[qi], ld, oa, (si), Q[qj], ld, ob, (sj), T, h, h); \
} \
} while (0) \
#define MK1(qi, si, Ma, Mb) do { \
int oa = wr[qi] ? effop(1, (si)) : effop(op, (si)); \
wr[qi] = 1; \
mmxR(h, &(Ma), &(Mb), Q[qi], ld, oa, T); \
} while (0)
{ /* M1 = (A11+A22)(B11+B22) -> +C11, +C22 */
Opd Ma = oadd(&A11, &A22, 1), Mb = oadd(&B11, &B22, 1);
MK(0, 1, 3, 1, Ma, Mb);
}
{ /* M2 = (A21+A22)B11 -> +C21, -C22 */
Opd Ma = oadd(&A21, &A22, 1), Mb = B11;
MK(2, 1, 3, 2, Ma, Mb);
}
{ /* M3 = A11(B12-B22) -> +C12, +C22 */
Opd Ma = A11, Mb = oadd(&B12, &B22, -1);
MK(1, 1, 3, 1, Ma, Mb);
}
{ /* M4 = A22(B21-B11) -> +C11, +C21 */
Opd Ma = A22, Mb = oadd(&B21, &B11, -1);
MK(0, 1, 2, 1, Ma, Mb);
}
{ /* M5 = (A11+A12)B22 -> -C11, +C12 */
Opd Ma = oadd(&A11, &A12, 1), Mb = B22;
MK(0, 2, 1, 1, Ma, Mb);
}
{ /* M6 = (A21-A11)(B11+B12) -> +C22 */
Opd Ma = oadd(&A21, &A11, -1), Mb = oadd(&B11, &B12, 1);
MK1(3, 1, Ma, Mb);
}
{ /* M7 = (A12-A22)(B21+B22) -> +C11 */
Opd Ma = oadd(&A12, &A22, -1), Mb = oadd(&B21, &B22, 1);
MK1(0, 1, Ma, Mb);
}
#undef MK
#undef MK1
}
static unsigned char PAD40[40u*1024u*1024u] __attribute__((aligned(4096)));
void matrix_multiply(int n, const long long *A, const long long *B, long long *C) {
static volatile long long g_pk6_0; g_pk6_0 = 0; (void)g_pk6_0;
static volatile long long g_pk6_1; g_pk6_1 = 1; (void)g_pk6_1;
static volatile long long g_pk6_2; g_pk6_2 = 2; (void)g_pk6_2;
static volatile long long g_pk6_3; g_pk6_3 = 3; (void)g_pk6_3;
static volatile long long g_pk6_4; g_pk6_4 = 4; (void)g_pk6_4;
static volatile long long g_pk6_5; g_pk6_5 = 5; (void)g_pk6_5;
Opd Ao = o1(A, n), Bo = o1(B, n);
if (n <= CUTL || (n & 1)) { leaf(n, &Ao, &Bo, C, n, 0, 0, 0, 0); return; }
mmxR(n, &Ao, &Bo, C, n, 0, g_scratch);
}
#pragma GCC pop_options
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 817.925 ms | 47 MB + 972 KB | Accepted | Score: 100 | 显示更多 |