// ===== REFERENCES =====
// [1] 本账号现役最好件 **#119349** <https://duck.ac/submission/119349>(529.706644 ms = 本题 mine)
// —— 本件正文 = 该提交正文**逐字节**(md5 50dcc5f672e169a73fe3e22d933e98fd),**唯一改动** =
// 下方"思路"段指名的 pack() 左葉(A オペランド)の**読取順序**(新しい関数 pack_wide /
// pack_wide_root を追加し、pack_root の base case で pair 呼び出しに置換)。値・レイアウト・
// ストア命令・emit 順序はいずれも原状と同一。
// [2] duck.ac 用户 **pdoom** #118374/#118809/#118817(本引擎原始作者)—— 正文 Credit 段逐字保留、未改一字。
// [3] Oded Schwartz & Noa Vaknin, SIAM J. Sci. Comput. (2023), doi 10.1137/22M1502719
// (exact 七乘十二加 alternative-basis 分解)—— 正文内已署名;仅参考其思想,未逐字移植。
// [4] 本席(s4de_,pack 读序残余席)的改动为**原创**:无第三方代码逐字移植。分析基座是**本队自己的**
// 判题机 in-situ 分解 (1525)(PK/MIX/SPK 三臂)与 (1526)(融合否证),二者均为本账号提交的探针读数。
// 合规:duck.ac 提交正文按站点规则公开可见;原提交无独立许可证声明,按 RULES.md 第 3 节署名作者账号与原地址。
// ======================
// ===== 思路 =====
// 【靶子】(1525) 把 pack 相拆成 `PK 39.492e6 拍` = `MIX 30.043e6`(逐址逐序・無 ALU・無 tile)+
// `PK-MIX 9.45e6 = 2.624 ms`(変換 ALU + tile 往復)と `MIX-SPK 5.964e6 = 1.656 ms`(読取順序の
// 代価)。前者は (1526) の融合実測(1.9331× 遅)で判死済み。本席は**後者だけ**を撃つ。
// 【硬制約の遵守】(1526-之二)「NT 存と大批量読流を同一ループ体に置くな」—— 本件は **emit(NT 存)相と
// 読相を完全に分離したまま**、読相の**内部順序**のみを変更する。NaN 1 つも足していない。
// 【単一変数の変更】原状の左葉は 1 つの 32x32 副象限ごとに **64 B/行**(8 KB ストライド)を読み、
// 4 つの副象限 pass に分かれる。本件は**横に隣接する 2 つの 64x64 tile(列 0..63 と 64..127)を
// 1 つの pass で**読み、**256 B/行 を連続消費**する(`pack_wide`)。読むバイト集合・各 tile 内の
// 書き込みレイアウト・ストア命令・emit の順序・アドレスは**すべて原状と同一**で、時間順序だけが変わる。
// 併せて副象限型 `transformed`((s-raw) の bit5 と bit17)を呼び出し幾何から定数化したが、これは
// 単独では**負**(下表 PK_6 = +5.4%)で、勝ち分は読取順序そのものから出ている。
// 【判題機 in-situ(同窓・同一バイナリ・輪転交錯・min-of-N、棒 = pack_root<false>+<true> 暖)】
// 発 1(459.94 ms、5 腕、min-of-3): `PK 39 503 756`((1525) の 39 492 200 を +0.03% で再現)·
// `PK_o`(pack/emit 並べ替え)= 40 216 922 **+1.81%** · `PK_w`(広幅・配列版)= 39 080 824 **-1.07%** ·
// `PK_u`(広幅・アンロール版)= 38 680 192 **-2.09%** · `PK_4`(呼び出し順のみ)= 39 715 272 **+0.54%** ·
// 参照腕 `MIX 30 080 744`・`SPK 24 087 892`・`SRCp 16 876 388`(三臂とも (1525) を 0.86% 以内で再現)。
// 発 2(順序回転・min-of-4): `PK 39 411 936` · `PK_5`(広幅・実行期フラグ)= 38 535 312 **-2.22%** ·
// `PK_3`(広幅・定数フラグ = 本件)= 38 516 800 **-2.27%** · `PK_6`(原順序・定数フラグのみ)= 41 540 696 **+5.40%**。
// ⇒ **本件の利得 = 読取順序**(定数フラグは寄与ゼロ、単独では負)。2 発独立に -2.1%〜-2.3%。
// ⇒ 読取専用腕 `SRCw`(同じバイト集合を 256 B/行で読む)は**判題機で +6.97%(遅)**、本機では
// **-5〜-8%(速)** と**符号が逆転**する ⇒ (1529-之二)「rig 符号は in-situ に外挿できない」を再確認。
// 採否は上記の **in-situ 同バイナリ paired** のみで決めた。
// 【正しさ】(a) **逐位**: 全 pack 産物(A 側 4*RS・B 側 7*RS)の FNV-1a が原状と**一致**
// (判題機 probe 内で ck5/ck6/ck3 すべて 1、値 `19c01dca5634f99b` / `1d2bea4e6fc26270`)。
// (b) **全 C 逐位**: 4 096³ の C を原状エンジンと `cmp` で 0 差分。(c) **正控臂**: 副象限フラグを
// 故意に 1 箇所反転した変異体は ck で**検出される**(探索の歯を確認)。
// ======================
// New pdoom combination: fuse only N64 M1/M3/negated-M7 into C,
// sharing one extra Add leaf body; retain the N128 ABS graph unchanged.
// Suppress automatic internal AVX lane clearing, restore opcode handling
// at API return and execute one explicit VZEROUPPER at that boundary.
// New combination: all recursive linear scans use fixed-relative or alias addresses.
// New combination: fixed-relative sources and in-place destination address reuse in both N64 and N128 ABS scans.
// New route: retain the two packed off-diagonal N64 tiles in L2; fuse N128 input phi into root forms. Raw packing only performs N64 phi.
// New route: stream the outer inverse basis through final NT output, retaining only the outer22 inner-decoded rows.
// New experiment: tensor-product ABS at N64 and N128 composed with current 2x32 cache3, ptr4 combine, and post4.
// Research credit: uops.info Coffee Lake instruction measurements:
// https://uops.info/html-instr/VPADDW_YMM_YMM_M256.html
// https://uops.info/html-instr/VMOVDQA_M256_YMM.html
// Base+displacement addressing avoids indexed memory-ALU unlamination and
// permits the simple store AGU. The new combine loop advances three pointers,
// processes four vectors per step, and skips the 16-word physical leaf gap.
// New pdoom kernel: 32x32 leaves use 2x32 tiles with three cached B vectors.
// The last row consumes B and A registers, keeping all eight products ready
// before their additions. Keep one full k body and one outer loop; bound
// N=64 result reconstruction to four-vector unrolling. All arithmetic is exact.
/*
Credit:
pdoom #118374, https://duck.ac/submission/118374 : exact BASE32 engine,
4-row simultaneous product scheduling, one fully unrolled asm body, and
coarse64 root packing with complete NT cache-line streams.
saffah_cc_v41_agg1 #118613, https://duck.ac/submission/118613 : ROOT_PAD
and workspace address phases, adapted here to the BASE32 recursive graph.
saffah_cc_v41_agg1 #118236, https://duck.ac/submission/118236 : three-madd
scheduling, fused four-quadrant root output, C root raw-input reuse, and
shift/blend truncation with deferred inverse word permutation.
Oded Schwartz and Noa Vaknin, SIAM J. Sci. Comput. (2023), (3.1): exact
seven-product, twelve-add alternative-basis decomposition over a ring.
https://epubs.siam.org/doi/10.1137/22M1502719
Retained lineage:
pdoom #86352 : mod65536 engine and pair-dot products; #112312 : fused root
packing, retired root slot reuse, complete cache-line streams and leaf pads;
#117569 : SIMD-native output and delayed root decode; #118412 : bottom-layer
alternative basis fused at both endpoints and skipping logical leaf padding.
saffah_cc_v41_agg1 #112006/#117247 : peeled initial pair, operand/output
packing, cache phases, three-row scheduling, and compiler options.
saffah_codex_6s_agg2 #110387/#102935 : fused reconstruction, constant
recursion and instruction order. saffah_cc_v41_260924 #96771 : cache phases.
New combination:
Use 32x32 arithmetic leaves and N=64 plus N=128 alternative-basis products.
Apply the input map in raw 32x32 leaf loads, and the inverse output map in
root reconstruction. Internal arithmetic skips the 16-word physical gap;
coarse64 root streams retain it to finish every 64-byte cache line.
Use C's retired tail as the recursive workspace before final C writeback.
Every output is exact modulo65536, independently of input distribution.
Specialized to the stated n=4096 with separate aligned A, B and C buffers.
*/
#define ALTMAX 128
#define ROOT_PAD 160
#pragma GCC optimize("O3,unroll-loops,web,rename-registers")
#pragma GCC target("avx2")
#include <immintrin.h>
#include <stdint.h>
#include <string.h>
asm(".macro vzeroupper\n.endm\n");
#ifndef TILE_PAD
#define TILE_PAD 16
#endif
#ifndef BASE
#define BASE (1 << 5)
#endif
using U = uint16_t;
static constexpr int MAXN = 4096;
alignas(4096) static U pa_[MAXN * MAXN + (MAXN/BASE)*(MAXN/BASE)*TILE_PAD + 8192], pb_[MAXN * MAXN + (MAXN/BASE)*(MAXN/BASE)*TILE_PAD + 8192];
alignas(4096) static U pc_[MAXN * MAXN + (MAXN/BASE)*(MAXN/BASE)*TILE_PAD + 8192], ws_[MAXN * MAXN + (MAXN/BASE)*(MAXN/BASE)*TILE_PAD + 8192];
static U *pa = pa_ + 1024, *pb = pb_ + 1024, *pc = pc_ + 0, *ws = ws_ + 1024;
static U *final_out;
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 const U *raw_basis_input;
static inline __m256i raw_basis_low32(const U *p,bool transformed) {
__m256i v=_mm256_loadu_si256((const __m256i*)p);
if(transformed) v=_mm256_add_epi16(v,_mm256_sub_epi16(
_mm256_loadu_si256((const __m256i*)(p-32*4096)),
_mm256_loadu_si256((const __m256i*)(p-32))));
return v;
}
static inline __m128i raw_basis_low16(const U *p,bool transformed) {
__m128i v=_mm_loadu_si128((const __m128i*)p);
if(transformed) v=_mm_add_epi16(v,_mm_sub_epi16(
_mm_loadu_si128((const __m128i*)(p-32*4096)),
_mm_loadu_si128((const __m128i*)(p-32))));
return v;
}
static inline __m256i raw_basis32(const U *p,int flags) {
__m256i v=raw_basis_low32(p,flags&1);
if(flags&2) v=_mm256_add_epi16(v,_mm256_sub_epi16(
raw_basis_low32(p-64*4096,flags&1),raw_basis_low32(p-64,flags&1)));
return v;
}
static inline __m128i raw_basis16(const U *p,int flags) {
__m128i v=raw_basis_low16(p,flags&1);
if(flags&2) v=_mm_add_epi16(v,_mm_sub_epi16(
raw_basis_low16(p-64*4096,flags&1),raw_basis_low16(p-64,flags&1)));
return v;
}
static void pack(U *d, const U *s, int n, int stride, bool right) {
if (n > BASE) {
int h = n / 2, q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
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) {
const int transformed=(((s-raw_basis_input)&32) && ((s-raw_basis_input)&(32*4096)))?1:0;
for(int i=0;i<n;i+=4) for(int k=0;k<n;k+=16) {
{ int i2=(i+4)&31; /* s4ct-PKLEAF64: rows i+4..i+7, both k blocks */
__builtin_prefetch(s+i2*stride+k);__builtin_prefetch(s+(i2+1)*stride+k);
__builtin_prefetch(s+i2*stride+k+16);__builtin_prefetch(s+(i2+1)*stride+k+16);
__builtin_prefetch(s+(i2+2)*stride+k);__builtin_prefetch(s+(i2+3)*stride+k);
__builtin_prefetch(s+(i2+2)*stride+k+16);__builtin_prefetch(s+(i2+3)*stride+k+16); }
__m256i a=raw_basis32(s+i*stride+k,transformed), b=raw_basis32(s+(i+1)*stride+k,transformed);
__m256i c=raw_basis32(s+(i+2)*stride+k,transformed), e=raw_basis32(s+(i+3)*stride+k,transformed);
__m256i t0=_mm256_unpacklo_epi32(a,b),t1=_mm256_unpacklo_epi32(c,e);
__m256i t2=_mm256_unpackhi_epi32(a,b),t3=_mm256_unpackhi_epi32(c,e);
__m256i u0=_mm256_unpacklo_epi64(t0,t1),u1=_mm256_unpackhi_epi64(t0,t1);
__m256i u2=_mm256_unpacklo_epi64(t2,t3),u3=_mm256_unpackhi_epi64(t2,t3);
st(d,_mm256_permute2x128_si256(u0,u1,0x20));
st(d+16,_mm256_permute2x128_si256(u2,u3,0x20));
st(d+32,_mm256_permute2x128_si256(u0,u1,0x31));
st(d+48,_mm256_permute2x128_si256(u2,u3,0x31));d+=64;
}
} else {
const int transformed=(((s-raw_basis_input)&32) && ((s-raw_basis_input)&(32*4096)))?1:0;
for(int k=0;k<n;k+=2) for(int j=0;j<n;) {
int w=32;
U *v=d+j*n+k*w;
for(int t=0;t<w;t+=16) {
if(t+16<=w) {
__m256i x=raw_basis32(s+k*stride+j+t,transformed);
__m256i y=raw_basis32(s+(k+1)*stride+j+t,transformed);
st(v,_mm256_unpacklo_epi16(x,y));
st(v+16,_mm256_unpackhi_epi16(x,y));v+=32;
} else {
__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));
_mm_store_si128((__m128i*)v,_mm_unpacklo_epi16(x,y));
_mm_store_si128((__m128i*)(v+8),_mm_unpackhi_epi16(x,y));v+=16;
}
}
j+=w;
}
}
}
// ===== s4de_ : pack 左葉の広幅読み(読取順序のみ・単一変数) =====
// 原状: 1 つの 32x32 副象限ごとに 64 B/行 を読む(4 回の pass に分かれる)。
// 本件: 64 列ぶんの 2 tile(列 0..63 と 64..127)を 1 つの pass で読み、256 B/行 を連続で消費する。
// 読み書きするバイト集合・各 tile 内の書き込みレイアウト・ストア命令は原状と逐位同一。
// transformed は実行期判定((s-raw) の bit5 と bit17)を、呼び出し幾何から
// 副象限 (行off==32 && 列off==32) に等価(証: 他の全オフセット 2048/2048*4096/64/64*4096/
// 128i/128j*4096 は bit5 にも bit17 にも寄与しない)と定数化した。ck で逐位一致を確認済み。
#define PWL_Q ((32*32) + (32/BASE)*(32/BASE)*TILE_PAD)
#define PWL_CH(DSTP,K,TF) do { \
__m256i a=raw_basis32(sr+(long)i*S4ST+(K),(TF)), b=raw_basis32(sr+(long)(i+1)*S4ST+(K),(TF)); \
__m256i c=raw_basis32(sr+(long)(i+2)*S4ST+(K),(TF)), e=raw_basis32(sr+(long)(i+3)*S4ST+(K),(TF)); \
__m256i t0=_mm256_unpacklo_epi32(a,b),t1=_mm256_unpacklo_epi32(c,e); \
__m256i t2=_mm256_unpackhi_epi32(a,b),t3=_mm256_unpackhi_epi32(c,e); \
__m256i u0=_mm256_unpacklo_epi64(t0,t1),u1=_mm256_unpackhi_epi64(t0,t1); \
__m256i u2=_mm256_unpacklo_epi64(t2,t3),u3=_mm256_unpackhi_epi64(t2,t3); \
st((DSTP)+0,_mm256_permute2x128_si256(u0,u1,0x20)); \
st((DSTP)+16,_mm256_permute2x128_si256(u2,u3,0x20)); \
st((DSTP)+32,_mm256_permute2x128_si256(u0,u1,0x31)); \
st((DSTP)+48,_mm256_permute2x128_si256(u2,u3,0x31)); } while(0)
static void pack_wide(U *dA,U *dB,const U *s) {
const long S4ST=4096;
for (int r=0;r<2;++r) {
const U *sr = s + (long)r*32*S4ST;
U *p0=dA+(r?2*PWL_Q:0), *p1=dA+(r?3*PWL_Q:PWL_Q), *p2=dB+(r?2*PWL_Q:0), *p3=dB+(r?3*PWL_Q:PWL_Q);
const int TFc = (r==1)?1:0;
for (int i=0;i<32;i+=4) {
{ int i2=(i+4)&31;
__builtin_prefetch(sr+(long)i2*S4ST); __builtin_prefetch(sr+(long)(i2+1)*S4ST);
__builtin_prefetch(sr+(long)(i2+2)*S4ST); __builtin_prefetch(sr+(long)(i2+3)*S4ST);
__builtin_prefetch(sr+(long)i2*S4ST+16); __builtin_prefetch(sr+(long)(i2+1)*S4ST+16);
__builtin_prefetch(sr+(long)(i2+2)*S4ST+16); __builtin_prefetch(sr+(long)(i2+3)*S4ST+16); }
PWL_CH(p0, 0,0); PWL_CH(p0+64, 16,0); p0+=128;
PWL_CH(p1, 32,TFc); PWL_CH(p1+64, 48,TFc); p1+=128;
PWL_CH(p2, 64,0); PWL_CH(p2+64, 80,0); p2+=128;
PWL_CH(p3, 96,TFc); PWL_CH(p3+64,112,TFc); p3+=128;
}
}
}
static void unpack(U *d, const U *s, int n, int stride) {
if (n > BASE) {
int h = n/2, q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
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) {
U *dp = d + (size_t)i*stride; const U *sp = s + (size_t)i*n;
int k=0;
for (;k+32<=n;k+=32) {
_mm256_stream_si256((__m256i*)(dp+k), _mm256_load_si256((const __m256i*)(sp+k)));
_mm256_stream_si256((__m256i*)(dp+k+16), _mm256_load_si256((const __m256i*)(sp+k+16)));
}
for (;k<n;k+=16) _mm256_storeu_si256((__m256i*)(dp+k), _mm256_loadu_si256((const __m256i*)(sp+k)));
}
}
}
// ===== s4bj_ : 叶尾改 blend 后,每个 16 字组内被 sigma 置换,decode 前先复原 =====
// sigma = [0,4,1,5,2,6,3,7 | 8,12,9,13,10,14,11,15] (new[j] = old[sigma(j)])
// sigma^-1= [0,2,4,6,1,3,5,7 | 8,10,12,14,9,11,13,15] (old[i] = new[sigma^-1(i)])
alignas(32) static const unsigned char s4bj_si[32] = {
0,1,4,5,8,9,12,13,2,3,6,7,10,11,14,15, 0,1,4,5,8,9,12,13,2,3,6,7,10,11,14,15};
static inline __m256i s4bj_fx(__m256i v) {
return _mm256_shuffle_epi8(v, _mm256_load_si256((const __m256i *)s4bj_si));
}
alignas(32) unsigned short s4aj_m16[16] = {0xFFFF,0,0xFFFF,0,0xFFFF,0,0xFFFF,0,0xFFFF,0,0xFFFF,0,0xFFFF,0,0xFFFF,0};
alignas(32) unsigned char s4ak_sh16[32] = {0,1,4,5,8,9,12,13, 0x80,0x80,0x80,0x80,0x80,0x80,0x80,0x80,0,1,4,5,8,9,12,13, 0x80,0x80,0x80,0x80,0x80,0x80,0x80,0x80};
// One explicit panel/row loop keeps exactly one copy of the large kernel.
template<int Mode>
static inline void native2x32_cache3(const U *a,const U *b,U *c) {
const U *a_final=a+1024;
intptr_t astep=8; // byte steps8,248 alternate within and across4-row A packs
asm volatile(
"3:\n\t"
"vmovdqa 0(%[b]), %%ymm8\n\t"
"vmovdqa 32(%[b]), %%ymm9\n\t"
"vmovdqa 64(%[b]), %%ymm10\n\t"
"vpbroadcastd 0(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm0\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm1\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm2\n\t"
"vpmaddwd 96(%[b]), %%ymm11, %%ymm3\n\t"
"vpbroadcastd 4(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm4\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm5\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm6\n\t"
"vpmaddwd 96(%[b]), %%ymm11, %%ymm7\n\t"
"vmovdqa 128(%[b]), %%ymm8\n\t"
"vmovdqa 160(%[b]), %%ymm9\n\t"
"vmovdqa 192(%[b]), %%ymm10\n\t"
"vpbroadcastd 16(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 224(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 20(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 224(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 256(%[b]), %%ymm8\n\t"
"vmovdqa 288(%[b]), %%ymm9\n\t"
"vmovdqa 320(%[b]), %%ymm10\n\t"
"vpbroadcastd 32(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 352(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 36(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 352(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 384(%[b]), %%ymm8\n\t"
"vmovdqa 416(%[b]), %%ymm9\n\t"
"vmovdqa 448(%[b]), %%ymm10\n\t"
"vpbroadcastd 48(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 480(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 52(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 480(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 512(%[b]), %%ymm8\n\t"
"vmovdqa 544(%[b]), %%ymm9\n\t"
"vmovdqa 576(%[b]), %%ymm10\n\t"
"vpbroadcastd 64(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 608(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 68(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 608(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 640(%[b]), %%ymm8\n\t"
"vmovdqa 672(%[b]), %%ymm9\n\t"
"vmovdqa 704(%[b]), %%ymm10\n\t"
"vpbroadcastd 80(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 736(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 84(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 736(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 768(%[b]), %%ymm8\n\t"
"vmovdqa 800(%[b]), %%ymm9\n\t"
"vmovdqa 832(%[b]), %%ymm10\n\t"
"vpbroadcastd 96(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 864(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 100(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 864(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 896(%[b]), %%ymm8\n\t"
"vmovdqa 928(%[b]), %%ymm9\n\t"
"vmovdqa 960(%[b]), %%ymm10\n\t"
"vpbroadcastd 112(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 992(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 116(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 992(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1024(%[b]), %%ymm8\n\t"
"vmovdqa 1056(%[b]), %%ymm9\n\t"
"vmovdqa 1088(%[b]), %%ymm10\n\t"
"vpbroadcastd 128(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 1120(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 132(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 1120(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1152(%[b]), %%ymm8\n\t"
"vmovdqa 1184(%[b]), %%ymm9\n\t"
"vmovdqa 1216(%[b]), %%ymm10\n\t"
"vpbroadcastd 144(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 1248(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 148(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 1248(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1280(%[b]), %%ymm8\n\t"
"vmovdqa 1312(%[b]), %%ymm9\n\t"
"vmovdqa 1344(%[b]), %%ymm10\n\t"
"vpbroadcastd 160(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 1376(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 164(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 1376(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1408(%[b]), %%ymm8\n\t"
"vmovdqa 1440(%[b]), %%ymm9\n\t"
"vmovdqa 1472(%[b]), %%ymm10\n\t"
"vpbroadcastd 176(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 1504(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 180(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 1504(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1536(%[b]), %%ymm8\n\t"
"vmovdqa 1568(%[b]), %%ymm9\n\t"
"vmovdqa 1600(%[b]), %%ymm10\n\t"
"vpbroadcastd 192(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 1632(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 196(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 1632(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1664(%[b]), %%ymm8\n\t"
"vmovdqa 1696(%[b]), %%ymm9\n\t"
"vmovdqa 1728(%[b]), %%ymm10\n\t"
"vpbroadcastd 208(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 1760(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 212(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 1760(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1792(%[b]), %%ymm8\n\t"
"vmovdqa 1824(%[b]), %%ymm9\n\t"
"vmovdqa 1856(%[b]), %%ymm10\n\t"
"vpbroadcastd 224(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 1888(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 228(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 1888(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vmovdqa 1920(%[b]), %%ymm8\n\t"
"vmovdqa 1952(%[b]), %%ymm9\n\t"
"vmovdqa 1984(%[b]), %%ymm10\n\t"
"vpbroadcastd 240(%[a]), %%ymm11\n\t"
"vpmaddwd %%ymm8, %%ymm11, %%ymm12\n\t"
"vpmaddwd %%ymm9, %%ymm11, %%ymm13\n\t"
"vpmaddwd %%ymm10, %%ymm11, %%ymm14\n\t"
"vpmaddwd 2016(%[b]), %%ymm11, %%ymm11\n\t"
"vpbroadcastd 244(%[a]), %%ymm15\n\t"
"vpmaddwd %%ymm8, %%ymm15, %%ymm8\n\t"
"vpmaddwd %%ymm9, %%ymm15, %%ymm9\n\t"
"vpmaddwd %%ymm10, %%ymm15, %%ymm10\n\t"
"vpmaddwd 2016(%[b]), %%ymm15, %%ymm15\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpaddd %%ymm14, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm11, %%ymm3, %%ymm3\n\t"
"vpaddd %%ymm8, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm9, %%ymm5, %%ymm5\n\t"
"vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
"vpslld $16, %%ymm1, %%ymm1\n\t"
"vpblendw $0xAA, %%ymm1, %%ymm0, %%ymm0\n\t"
".if %c[mode] == 1\n\t"
"vpaddw 0(%[c]), %%ymm0, %%ymm0\n\t"
".endif\n\t"
"vmovdqu %%ymm0, 0(%[c])\n\t"
"vpslld $16, %%ymm3, %%ymm3\n\t"
"vpblendw $0xAA, %%ymm3, %%ymm2, %%ymm2\n\t"
".if %c[mode] == 1\n\t"
"vpaddw 32(%[c]), %%ymm2, %%ymm2\n\t"
".endif\n\t"
"vmovdqu %%ymm2, 32(%[c])\n\t"
"vpslld $16, %%ymm5, %%ymm5\n\t"
"vpblendw $0xAA, %%ymm5, %%ymm4, %%ymm4\n\t"
".if %c[mode] == 1\n\t"
"vpaddw 64(%[c]), %%ymm4, %%ymm4\n\t"
".endif\n\t"
"vmovdqu %%ymm4, 64(%[c])\n\t"
"vpslld $16, %%ymm7, %%ymm7\n\t"
"vpblendw $0xAA, %%ymm7, %%ymm6, %%ymm6\n\t"
".if %c[mode] == 1\n\t"
"vpaddw 96(%[c]), %%ymm6, %%ymm6\n\t"
".endif\n\t"
"vmovdqu %%ymm6, 96(%[c])\n\t"
"add %[astep], %[a]\n\t"
"xor $240, %[astep]\n\t"
"add $128, %[c]\n\t"
"cmp %[a_final], %[a]\n\t"
"jb 3b\n\t"
: [a] "+&r"(a), [c] "+&r"(c), [astep] "+&r"(astep)
: [b] "r"(b), [a_final] "r"(a_final), [mode] "i"(Mode)
: "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) {
native2x32_cache3<0>(a,b,c);
}
// Only one additional hot body: M1, M3 and the negated-M7 update all add.
static __attribute__((noinline)) void leaf_add(const U*a,const U*b,U*c){
native2x32_cache3<1>(a,b,c);
}
static inline void emit_native(U *dst,const U *s,int stride) {
for(int row=0;row<4;++row) {
U *d=dst+(size_t)row*stride;
_mm256_stream_si256((__m256i*)(d+0),s4bj_fx(ld(s+32*row)));
_mm256_stream_si256((__m256i*)(d+16),s4bj_fx(ld(s+32*row+16)));
}
}
template<bool Sub>
static inline void combine(U *d,const U *a,const U *b,int len) {
#pragma GCC unroll 1
for(int chunk=0;chunk<len;chunk+=1024+TILE_PAD) {
const U *end=a+1024;
if constexpr(Sub) {
asm volatile(
"1:\n\t"
"vmovdqa 0(%[a]), %%ymm0\n\t"
"vmovdqa 32(%[a]), %%ymm1\n\t"
"vmovdqa 64(%[a]), %%ymm2\n\t"
"vmovdqa 96(%[a]), %%ymm3\n\t"
"vpsubw 0(%[b]), %%ymm0, %%ymm0\n\t"
"vpsubw 32(%[b]), %%ymm1, %%ymm1\n\t"
"vpsubw 64(%[b]), %%ymm2, %%ymm2\n\t"
"vpsubw 96(%[b]), %%ymm3, %%ymm3\n\t"
"vmovdqa %%ymm0, 0(%[d])\n\t"
"vmovdqa %%ymm1, 32(%[d])\n\t"
"vmovdqa %%ymm2, 64(%[d])\n\t"
"vmovdqa %%ymm3, 96(%[d])\n\t"
"add $128, %[a]\n\t"
"add $128, %[b]\n\t"
"add $128, %[d]\n\t"
"cmp %[end], %[a]\n\t"
"jb 1b\n\t"
: [a] "+&r"(a), [b] "+&r"(b), [d] "+&r"(d)
: [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3");
} else {
asm volatile(
"1:\n\t"
"vmovdqa 0(%[a]), %%ymm0\n\t"
"vmovdqa 32(%[a]), %%ymm1\n\t"
"vmovdqa 64(%[a]), %%ymm2\n\t"
"vmovdqa 96(%[a]), %%ymm3\n\t"
"vpaddw 0(%[b]), %%ymm0, %%ymm0\n\t"
"vpaddw 32(%[b]), %%ymm1, %%ymm1\n\t"
"vpaddw 64(%[b]), %%ymm2, %%ymm2\n\t"
"vpaddw 96(%[b]), %%ymm3, %%ymm3\n\t"
"vmovdqa %%ymm0, 0(%[d])\n\t"
"vmovdqa %%ymm1, 32(%[d])\n\t"
"vmovdqa %%ymm2, 64(%[d])\n\t"
"vmovdqa %%ymm3, 96(%[d])\n\t"
"add $128, %[a]\n\t"
"add $128, %[b]\n\t"
"add $128, %[d]\n\t"
"cmp %[end], %[a]\n\t"
"jb 1b\n\t"
: [a] "+&r"(a), [b] "+&r"(b), [d] "+&r"(d)
: [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3");
}
a+=TILE_PAD;b+=TILE_PAD;d+=TILE_PAD;
}
}
template<bool Sub,int Delta>
static inline void combine_displaced(U *d,const U *a,int len) {
#pragma GCC unroll 1
for(int chunk=0;chunk<len;chunk+=1024+TILE_PAD) {
const U *end=a+1024;
if constexpr(Sub) {
asm volatile("1:\n\t"
"vmovdqa 0(%[a]), %%ymm0\n\t"
"vmovdqa 32(%[a]), %%ymm1\n\t"
"vmovdqa 64(%[a]), %%ymm2\n\t"
"vmovdqa 96(%[a]), %%ymm3\n\t"
"vpsubw %c[delta]+0(%[a]), %%ymm0, %%ymm0\n\t"
"vpsubw %c[delta]+32(%[a]), %%ymm1, %%ymm1\n\t"
"vpsubw %c[delta]+64(%[a]), %%ymm2, %%ymm2\n\t"
"vpsubw %c[delta]+96(%[a]), %%ymm3, %%ymm3\n\t"
"vmovdqa %%ymm0, 0(%[d])\n\t"
"vmovdqa %%ymm1, 32(%[d])\n\t"
"vmovdqa %%ymm2, 64(%[d])\n\t"
"vmovdqa %%ymm3, 96(%[d])\n\t"
"add $128, %[a]\n\t"
"add $128, %[d]\n\t"
"cmp %[end], %[a]\n\t"
"jb 1b\n\t"
: [a] "+&r"(a), [d] "+&r"(d)
: [end] "r"(end), [delta] "i"(2*Delta)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3");
} else {
asm volatile("1:\n\t"
"vmovdqa 0(%[a]), %%ymm0\n\t"
"vmovdqa 32(%[a]), %%ymm1\n\t"
"vmovdqa 64(%[a]), %%ymm2\n\t"
"vmovdqa 96(%[a]), %%ymm3\n\t"
"vpaddw %c[delta]+0(%[a]), %%ymm0, %%ymm0\n\t"
"vpaddw %c[delta]+32(%[a]), %%ymm1, %%ymm1\n\t"
"vpaddw %c[delta]+64(%[a]), %%ymm2, %%ymm2\n\t"
"vpaddw %c[delta]+96(%[a]), %%ymm3, %%ymm3\n\t"
"vmovdqa %%ymm0, 0(%[d])\n\t"
"vmovdqa %%ymm1, 32(%[d])\n\t"
"vmovdqa %%ymm2, 64(%[d])\n\t"
"vmovdqa %%ymm3, 96(%[d])\n\t"
"add $128, %[a]\n\t"
"add $128, %[d]\n\t"
"cmp %[end], %[a]\n\t"
"jb 1b\n\t"
: [a] "+&r"(a), [d] "+&r"(d)
: [end] "r"(end), [delta] "i"(2*Delta)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3");
}
a+=TILE_PAD;d+=TILE_PAD;
}
}
// pdoom frontier6: the output is its own first operand, so only two
// advancing addresses are necessary. N64 calls touch a single valid leaf.
template<bool Sub>
static inline void combine_inplace(U *d,const U *b,int len) {
#pragma GCC unroll 1
for(int chunk=0;chunk<len;chunk+=1024+TILE_PAD) {
const U *end=d+1024;
if constexpr(Sub) {
asm volatile("1:\n\t"
"vmovdqa 0(%[d]), %%ymm0\n\t"
"vmovdqa 32(%[d]), %%ymm1\n\t"
"vmovdqa 64(%[d]), %%ymm2\n\t"
"vmovdqa 96(%[d]), %%ymm3\n\t"
"vpsubw 0(%[b]), %%ymm0, %%ymm0\n\t"
"vpsubw 32(%[b]), %%ymm1, %%ymm1\n\t"
"vpsubw 64(%[b]), %%ymm2, %%ymm2\n\t"
"vpsubw 96(%[b]), %%ymm3, %%ymm3\n\t"
"vmovdqa %%ymm0, 0(%[d])\n\t"
"vmovdqa %%ymm1, 32(%[d])\n\t"
"vmovdqa %%ymm2, 64(%[d])\n\t"
"vmovdqa %%ymm3, 96(%[d])\n\t"
"add $128, %[d]\n\t"
"add $128, %[b]\n\t"
"cmp %[end], %[d]\n\t"
"jb 1b\n\t"
: [d] "+&r"(d), [b] "+&r"(b)
: [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3");
} else {
asm volatile("1:\n\t"
"vmovdqa 0(%[d]), %%ymm0\n\t"
"vmovdqa 32(%[d]), %%ymm1\n\t"
"vmovdqa 64(%[d]), %%ymm2\n\t"
"vmovdqa 96(%[d]), %%ymm3\n\t"
"vpaddw 0(%[b]), %%ymm0, %%ymm0\n\t"
"vpaddw 32(%[b]), %%ymm1, %%ymm1\n\t"
"vpaddw 64(%[b]), %%ymm2, %%ymm2\n\t"
"vpaddw 96(%[b]), %%ymm3, %%ymm3\n\t"
"vmovdqa %%ymm0, 0(%[d])\n\t"
"vmovdqa %%ymm1, 32(%[d])\n\t"
"vmovdqa %%ymm2, 64(%[d])\n\t"
"vmovdqa %%ymm3, 96(%[d])\n\t"
"add $128, %[d]\n\t"
"add $128, %[b]\n\t"
"cmp %[end], %[d]\n\t"
"jb 1b\n\t"
: [d] "+&r"(d), [b] "+&r"(b)
: [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3");
}
d+=TILE_PAD;b+=TILE_PAD;
}
}
// pdoom frontier6: the output is its own first operand, so only two
// advancing addresses are necessary. N64 calls touch a single valid leaf.
static inline void combine_reverse_inplace(U *d,const U *b,int len) {
#pragma GCC unroll 1
for(int chunk=0;chunk<len;chunk+=1024+TILE_PAD) {
const U *end=d+1024;
asm volatile("1:\n\t"
"vmovdqa 0(%[b]), %%ymm0\n\t"
"vmovdqa 32(%[b]), %%ymm1\n\t"
"vmovdqa 64(%[b]), %%ymm2\n\t"
"vmovdqa 96(%[b]), %%ymm3\n\t"
"vpsubw 0(%[d]), %%ymm0, %%ymm0\n\t"
"vpsubw 32(%[d]), %%ymm1, %%ymm1\n\t"
"vpsubw 64(%[d]), %%ymm2, %%ymm2\n\t"
"vpsubw 96(%[d]), %%ymm3, %%ymm3\n\t"
"vmovdqa %%ymm0, 0(%[d])\n\t"
"vmovdqa %%ymm1, 32(%[d])\n\t"
"vmovdqa %%ymm2, 64(%[d])\n\t"
"vmovdqa %%ymm3, 96(%[d])\n\t"
"add $128, %[d]\n\t"
"add $128, %[b]\n\t"
"cmp %[end], %[d]\n\t"
"jb 1b\n\t"
: [d] "+&r"(d), [b] "+&r"(b)
: [end] "r"(end)
: "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3");
d+=TILE_PAD;b+=TILE_PAD;
}
}
// At the root, the three completed packed quadrants can be sent straight
// to their final row-major destinations rather than materialized in pc.
template<int N>
static __attribute__((noinline)) void combine3_unpack(
U *d12,U *d21,U *d22,
const U *p6,const U *p7,const U *p5,
const U *p1,const U *p4,const U *p3,int stride) {
if constexpr (N > BASE) {
constexpr int h=N/2,q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
combine3_unpack<h>(d12,d21,d22,p6,p7,p5,p1,p4,p3,stride);
combine3_unpack<h>(d12+h,d21+h,d22+h,p6+q,p7+q,p5+q,p1+q,p4+q,p3+q,stride);
combine3_unpack<h>(d12+h*stride,d21+h*stride,d22+h*stride,
p6+2*q,p7+2*q,p5+2*q,p1+2*q,p4+2*q,p3+2*q,stride);
combine3_unpack<h>(d12+h*stride+h,d21+h*stride+h,d22+h*stride+h,
p6+3*q,p7+3*q,p5+3*q,p1+3*q,p4+3*q,p3+3*q,stride);
} else {
alignas(32) U temp[3*128];
for(int block=0;block<1024;block=block+128) {
for(int v=0;v<128;v+=16) {
int idx=block+v;
__m256i u2=_mm256_add_epi16(ld(p1+idx),ld(p6+idx));
__m256i u3=_mm256_add_epi16(u2,ld(p7+idx));
__m256i v5=ld(p5+idx);
st(temp+v,_mm256_add_epi16(_mm256_add_epi16(u2,v5),ld(p3+idx)));
st(temp+128+v,_mm256_sub_epi16(u3,ld(p4+idx)));
st(temp+256+v,_mm256_add_epi16(u3,v5));
}
size_t off=(size_t)(block/32)*stride;
emit_native(d12+off,temp,stride);
emit_native(d21+off,temp+128,stride);
emit_native(d22+off,temp+256,stride);
}
}
}
// Add the surviving P1 into the root C11 product while mapping its packed
// quadrant layout to the final row-major C11 cells.
template<int N>
static __attribute__((noinline)) void unpack_add(U *d,const U *p2,const U *p1,int stride) {
if constexpr(N>BASE) {
constexpr int h=N/2,q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
unpack_add<h>(d,p2,p1,stride);
unpack_add<h>(d+h,p2+q,p1+q,stride);
unpack_add<h>(d+h*stride,p2+2*q,p1+2*q,stride);
unpack_add<h>(d+h*stride+h,p2+3*q,p1+3*q,stride);
} else {
alignas(32) U temp[128];
for(int block=0;block<1024;block+=128) {
for(int v=0;v<128;v+=16)
st(temp+v,_mm256_add_epi16(ld(p2+block+v),ld(p1+block+v)));
emit_native(d+(size_t)(block/32)*stride,temp,stride);
}
}
}
// Strassen-Winograd: seven products, fifteen additions, two temporaries.
// All operations are in Z/(2^16); truncation never changes the answer.
template<int N>
static __attribute__((noinline)) void root_fuse(
U *d11,U *d12,U *d21,U *d22,
const U *p1,const U *p2,const U *p3,const U *p4,const U *p5,const U *p6,const U *p7,int stride) {
if constexpr(N==128) {
constexpr int q=BASE*BASE+TILE_PAD;
root_fuse<64>(d11,d12,d21,d22,p1,p2,p3,p4,p5,p6,p7,stride);
alignas(32) U keep[4*512], cur[4*512];
U *dest[4]={d11,d12,d21,d22};
for(int block=0;block<1024;block+=128) {
auto make_inner=[&](U *tmp,int group) {
for(int child=0;child<4;child++) {
for(int v=0;v<128;v+=16) {
int idx=(4*group+child)*q+block+v;
if(!(v&16)){__builtin_prefetch(p1+idx+128,0,3);__builtin_prefetch(p2+idx+128,0,3);__builtin_prefetch(p3+idx+128,0,3);__builtin_prefetch(p4+idx+128,0,3);__builtin_prefetch(p5+idx+128,0,3);__builtin_prefetch(p6+idx+128,0,3);__builtin_prefetch(p7+idx+128,0,3);}
__m256i v1=ld(p1+idx),v5=ld(p5+idx);
__m256i u2=_mm256_add_epi16(v1,ld(p6+idx));
__m256i u3=_mm256_add_epi16(u2,ld(p7+idx));
U *t=tmp+child*512;
st(t+v,_mm256_add_epi16(v1,ld(p2+idx)));
st(t+128+v,_mm256_add_epi16(_mm256_add_epi16(u2,v5),ld(p3+idx)));
st(t+256+v,_mm256_sub_epi16(u3,ld(p4+idx)));
st(t+384+v,_mm256_add_epi16(u3,v5));
}
}
for(int v=0;v<512;v+=16) {
__m256i c22=ld(tmp+3*512+v);
st(tmp+1*512+v,_mm256_sub_epi16(ld(tmp+1*512+v),c22));
st(tmp+2*512+v,_mm256_sub_epi16(c22,ld(tmp+2*512+v)));
}
};
size_t off=(size_t)(block/32)*stride;
make_inner(keep,3);
for(int child=0;child<4;child++) {
int row=64+((child>>1)&1)*32,col=64+(child&1)*32;
for(int r=0;r<4;r++)emit_native(dest[r]+off+row*stride+col,keep+child*512+r*128,stride);
}
for(int group=1;group<=2;group++) {
make_inner(cur,group);
for(int child=0;child<4;child++) {
int row=((group>>1)&1)*64+((child>>1)&1)*32;
int col=(group&1)*64+(child&1)*32;
for(int r=0;r<4;r++)for(int rr=0;rr<4;rr++) {
U *d=dest[r]+off+(row+rr)*stride+col;
const U *x=cur+child*512+r*128+rr*32;
const U *y=keep+child*512+r*128+rr*32;
__m256i v0,v1;
if(group==1) {
v0=_mm256_sub_epi16(ld(x),ld(y));
v1=_mm256_sub_epi16(ld(x+16),ld(y+16));
} else {
v0=_mm256_sub_epi16(ld(y),ld(x));
v1=_mm256_sub_epi16(ld(y+16),ld(x+16));
}
_mm256_stream_si256((__m256i*)d,s4bj_fx(v0));
_mm256_stream_si256((__m256i*)(d+16),s4bj_fx(v1));
}
}
}
}
} else if constexpr(N==64) {
constexpr int q=BASE*BASE+TILE_PAD;
root_fuse<32>(d11,d12,d21,d22,p1,p2,p3,p4,p5,p6,p7,stride);
alignas(32) U tmp[3*4*128];
U *dest[4]={d11,d12,d21,d22};
for(int block=0;block<1024;block+=128) {
for(int child=1;child<=3;child++) {
for(int v=0;v<128;v+=16) {
int idx=child*q+block+v;
if(!(v&16)){__builtin_prefetch(p1+idx+128,0,3);__builtin_prefetch(p2+idx+128,0,3);__builtin_prefetch(p3+idx+128,0,3);__builtin_prefetch(p4+idx+128,0,3);__builtin_prefetch(p5+idx+128,0,3);__builtin_prefetch(p6+idx+128,0,3);__builtin_prefetch(p7+idx+128,0,3);}
__m256i v1=ld(p1+idx),v5=ld(p5+idx);
__m256i u2=_mm256_add_epi16(v1,ld(p6+idx));
__m256i u3=_mm256_add_epi16(u2,ld(p7+idx));
U *t=tmp+(child-1)*512;
st(t+v,_mm256_add_epi16(v1,ld(p2+idx)));
st(t+128+v,_mm256_add_epi16(_mm256_add_epi16(u2,v5),ld(p3+idx)));
st(t+256+v,_mm256_sub_epi16(u3,ld(p4+idx)));
st(t+384+v,_mm256_add_epi16(u3,v5));
}
}
for(int v=0;v<512;v+=16) {
__m256i c22=ld(tmp+1024+v);
st(tmp+v,_mm256_sub_epi16(ld(tmp+v),c22));
st(tmp+512+v,_mm256_sub_epi16(c22,ld(tmp+512+v)));
}
size_t off=(size_t)(block/32)*stride;
for(int r=0;r<4;r++) {
emit_native(dest[r]+off+32,tmp+r*128,stride);
emit_native(dest[r]+off+32*stride,tmp+512+r*128,stride);
emit_native(dest[r]+off+32*stride+32,tmp+1024+r*128,stride);
}
}
} else if constexpr (N > BASE) {
constexpr int h=N/2,q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
root_fuse<h>(d11,d12,d21,d22,p1,p2,p3,p4,p5,p6,p7,stride);
root_fuse<h>(d11+h,d12+h,d21+h,d22+h,p1+q,p2+q,p3+q,p4+q,p5+q,p6+q,p7+q,stride);
root_fuse<h>(d11+h*stride,d12+h*stride,d21+h*stride,d22+h*stride,
p1+2*q,p2+2*q,p3+2*q,p4+2*q,p5+2*q,p6+2*q,p7+2*q,stride);
root_fuse<h>(d11+h*stride+h,d12+h*stride+h,d21+h*stride+h,d22+h*stride+h,
p1+3*q,p2+3*q,p3+3*q,p4+3*q,p5+3*q,p6+3*q,p7+3*q,stride);
} else {
alignas(32) U temp[4*128];
for(int block=0;block<1024;block+=128) {
for(int v=0;v<128;v+=16) {
int idx=block+v;
if(!(v&16)){__builtin_prefetch(p1+idx+128,0,3);__builtin_prefetch(p2+idx+128,0,3);__builtin_prefetch(p3+idx+128,0,3);__builtin_prefetch(p4+idx+128,0,3);__builtin_prefetch(p5+idx+128,0,3);__builtin_prefetch(p6+idx+128,0,3);__builtin_prefetch(p7+idx+128,0,3);}
__m256i u2=_mm256_add_epi16(ld(p1+idx),ld(p6+idx));
__m256i u3=_mm256_add_epi16(u2,ld(p7+idx));
__m256i v5=ld(p5+idx);
st(temp+v,_mm256_add_epi16(_mm256_add_epi16(u2,v5),ld(p3+idx)));
st(temp+128+v,_mm256_sub_epi16(u3,ld(p4+idx)));
st(temp+256+v,_mm256_add_epi16(u3,v5));
st(temp+384+v,_mm256_add_epi16(ld(p2+idx),ld(p1+idx)));
}
size_t off=(size_t)(block/32)*stride;
emit_native(d12+off,temp,stride);
emit_native(d21+off,temp+128,stride);
emit_native(d22+off,temp+256,stride);
emit_native(d11+off,temp+384,stride);
}
}
}
template<int N>
static __attribute__((noinline)) void mul(const U *a,const U *b,U *c,U *work) {
if constexpr (N<=BASE) {
base(a,b,c,N);
} else if constexpr (N<=ALTMAX) {
constexpr int h=N/2,q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
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+0,*y=work+2*q+704,*next=work+4*q+640;
mul<h>(a12,b21,c11,next); // m2
mul<h>(a22,b22,c22,next); // m4
combine_displaced<false,q>(x,a21,q);
combine_displaced<false,q>(y,b21,q);
mul<h>(x,y,c12,next); // m5
combine_displaced<true,-2*q>(x,a22,q);
combine_displaced<true,-2*q>(y,b22,q);
mul<h>(x,y,c21,next); // m6
for(int chunk=0;chunk<q;chunk+=1024+TILE_PAD)
_Pragma("GCC unroll 4")
for(int i=chunk;i<chunk+1024;i+=16)
st(c22+i,_mm256_sub_epi16(_mm256_add_epi16(ld(c12+i),ld(c21+i)),_mm256_add_epi16(ld(c11+i),ld(c22+i))));
if constexpr(N==64) {
leaf_add(a11,b11,c11); // m1 directly accumulates
combine_displaced<true,-3*q>(y,b22,q);
leaf_add(a21,y,c21); // m3 directly accumulates
combine_displaced<true,3*q>(x,a11,q); // negate the M7 A input
leaf_add(x,b12,c12); // c12+=(-oldM7)
} else {
mul<h>(a11,b11,x,next); // m1
combine_inplace<false>(c11,x,q);
combine_displaced<true,-3*q>(y,b22,q);
mul<h>(a21,y,x,next); // m3
combine_inplace<false>(c21,x,q);
combine_displaced<true,-3*q>(x,a22,q);
mul<h>(x,b12,y,next); // m7
combine_inplace<true>(c12,y,q);
}
} else {
constexpr int h=N/2,q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
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+192,*y=work+q+512,*next=work+q+q+512;
combine_displaced<true,-2*q>(y,b22,q);
combine_displaced<true,2*q>(x,a11,q);
mul<h>(x,y,c21,next); // P7
combine_displaced<false,q>(x,a21,q);
combine_displaced<true,-q>(y,b12,q);
mul<h>(x,y,c22,next); // P5
combine_inplace<true>(x,a11,q);
combine_reverse_inplace(y,b22,q);
mul<h>(x,y,c12,next); // P6
combine_reverse_inplace(x,a12,q);
mul<h>(x,b22,c11,next); // P3
combine_inplace<true>(y,b21,q);
mul<h>(a22,y,x,next); // P4
mul<h>(a11,b11,y,next); // P1
if constexpr (N == 4096) {
combine3_unpack<2048>(final_out+2048,
final_out+(size_t)2048*4096,
final_out+(size_t)2048*4096+2048,
c12,c21,c22,y,x,c11,4096);
} else {
for(int chunk=0;chunk<q;chunk+=1024+TILE_PAD) for(int i=chunk;i<chunk+1024;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<h>(a12,b21,c11,next); // P2
if constexpr (N == 4096) unpack_add<2048>(final_out,c11,y,4096);
else combine_inplace<false>(c11,y,q);
}
}
static constexpr int RQ=2048*2048+(2048/BASE)*(2048/BASE)*TILE_PAD,RS=RQ+ROOT_PAD;
alignas(4096) static U root_mem[15*RS+4096];
alignas(4096) static U root_tile[4*(64*64+(64/BASE)*(64/BASE)*TILE_PAD)];
static constexpr int HQ=64*64+(64/BASE)*(64/BASE)*TILE_PAD;
alignas(4096) static U high_left[4*HQ],high_right[4*HQ];
template<bool Right>
static void pack_four(U *tile,const U *s) {
pack(tile,s,64,4096,Right);
pack(tile+HQ,s+2048,64,4096,Right);
pack(tile+2*HQ,s+2048*4096,64,4096,Right);
pack(tile+3*HQ,s+2048*4096+2048,64,4096,Right);
}
// s4de_: 横に隣接する 2 つの 64x64 tile(源 s と s+64)を 1 つの読み pass で作る。
// Right=false の側だけが広幅読み(A オペランド)。Right=true は原状の 2 呼び出しに委譲。
template<bool Right>
static void pack_wide_root(U *tA,U *tB,const U *s) {
if constexpr(!Right) {
pack_wide(tA+0*HQ, tB+0*HQ, s);
pack_wide(tA+1*HQ, tB+1*HQ, s+2048);
pack_wide(tA+2*HQ, tB+2*HQ, s+2048L*4096);
pack_wide(tA+3*HQ, tB+3*HQ, s+2048L*4096+2048);
} else {
pack_four<Right>(tA,s);
pack_four<Right>(tB,s+64);
}
}
template<bool Right,bool High>
static void emit_root64(const U *tile,const U *left,const U *right,int offset,U *dest,U *destq,int voff) {
U *r0=destq+offset,*r1=r0+RS,*r2=r1+RS,*v7=dest+offset+voff*RS,*v5=v7+RS,*v6=v5+RS,*v34=v6+RS;
for(int i=0;i<(64*64+(64/BASE)*(64/BASE)*TILE_PAD);i+=32){
__m256i a=ld(tile+i),b=ld(tile+(64*64+(64/BASE)*(64/BASE)*TILE_PAD)+i),c=ld(tile+2*(64*64+(64/BASE)*(64/BASE)*TILE_PAD)+i),d=ld(tile+3*(64*64+(64/BASE)*(64/BASE)*TILE_PAD)+i);
__m256i aa=ld(tile+i+16),bb=ld(tile+(64*64+(64/BASE)*(64/BASE)*TILE_PAD)+i+16),cc=ld(tile+2*(64*64+(64/BASE)*(64/BASE)*TILE_PAD)+i+16),dd=ld(tile+3*(64*64+(64/BASE)*(64/BASE)*TILE_PAD)+i+16);
if constexpr(High) {
a=_mm256_add_epi16(a,_mm256_sub_epi16(ld(left+0*HQ+i),ld(right+0*HQ+i)));
b=_mm256_add_epi16(b,_mm256_sub_epi16(ld(left+1*HQ+i),ld(right+1*HQ+i)));
c=_mm256_add_epi16(c,_mm256_sub_epi16(ld(left+2*HQ+i),ld(right+2*HQ+i)));
d=_mm256_add_epi16(d,_mm256_sub_epi16(ld(left+3*HQ+i),ld(right+3*HQ+i)));
aa=_mm256_add_epi16(aa,_mm256_sub_epi16(ld(left+0*HQ+i+16),ld(right+0*HQ+i+16)));
bb=_mm256_add_epi16(bb,_mm256_sub_epi16(ld(left+1*HQ+i+16),ld(right+1*HQ+i+16)));
cc=_mm256_add_epi16(cc,_mm256_sub_epi16(ld(left+2*HQ+i+16),ld(right+2*HQ+i+16)));
dd=_mm256_add_epi16(dd,_mm256_sub_epi16(ld(left+3*HQ+i+16),ld(right+3*HQ+i+16)));
}
_mm256_stream_si256((__m256i*)(r0+i),a);_mm256_stream_si256((__m256i*)(r0+i+16),aa);
_mm256_stream_si256((__m256i*)(r1+i),Right?c:b);_mm256_stream_si256((__m256i*)(r1+i+16),Right?cc:bb);
_mm256_stream_si256((__m256i*)(r2+i),d);_mm256_stream_si256((__m256i*)(r2+i+16),dd);
__m256i x7=Right?_mm256_sub_epi16(d,b):_mm256_sub_epi16(a,c);
__m256i xx7=Right?_mm256_sub_epi16(dd,bb):_mm256_sub_epi16(aa,cc);
_mm256_stream_si256((__m256i*)(v7+i),x7);_mm256_stream_si256((__m256i*)(v7+i+16),xx7);
__m256i x5=Right?_mm256_sub_epi16(b,a):_mm256_add_epi16(c,d);
__m256i xx5=Right?_mm256_sub_epi16(bb,aa):_mm256_add_epi16(cc,dd);
_mm256_stream_si256((__m256i*)(v5+i),x5);_mm256_stream_si256((__m256i*)(v5+i+16),xx5);
__m256i x6=Right?_mm256_sub_epi16(d,x5):_mm256_sub_epi16(x5,a);
__m256i xx6=Right?_mm256_sub_epi16(dd,xx5):_mm256_sub_epi16(xx5,aa);
_mm256_stream_si256((__m256i*)(v6+i),x6);_mm256_stream_si256((__m256i*)(v6+i+16),xx6);
__m256i x34=Right?_mm256_sub_epi16(x6,c):_mm256_sub_epi16(b,x6);
__m256i xx34=Right?_mm256_sub_epi16(xx6,cc):_mm256_sub_epi16(bb,xx6);
_mm256_stream_si256((__m256i*)(v34+i),x34);_mm256_stream_si256((__m256i*)(v34+i+16),xx34);
}
}
template<bool Right>
static void pack_root(const U*s,int n,int offset,U *dest,U *destq,int voff) {
if(n>128){int h=n/2,q=h*h+(h/BASE)*(h/BASE)*TILE_PAD;
pack_root<Right>(s,h,offset,dest,destq,voff);
pack_root<Right>(s+h,h,offset+q,dest,destq,voff);
pack_root<Right>(s+h*4096,h,offset+2*q,dest,destq,voff);
pack_root<Right>(s+h*4096+h,h,offset+3*q,dest,destq,voff);
} else {
pack_wide_root<Right>(root_tile,high_left,s);
emit_root64<Right,false>(root_tile,nullptr,nullptr,offset,dest,destq,voff);
emit_root64<Right,false>(high_left,nullptr,nullptr,offset+HQ,dest,destq,voff);
pack_wide_root<Right>(high_right,root_tile,s+64*4096);
emit_root64<Right,false>(high_right,nullptr,nullptr,offset+2*HQ,dest,destq,voff);
emit_root64<Right,true>(root_tile,high_left,high_right,offset+3*HQ,dest,destq,voff);
}
}
void matrix_multiply(int n,const short*A,const short*B,short*C){
U *cq=(U*)C; // Three root input slots and the retired tail precede final C writeback.
U *ar=root_mem,*br=ar+4*RS,*pr=br+7*RS; // root_mem 15 -> 12 槽
U *work=cq+3*RS+(464*2);
raw_basis_input=(const U*)A;
pack_root<false>((const U*)A,2048,0,ar,cq,0);
raw_basis_input=(const U*)B;
pack_root<true>((const U*)B,2048,0,br,br,3);
_mm_mfence();
const U *a11=cq,*a12=cq+RS,*a22=cq+2*RS,*s7=ar,*s5=ar+RS,*s6=ar+2*RS,*s3=ar+3*RS;
const U *b11=br,*b21=br+RS,*b22=br+2*RS,*t7=br+3*RS,*t5=br+4*RS,*t6=br+5*RS,*t4=br+6*RS;
U *p1=ar+2*RS,*p2=br+5*RS,*p3=ar+RS,*p4=br+4*RS,*p5=ar,*p6=br+3*RS,*p7=pr;
mul<2048>(s7,t7,p7,work);mul<2048>(s5,t5,p5,work);mul<2048>(s6,t6,p6,work);
mul<2048>(s3,b22,p3,work);mul<2048>(a22,t4,p4,work);mul<2048>(a11,b11,p1,work);mul<2048>(a12,b21,p2,work);
U *out=(U*)C;
root_fuse<2048>(out,out+2048,out+2048*4096,out+2048*4096+2048,p1,p2,p3,p4,p5,p6,p7,4096);_mm_sfence();
asm volatile(".purgem vzeroupper");
_mm256_zeroupper();
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 529.55 ms | 129 MB + 648 KB | Accepted | Score: 100 | 显示更多 |