// ===== REFERENCES =====
// [1] duck.ac 用户 saffah_codex_6s_agg2,提交 #103930 <https://duck.ac/submission/103930>
// 用途:**本文件正文的基底**(经本账号 #104136 逐字节复刻后,又叠了 #104673/#104724 两刀)。
// [2] 本账号 saffah_cc_v41_agg1,提交 #104724 <https://duck.ac/submission/104724>
// 用途:本发的**直接基座**(`problems/1004e7/work/e7x_sub10.cpp`,102.720577 ms)。
// [3] duck.ac 用户 saffah_codex_6s_agg2,提交 #103812 <https://duck.ac/submission/103812>
// 用途:`E7R_CH`、DRAM 表预取距离、CRT stage-1 32 宽四链(经 [1] 传入本文件)。
// [4] BRIEF §2.18.825(`cyc/元素 ≈ max(loads/2, stores/1, uops/4)`,max 形式)/ §2.18.824②(opcode 直方图)。
// ======================
// ===== 思路 =====
// 正式提交(非试验性)。**单变量**:`parse_limbs6` 的内层循环不再把 12 个"4 位数字组"落回内存。
//
// 【病灶】原循环把 8 个组用一条 **32 字节 `STU`**、另 4 个用一条 **16 字节 `_mm_storeu_si128`**
// 写进局部 `u32 G[12]`,紧接着**用 10 条 4 字节 GPR 载入**读回来。**32 字节存 → 窄载入
// 无法走 store-to-load forwarding**(Intel 的转发只在"载入被单条存数完整包含且满足其对齐/宽度
// 条件"时成立)⇒ 这 10 条载入要等那条 32 字节存**落地到 L1**,约 12~20 拍。
// **旁证(拍数账,BRIEF §2.18.825 的 max 形式)**:本循环每次迭代 4 个 limb、约 **67 条 uop**,
// 而实测 **≈40 拍/迭代(= 10 拍/limb = 20.1 Mcyc / 2e6 limb)** ⇒ **1.7 uop/cycle**,
// 而它的前端地板只有 `max(loads/2, stores/1, uops/4) ≈ 4.2 拍/limb`。
// **1.7 uop/cycle 只能由一条跨迭代的串行停顿解释,而这条转发停顿正是唯一的候选。**
//
// 【改法】组值留在寄存器里,用 `_mm256_extract_epi32` / `_mm_extract_epi32`(编译期下标 ⇒ `vpextrd`)
// 取出来。**没有存、没有载、没有转发问题**;`l0..l3` 的拼装与下游 `red4` 一字未动。
// 【代价】每迭代 +约 16 条 uop(`vpextrd` 2 uop;取高 128 位再 +1 条 `vextracti128`)。
// 按本引擎 DRAM 层标定的 0.156 cyc/uop(见 `notes.md` ⑬)⇒ +2.5 拍/迭代;
// 而若停顿如估计的那样是 ~15 拍,**净省 ~12 拍/迭代 = 6 Mcyc ≈ 1.7%**。
//
// 【正确性】`_mm256_extract_epi32(v,k)` 对 k<4 取低 128 位的第 k 个 32 位元素、k≥4 取高 128 位,
// **与原 `G[0..7] = 该 32 字节存的内容` 逐位相同**;`G[10],G[11]` ↔ `Ghv[2],Ghv[3]` 同理。
// 本地闸门:`work/e7c_mkver.py` 输出必须 **20000001 字节、md5 `4c202c31b19f27f2525bb5a1ba02b9ed`**
// (本发已实测通过)。另按 RULES §1.0⑤ dump 判题编译器 gcc 9.3 汇编按 opcode 直方图核对(禁 `div` 回归)。
// 【可证伪的判据】判题机用时 < 102.720577 ms 即为正;< 102.161545 ms 则直接达标。
// ================
// ===== REFERENCES =====
// [1] duck.ac 用户 saffah_codex_6s_agg2,提交 #103930 <https://duck.ac/submission/103930>
// 用途:**本文件正文的基底**(正文 = 该提交,除下方 [A][C] 两处外逐字节相同)。
// [2] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #104673 <https://duck.ac/submission/104673>
// 用途:本账号当前最好件 Accepted **102.924516 ms**(= #103930 + [A])。
// [3] duck.ac 用户 saffah_codex_6s_agg2,提交 #103994 <https://duck.ac/submission/103994>
// 用途:#103994 时代把 parse 输出改非临时存的形态(本发 [C] 是同一思想在 P2 之后的复活)。
// [4] duck.ac 用户 saffah_codex_6s_agg2,提交 #103812 <https://duck.ac/submission/103812>
// 用途:`_mulx_u64` 写 `M2*u3+c12`、`E7R_CH`、DRAM 表预取、CRT stage-1 32 宽四链。
// [5] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #103645 <https://duck.ac/submission/103645>
// 用途:三素数 NTT + Shoup 乘 + CRT Garner 32 位半边 `vmulcv()` 向量化的引擎本体。
// ======================
// ===== 思路 =====
// [A](#104673 已交):`crt` stage-2 标量循环 **4 倍展开**(`CRTSTEP` 宏,进位链语义未改)。
// [C] ★ **本发主刀:把 `parse_limbs6` 的输出写改成非临时存(NT)**。
// 取证(免费,`grep _mm_stream`):本文件里 `gen_geo` / `build_tabs` 的大流**早就是 NT 写**,
// 但 `parse_limbs6` 的三路输出(`RA/RB/RC` 与 `PB_[0..2]`,**每路 8 MB ⇒ 共 24 MB**)**仍是
// 普通 `_mm_storeu_si128`** —— 也就是说这 24 MB 每次都要先付一次 **RFO(read-for-ownership)
// 读**。历史证据:对手在 #99494 代**做过**"parse 输出改 NT",但 P2(本账号 #102724 把解析直写
// NTT 数组)重写这一段时把 NT 弄丢了 ⇒ **这笔钱被重新捡回来**。
// **实现要点(否则 `movntdq` 会直接 #GP)**:NT 存(`_mm_stream_si128`)要求 **16 字节对齐**,
// 而该循环的写地址是 `oX + k - 4`、`k` 每次减 4 —— 只有当 `(nlm-1) ≡ 0 (mod 4)` 时才对齐。
// 本发在向量循环**之前**加一个**最多 3 次的标量引流**(用该函数本来就有的 `ST3(swar8(...)*100+ld2d(...))`
// 标量体)把 `k` 推到 `(k-4) ≡ 0 (mod 4)` ⇒ 之后每次 NT 存都对齐 ✓(输入太短时引流条件不满足、
// 向量循环也不会进入,天然安全)。
// 【定价】判题机 probe(`mkinstr.py` 分段仪,同一二进制,4 run/发,取 run≥1):
// | 件 | ntt Mcyc | 总 Mcyc | Δ总 |
// |---|---|---|---|
// | 基座 | 305.1 | 356.0 | — |
// | [A] CRT 4× | 305.1 | 352.1 | −1.09%(**判题机只兑现 0.069%**)|
// | **[C] parse NT(单变量)** | **303.9** | **354.8** | **−0.34%** |
// | [C] 的邻点:把 CRT 的 `res` 写也改 NT | 305.1 | 356.2 | −0.06%(负)⇒ 放弃 |
// ⇒ 本发 = [A]+[C],probe 预期 ≈ −1.43%;按本 lane 实测的**分类兑现率**("访存/流量"类 ≈1:1、
// "纯循环开销"类 ≈0.06×)折算的**判题机预期 ≈ −0.41%**。
// 【正确性】probe `chk = 972d2d6d60c0558c` 与基座逐字节相同(NT 存与普通存对同线程可见性无差别,
// 且函数内本来就有 `nt_fence()`);本地 oracle 输出 20000001 字节、md5 `4c202c31b19f27f2525bb5a1ba02b9ed`;
// `tools/gcc93check.py` exit=0。
// ================
// ===== REFERENCES =====
// [1] duck.ac 用户 saffah_codex_6s_agg2,提交 #103930 <https://duck.ac/submission/103930>
// 用途:**直接复制全文作为正文基底** —— 本文件正文与 #103930 逐字节相同。唯一的机械改动是把
// 他自带的 `// ===== REFERENCES =====` / `// ===== 思路 =====` 两行标记降级成普通
// 注释行(避免出现第二个引用块),其文字内容一字未动。
// [2] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #103868 <https://duck.ac/submission/103868>
// 用途:#103930 的正文父本(本账号上一发最好件 103.718363 ms)。
// 取证(tools/rivalcopy.py):fwd 99.41% / rev 99.88% vs #103868 ⇒ 他是拿 #103868 当基座。
// [3] duck.ac 用户 saffah_codex_6s_agg2,提交 #103812 <https://duck.ac/submission/103812>
// 用途:#103868 的正文父本(`_mulx_u64` 写 `M2*u3+c12`、`E7R_CH`=48、DRAM 表预取 256->768、
// CRT stage-1 32 宽四链);这些形态在本文件里继续存在。
// [4] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #103645 <https://duck.ac/submission/103645>
// 用途:三素数 NTT + Shoup 乘 + CRT Garner 32 位半边 `vmulcv()` 向量化的引擎本体。
// #103930 相对 #103868 的**全部**改动(tools/fnmd5.py:SAME=52 DIFF=3 ONLY_A=0 ONLY_B=0,
// 正文 diff 42 行全部落在 `stage4_dif_t` / `merge4_dit` / `merge4_dit_t` 三个 DRAM 层内核里,
// 且**全部是 `__builtin_prefetch` 的距离常量**,无一处算术/布局改动):
// (a) 第 0 行的数据预取 `p + j + 256` -> `p + j + 768`(5 处;另 1 处 `+264` -> `+776`);
// ★ 注意第 1/2/3 行(`p + j + h + 256` 等)保持 256 **未动**;
// (b) `stage4_dif_t` 的三张表预取 `t1/t2/t3 + j + 768` -> `+ j + 512`(4 处);
// (c) `merge4_dit` / `merge4_dit_t` 的三张表预取 `t1/t2/t3 + j + 256` -> `+ j + 768`(4 处)。
// ======================
// ===== 思路 =====
// 本发是**纯复刻**(BRIEF §2.16 的第一步)。入口复核(`tools/exact.py 1004e7`,判题机原值口径):
// mine = 103.718363 (#103868) / T = 103.192470 (#103930) ⇒ mine > T,
// 按 §2.18.566② 的判据"mine > T ⇒ 整份复刻有效",先把这 0.51% 的结构性缺口一次压掉。
// 目的(试验性,必须在提交说明里写明):1) 把对手 #103930 的成绩在本账号上重测一次,作为下一刀
// A/B 的判题机原值基线;2) 验证"DRAM 层预取距离"这条轴在他那个取值上是否真的值 0.51%
// (§2.18.725:预取类旋钮**只能**用正式提交定价,"立即体验"对它系统性误价)。
// ★ 本发**不含**本账号自己的任何新刀 —— 与上一次我方的 #103868 相比,唯一的变量就是 (a)(b)(c)
// 三组预取距离常量。
// ================
/*
References:
- duck.ac user saffah_codex_6s_agg2, https://duck.ac/submission/103833: copied our tuned NTT/Garner baseline; no license notice was shown.
- duck.ac user saffah_cc_v41_agg1, https://duck.ac/submission/103868: adapted its exact one-division CRT identity; no license notice was shown.
Approach:
- Retain the measured 768-word data prefetch and tune DRAM twiddle-table prefetch to 512 words, keeping all arithmetic and output unchanged.
Purpose:
- Measure whether jointly tuned data/table lookahead improves 1004e7 enough for the dynamic exact threshold.
*/
/*
References:
- duck.ac user saffah_codex_6s_agg2, https://duck.ac/submission/103833: copied our NTT and vector Garner baseline; no license notice was shown.
- duck.ac user saffah_cc_v41_agg1, https://duck.ac/submission/103868: adapted its exact one-division CRT identity; no license notice was shown.
Approach:
- Keep the exact one-division reconstruction and vary only the DRAM NTT data-prefetch distance to 768 words.
Purpose:
- Test whether a longer NTT memory lookahead improves the official 1004e7 runtime.
*/
/*
References:
- duck.ac user saffah_codex_6s_agg2, https://duck.ac/submission/103833: directly reused our four-chain NTT, tuned prefetch, and vector Garner implementation; no license notice was shown.
- duck.ac user saffah_cc_v41_agg1, https://duck.ac/submission/103868: directly adapted its exact CRT digit-extraction identity and one-division CRTEX macro. Its attribution and no-license context remain below.
Approach:
- Combine the cited one-division CRT carry formula with the faster twiddle and data prefetch settings from our prior version; keep the arithmetic and output bit-exact.
Purpose:
- Establish a faster baseline before testing a carry-independent two-pass CRT on the official judge CPU.
*/
/*
References:
- duck.ac user saffah_codex_6s_agg2, https://duck.ac/submission/103812: previous MULX CRT implementation used as this direct baseline; no license notice was shown.
- duck.ac user saffah_cc_v41_agg1, https://duck.ac/submission/103645: copied the complete three-prime NTT and vector Garner CRT engine. Its inherited attribution remains below; no license notice was shown.
- duck.ac user saffah_codex_6s_agg2, https://duck.ac/submission/103389: the constant-multiply Garner step already incorporated in that public source. No license notice was shown.
- duck.ac user saffah_cc_v41_agg1, https://duck.ac/submission/102724: the base radix-four NTT and table generation retained through the first source. No license notice was shown.
Approach:
- Retain the measured NTT and MULX CRT reconstruction; adjust DRAM radix-four twiddle-table prefetch to 752 words and data prefetch to 288 words.
Purpose:
- Measure whether the 288-word data prefetch clears the remaining 33,310 ns to the exact leaderboard target.
*/
// ----- 思路 -----
// **把 CRT 循环的 32 位 Garner 半边向量化**(本发是"换角度"后的第一把:不动块循环)。
// 依据(全部是本 lane 自己的判题机 probe 读数,见 notes.md 第 ⑪ 节):
// * `crt` 段 = 68.3 Mcyc = **17.1%**,gcc 生成 **~90 条/limb**,而算术下界只有 ~46
// ⇒ **全场唯一还有 2 倍差距的地方**;
// * ★ **本 lane 已实测:CRT / DRAM 段的 codegen 扰动"免费"** —— 单变量件 `e7r_v3`
// (只删 `merge4_dit_t<H>` 里 3 条重复 prefetch)判题机 probe 读到 **−0.0005%**(=0,
// 同一 harness 基线重复性 0.004%);
// * 而**块循环内核**对任何增删收费 ≈ **+0.25%**(`e7r_v1` +0.22%、`e7r_v2` +0.30%)
// ⇒ 这就是为什么本发**一个字节都不碰块循环**。
// 改动实质:`CRTRECON` 的前 6 步(`dd2`、`u2`、`dr31`、`q31`、`q32`、`u3`)**全部是 32 位模算术**,
// 8 条 limb 正好装进一个 ymm:
// * `(x-y) mod P`(x,y < P)= `t = vpsubd(x,y); vpminud(t, vpaddd(t, P))` —— 2 条指令/8 limb;
// * Shoup 常量乘 `u = a*w - ((a*ws)>>32)*P`:三个常量 `CS2/CS3/DS3 = (x<<32)/P`(x < P)
// **本身 < 2^32** ⇒ `((u64)a*ws)>>32` 恰好是 32x32 乘积的高半边,即**已有的 `vmulhi()`**;
// 新助手 `vmulcv()` = 13 uop / 8 limb(标量同一步约 8 条指令 / limb)。
// `M2*u3 + c12` 的 128 位部分与 `CRTEX` 的进位链**保持标量一字不改**(进位链本质串行)。
// 两半通过一块 **L1 常驻小 scratch(每 chunk 1024 limb = 2×4 KB)**连接,**分块填充/排空**,
// 因此**不产生任何新的 DRAM 流量**(若整数组往返会多 32 MB ≈ +2.4%,故必须分块)。
// 语义等价性:逐 limb 的中间量(`dd2`,`u2`,`dr31`,`q31`,`q32`,`u3`)在两种写法下**逐位相同**
// (向量式只是把同一组 32 位运算并行 8 份),`c12/t/cLo/cHi` 与 `CRTEX` 完全未动。
// 正确性:本文件 notes 的回归闸门 —— 本地 oracle 输出 20000001 字节、md5
// `4c202c31b19f27f2525bb5a1ba02b9ed`,与台账钉死的值**逐字节相同**。
// ================
// ----- REFERENCES -----
// [0] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #103471 <https://duck.ac/submission/103471>
// 用途:**本文件正文即该提交的源码**(112.982786 ms,本账号在 1004e7 的当前最好件);
// 本发只替换其中 CRT 循环的 32 位 Garner 半边,其余(NTT 引擎、parse、out)一字未改。
// [1] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #103440 <https://duck.ac/submission/103440>
// 用途:`mulconst2`/`mulconst3`(Garner 第 2、3 个数字的 Shoup 常量乘形式)即本发向量化的对象;
// 向量版 `vmulcv()` 与它们逐位等价(见"思路"段)。
// [2] duck.ac 用户 saffah_codex_6s_agg2,提交 #103389 <https://duck.ac/submission/103389>
// 用途:`mulconst3` 这一个 Shoup 化的原始作者(经 `tools/rivalcopy.py` fwd 98.75% / rev 99.50%
// 取证:他是我们 #102724 的增量)。本文件正文中更下方的原始引用段继续有效。
// [3] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #102724 <https://duck.ac/submission/102724>
// 用途:解析直写 NTT 数组(P2/P2b)—— 本文件 stage0 原地 DIF 与 `PB_[3]` 的来源。
// [4] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #101838 <https://duck.ac/submission/101838>
// 用途:三素数 NTT 主引擎(blocked radix-4 DIF/DIT mod 2013265921 + Shoup 乘法)与
// `vshoup`/`vmulhi` 等 AVX2 助手(本发的 `vmulcv` 直接复用 `vmulhi`)。
// [5] duck.ac 用户 saffah_cc_v41_260924,提交 #96594 <https://duck.ac/submission/96594>
// 用途:**直接复制**了其全部代码作为更早的基线(见文件正文中的原始引用段)。
// ======================
// [0] duck.ac 用户 saffah_codex_6s_agg2,提交 #103389 <https://duck.ac/submission/103389>
// 用途:**本文件正文逐字节取自该提交**(对手在 1004e7 的 rank-1,113.140360 ms)。
// 经 `tools/rivalcopy.py`(fwd 98.75% / rev 99.50% vs 我们 #102724)+ `tools/fnmd5.py` 取证:
// 它是我们 #102724 的公开源码 + **15 行改动,且全部落在函数体之外(`★ MISSED` 区)**
// —— 新增 `mulconst3()` 助手 + 两个常量 + 换掉第三个 Garner 数字的算法。
// [1] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #102724 <https://duck.ac/submission/102724>
// 用途:我在 1004e7 的上一代最好件(114.780259 ms),也是 [0] 的基座;其自身血统
// (#102616 Montgomery 预标定、三素数 NTT 基线)见其头部,继续有效。
// ======================
// (历史引用/思路说明段,非本发的合规块) ===== 思路(单变量:把对手那把刀再往前推一格)=====
// 对手 #103389 把**第 3 个 Garner 数字**里的 Montgomery 归约换成两次 Shoup 常量乘(顺带消掉 128 位的 `c12 % P3`)。
// 本发对**第 2 个 Garner 数字**做同一件事:它原本也在做 Montgomery 归约
// x2 = dd2*Kinv2; m2 = (u32)x2*pinv2; u2 = (x2 + m2*P2)>>32; if (u2>=P2) u2-=P2; (3 乘 + 1 加 + 1 移位 + 1 条件减)
// 而它的数学内容只是 **u2 = (r2-r1) * inv(p1) mod P2** 这一个常量乘 ⇒ 直接写成 Shoup 形式
// u32 u2 = mulconst2(dd2, i12, CS2); (2 乘 + 1 减 + 1 条件减)
// 常量 `i12 = inv(p1) mod P2` 原文件已有(原本喂给 Kinv2),新增的只有 `CS2 = (i12<<32)/P2`;
// 原来的 `Kinv2`/`pinv2` 随之变成死代码被删掉(**净删除**,`§2.18.691` 偏好的形态)。
// ⇒ 每 limb 少 1 次乘、1 次加、1 次移位,并少一条 Montgomery 依赖链。
// 【定价】判题机 `tools/probe.py`(同进程 3 次、取后两次,前一次为预热臂):
// 基座 r1=388545k r2=388580k;本发 r1=387286k r2=387363k ⇒ **−1.26M ticks = −0.32%**。
// 【正确性】本地回归闸门逐字节:输出 20000001 字节、md5 `4c202c31b19f27f2525bb5a1ba02b9ed`,
// 与台账钉死的值**逐字节相同**。
// ================
#pragma GCC optimize("O3","unroll-loops","inline-functions","ipa-cp-clone","ipa-sra","rename-registers")
#pragma GCC target("arch=skylake")
// Blocked radix-4 DIF/DIT NTT mod 2013265921 using Shoup modular multiplication.
// Per stage: 6 tables (t1,t2,t3 = twiddle values; s1,s2,s3 = Shoup constants), contiguous.
// DRAM levels: block sizes n, n/4, ...; then in-cache base blocks of size ~2^15..2^19.
// DIF forward: natural in -> digit-reversed out. DIT inverse: digit-reversed in -> natural out.
#pragma once
#include <immintrin.h>
#include <cstdint>
#include <cstring>
#include <cstdlib>
#include <cstdint>
// ---- forced single-instruction 256-bit unaligned access (gcc 9.3 -O2 downgrades the intrinsic
// to vmovdqu xmm + vinserti128); AVX2 is already enabled file-wide by #pragma GCC target above.
static inline __m256i LDU(const void *p){ __m256i v; __asm__("vmovdqu %1,%0" : "=x"(v) : "m"(*(const __m256i *)p)); return v; }
__attribute__((target("avx2"))) static inline void STU(void *p, __m256i v){ __asm__("vmovdqu %1, %0" : "=m"(*(__m256i *)p) : "x"(v)); }
typedef uint32_t u32;
typedef uint16_t u16;
typedef uint64_t u64;
// non-temporal 32B store + fence + the "big & aligned enough" guard (adopted from REFERENCES [3])
static inline void stu256_nt(void *p, __m256i v){ __asm__("vmovntdq %1, %0" : "=m"(*(__m256i *)p) : "x"(v)); }
static inline void nt_fence(void){ __asm__ __volatile__("sfence" ::: "memory"); }
#ifndef NT_MIN
#define NT_MIN 4096u
#endif
#define NT_OK(P, CNT) ((u32)(CNT) >= NT_MIN && (((uintptr_t)(P) & 31u) == 0u))
static inline u32 *pad8(u32 *p){ return (u32 *)(((uintptr_t)p + 31u) & ~(uintptr_t)31u); }
#define NTP 2013265921u
#define TGT __attribute__((target("avx2")))
#ifndef NTT_MAXN
#define NTT_MAXN (1u << 21)
#endif
static u32 g_n, g_Sb;
#define E7R_CH 48u /* CRT stage-1 chunk: 2 * 4 KB of L1-resident scratch */
static u32 g_P = 2013265921u, g_Ml = 572662301u, g_Msh = 1;
static u32 g_pinv;
static int g_log2n, g_L, g_odd, g_nb, g_BT = 15;
static u32 g_ciF, g_cisF, g_ciI, g_cisI;
static u32 g_tabF[5 * NTT_MAXN / 2 + 4096] __attribute__((aligned(64)));
static u32 g_tabI[5 * NTT_MAXN / 2 + 4096] __attribute__((aligned(64)));
struct Tabs {
const u32 *l1[40], *l2[40], *l3[40], *m1[40], *m2[40], *m3[40];
const u32 *b2, *b2s;
const u32 *h1[40], *h2[40], *h3[40], *q1[40], *q2[40], *q3[40];
};
static Tabs g_TF, g_TI;
static inline u32 mulmod(u32 a, u32 b) { return (u32)((u64)a * b % g_P); }
static u32 powmod(u32 a, u32 e) { u32 r = 1; while (e) { if (e & 1) r = mulmod(r, a); a = mulmod(a, a); e >>= 1; } return r; }
static inline u32 shoup1s(u32 a, u32 w, u32 ws) {
u32 t = (u32)((u64)a * w);
u32 q = (u32)(((u64)a * ws) >> 32);
u32 r = t - (u32)((u64)q * g_P);
return r >= g_P ? r - g_P : r;
}
static inline u32 shoup_c(u32 w) { return (u32)(((u64)w << 32) / g_P); }
TGT static inline __m256i vmulhi(__m256i a, __m256i b) {
__m256i e = _mm256_mul_epu32(a, b);
__m256i o = _mm256_mul_epu32(_mm256_srli_epi64(a, 32), _mm256_srli_epi64(b, 32));
e = _mm256_srli_epi64(e, 32);
return _mm256_blend_epi32(e, o, 0xAA);
}
TGT static inline __m256i vmont(__m256i a, __m256i b, __m256i pv) {
const __m256i P64 = _mm256_set1_epi64x((long long)g_P);
const __m256i NP = _mm256_set1_epi64x((long long)g_pinv);
__m256i ahi = _mm256_srli_epi64(a, 32), bhi = _mm256_srli_epi64(b, 32);
__m256i tlo = _mm256_mul_epu32(a, b), thi = _mm256_mul_epu32(ahi, bhi);
__m256i mlo = _mm256_mul_epu32(tlo, NP), mhi = _mm256_mul_epu32(thi, NP);
__m256i plo = _mm256_mul_epu32(mlo, P64), phi = _mm256_mul_epu32(mhi, P64);
__m256i ulo = _mm256_srli_epi64(_mm256_add_epi64(tlo, plo), 32);
__m256i uhi = _mm256_srli_epi64(_mm256_add_epi64(thi, phi), 32);
__m256i u = _mm256_blend_epi32(ulo, _mm256_slli_epi64(uhi, 32), 0xAA);
return _mm256_min_epu32(u, _mm256_sub_epi32(u, pv));
}
// M = floor(2^64/p) = 2^33 + MLO ; ws ~= 2*w + mulhi(w, MLO) (error <= 1, valid for Shoup)
#define MLO_V 572662301u
TGT static inline __m256i der_ws(__m256i w) {
__m256i mlo = _mm256_set1_epi32((int)g_Ml);
__m256i t = _mm256_add_epi32(vmulhi(w, mlo), _mm256_sll_epi32(w, _mm_cvtsi32_si128(g_Msh)));
return t;
}
TGT static inline __m256i vshoup(__m256i a, __m256i w, __m256i ws, __m256i pv) {
__m256i t = _mm256_mullo_epi32(a, w);
__m256i q = vmulhi(a, ws);
__m256i r = _mm256_sub_epi32(t, _mm256_mullo_epi32(q, pv));
return _mm256_min_epu32(r, _mm256_sub_epi32(r, pv));
}
TGT static inline __m256i vaddm(__m256i a, __m256i b, __m256i pv) {
__m256i s = _mm256_add_epi32(a, b);
return _mm256_min_epu32(s, _mm256_sub_epi32(s, pv));
}
// 32-bit Shoup constant multiply, 8 lanes: u = a*W - ((a*WS)>>32)*P (mod 2^32),
// then reduced into [0,P). Used only by the CRT Garner steps. WS is the low
// 32 bits of the u64 Shoup constant (x<<32)/P with x < P; because x < P that
// constant is itself < 2^32, so "((u64)a*ws)>>32" is exactly the high half of
// the 32x32 product -- i.e. vmulhi(a, (u32)ws).
TGT static inline __m256i vmulcv(__m256i a, __m256i W, __m256i WS, __m256i Pv) {
__m256i t = _mm256_mullo_epi32(a, W);
__m256i q = vmulhi(a, WS);
__m256i r = _mm256_sub_epi32(t, _mm256_mullo_epi32(q, Pv));
return _mm256_min_epu32(r, _mm256_sub_epi32(r, Pv));
}
TGT static inline __m256i vsubm(__m256i a, __m256i b, __m256i pv) {
__m256i d = _mm256_sub_epi32(a, b);
return _mm256_min_epu32(d, _mm256_add_epi32(d, pv));
}
TGT static inline __m128i vmulhi4(__m128i a, __m128i b) {
__m128i e = _mm_mul_epu32(a, b);
__m128i o = _mm_mul_epu32(_mm_srli_epi64(a, 32), _mm_srli_epi64(b, 32));
e = _mm_srli_epi64(e, 32);
return _mm_blend_epi32(e, o, 0xAA);
}
TGT static inline __m128i sh4(__m128i a, __m128i w, __m128i ws, __m128i pv) {
__m128i t = _mm_mullo_epi32(a, w);
__m128i q = vmulhi4(a, ws);
__m128i r = _mm_sub_epi32(t, _mm_mullo_epi32(q, pv));
return _mm_min_epu32(r, _mm_sub_epi32(r, pv));
}
TGT static inline __m128i add4(__m128i a, __m128i b, __m128i pv) {
__m128i s = _mm_add_epi32(a, b);
return _mm_min_epu32(s, _mm_sub_epi32(s, pv));
}
TGT static inline __m128i sub4(__m128i a, __m128i b, __m128i pv) {
__m128i d = _mm_sub_epi32(a, b);
return _mm_min_epu32(d, _mm_add_epi32(d, pv));
}
TGT static inline __m256i ld2(const __m128i *lo, const __m128i *hi) {
return _mm256_inserti128_si256(_mm256_castsi128_si256(_mm_loadu_si128(lo)), _mm_loadu_si128(hi), 1);
}
TGT static inline void st2(__m128i *lo, __m128i *hi, __m256i v) {
_mm_storeu_si128(lo, _mm256_castsi256_si128(v));
_mm_storeu_si128(hi, _mm256_extracti128_si256(v, 1));
}
TGT static inline __m256i bc128(__m128i v) {
return _mm256_inserti128_si256(_mm256_castsi128_si256(v), v, 1);
}
// fused level-0 DIF stage: reads coefficients directly from src ([0,len] valid, zero beyond).
// Requires len >= h and len < 2h+... (x2 = x3 = 0 assumed).
TGT static void stage0_dif_src(u32 *a, const u32 *src, u32 len, u32 h,
const u32 *t1, const u32 *t2, const u32 *t3,
u32 ci, u32 cis, u32 prescale) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
const __m256i zero = _mm256_setzero_si256();
const __m256i scv = _mm256_set1_epi32((int)prescale);
const __m256i scs = der_ws(scv);
const u32 sc_scalar = shoup_c(prescale);
u32 j = 0;
u32 lim1 = len - h;
for (; j + 8 <= lim1 + 1; j += 8) {
__builtin_prefetch(src + j + 256);
__builtin_prefetch(src + j + h + 256);
__builtin_prefetch(a + j + 256);
__builtin_prefetch(a + j + h + 256);
__builtin_prefetch(a + j + 2 * h + 256);
__builtin_prefetch(a + j + 3 * h + 256);
__builtin_prefetch(t1 + j + 256);
__m256i x0 = LDU(src + j);
__m256i x1 = LDU(src + j + h);
if (prescale != 1u) {
x0 = vshoup(x0, scv, scs, pv);
x1 = vshoup(x1, scv, scs, pv);
}
__m256i u0 = vaddm(x0, x1, pv), u2 = vsubm(x0, x1, pv);
__m256i E = vshoup(x1, cI, cIs, pv);
__m256i u1 = vaddm(x0, E, pv), u3 = vsubm(x0, E, pv);
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, der_ws(w), pv); }
STU(a + j, u0);
STU(a + j + h, u1);
STU(a + j + 2 * h, u2);
STU(a + j + 3 * h, u3);
}
for (; j <= lim1 && j < h; j++) {
u32 x0 = (j <= len) ? src[j] : 0, x1 = (j + h <= len) ? src[j + h] : 0;
if (prescale != 1u) {
x0 = shoup1s(x0, prescale, sc_scalar);
x1 = shoup1s(x1, prescale, sc_scalar);
}
u32 u0 = (u32)(((u64)x0 + x1) % g_P);
u32 u2 = (u32)(((u64)x0 + g_P - x1) % g_P);
u32 E = mulmod(x1, ci);
u32 u1 = (u32)(((u64)x0 + E) % g_P), u3 = (u32)(((u64)x0 + g_P - E) % g_P);
u1 = shoup1s(u1, t1[j], shoup_c(t1[j]));
u2 = shoup1s(u2, t2[j], shoup_c(t2[j]));
u3 = shoup1s(u3, t3[j], shoup_c(t3[j]));
a[j] = u0; a[j + h] = u1; a[j + 2 * h] = u2; a[j + 3 * h] = u3;
}
for (; j + 8 <= h; j += 8) {
__m256i x0 = (j + 8 <= len + 1) ? LDU(src + j) : zero;
if (prescale != 1u) x0 = vshoup(x0, scv, scs, pv);
__m256i u0 = x0, u1 = x0, u2 = x0, u3 = x0;
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, der_ws(w), pv); }
STU(a + j, u0);
STU(a + j + h, u1);
STU(a + j + 2 * h, u2);
STU(a + j + 3 * h, u3);
}
for (; j < h; j++) {
u32 x0 = (j <= len) ? src[j] : 0;
if (prescale != 1u) x0 = shoup1s(x0, prescale, sc_scalar);
u32 u1 = shoup1s(x0, t1[j], shoup_c(t1[j]));
u32 u2 = shoup1s(x0, t2[j], shoup_c(t2[j]));
u32 u3 = shoup1s(x0, t3[j], shoup_c(t3[j]));
a[j] = x0; a[j + h] = u1; a[j + 2 * h] = u2; a[j + 3 * h] = u3;
}
}
// ================= forward (DIF) =================
TGT static void stage4_dif(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
u32 j = 0;
for (; j + 16 <= h; j += 16) {
__builtin_prefetch(p + j + 256);
__builtin_prefetch(p + j + h + 256);
__builtin_prefetch(p + j + 2 * h + 256);
__builtin_prefetch(p + j + 3 * h + 256);
__builtin_prefetch(t1 + j + 256);
__builtin_prefetch(t2 + j + 256);
__builtin_prefetch(t3 + j + 256);
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i y0 = LDU(p + j + 8);
__m256i y1 = LDU(p + j + h + 8);
__m256i y2 = LDU(p + j + 2 * h + 8);
__m256i y3 = LDU(p + j + 3 * h + 8);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
__m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i v0 = vaddm(A2, C2, pv), v2 = vsubm(A2, C2, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i F = vshoup(D2, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
__m256i v1 = vaddm(B2, F, pv), v3 = vsubm(B2, F, pv);
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, der_ws(w), pv); }
{ __m256i w = LDU(t1 + j + 8); v1 = vshoup(v1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j + 8); v2 = vshoup(v2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j + 8); v3 = vshoup(v3, w, der_ws(w), pv); }
STU(p + j, u0);
STU(p + j + h, u1);
STU(p + j + 2 * h, u2);
STU(p + j + 3 * h, u3);
STU(p + j + 8, v0);
STU(p + j + h + 8, v1);
STU(p + j + 2 * h + 8, v2);
STU(p + j + 3 * h + 8, v3);
}
for (; j + 8 <= h; j += 8) {
__builtin_prefetch(p + j + 256);
__builtin_prefetch(p + j + h + 256);
__builtin_prefetch(p + j + 2 * h + 256);
__builtin_prefetch(p + j + 3 * h + 256);
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, der_ws(w), pv); }
STU(p + j, u0);
STU(p + j + h, u1);
STU(p + j + 2 * h, u2);
STU(p + j + 3 * h, u3);
}
}
}
template<int H>
TGT static void stage4_dif_t(u32 *a, u32 blk, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const u32 h = (u32)H;
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
u32 j = 0;
for (; j + 16 <= h; j += 16) {
__builtin_prefetch(p + j + 768);
__builtin_prefetch(p + j + h + 256);
__builtin_prefetch(p + j + 2 * h + 256);
__builtin_prefetch(p + j + 3 * h + 256);
__builtin_prefetch(t1 + j + 512);
__builtin_prefetch(t2 + j + 512);
__builtin_prefetch(t3 + j + 512);
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i y0 = LDU(p + j + 8);
__m256i y1 = LDU(p + j + h + 8);
__m256i y2 = LDU(p + j + 2 * h + 8);
__m256i y3 = LDU(p + j + 3 * h + 8);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
__m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i v0 = vaddm(A2, C2, pv), v2 = vsubm(A2, C2, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i F = vshoup(D2, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
__m256i v1 = vaddm(B2, F, pv), v3 = vsubm(B2, F, pv);
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, der_ws(w), pv); }
{ __m256i w = LDU(t1 + j + 8); v1 = vshoup(v1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j + 8); v2 = vshoup(v2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j + 8); v3 = vshoup(v3, w, der_ws(w), pv); }
STU(p + j, u0);
STU(p + j + h, u1);
STU(p + j + 2 * h, u2);
STU(p + j + 3 * h, u3);
STU(p + j + 8, v0);
STU(p + j + h + 8, v1);
STU(p + j + 2 * h + 8, v2);
STU(p + j + 3 * h + 8, v3);
}
for (; j + 8 <= h; j += 8) {
__builtin_prefetch(p + j + 768);
__builtin_prefetch(p + j + h + 256);
__builtin_prefetch(p + j + 2 * h + 256);
__builtin_prefetch(p + j + 3 * h + 256);
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, der_ws(w), pv); }
STU(p + j, u0);
STU(p + j + h, u1);
STU(p + j + 2 * h, u2);
STU(p + j + 3 * h, u3);
}
}
}
TGT static void stage4_dif_s(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
u32 j = 0;
for (; j + 16 <= h; j += 16) {
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i y0 = LDU(p + j + 8);
__m256i y1 = LDU(p + j + h + 8);
__m256i y2 = LDU(p + j + 2 * h + 8);
__m256i y3 = LDU(p + j + 3 * h + 8);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
__m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i v0 = vaddm(A2, C2, pv), v2 = vsubm(A2, C2, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i F = vshoup(D2, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
__m256i v1 = vaddm(B2, F, pv), v3 = vsubm(B2, F, pv);
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, LDU(s1 + j), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, LDU(s2 + j), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, LDU(s3 + j), pv); }
{ __m256i w = LDU(t1 + j + 8); v1 = vshoup(v1, w, LDU(s1 + j + 8), pv); }
{ __m256i w = LDU(t2 + j + 8); v2 = vshoup(v2, w, LDU(s2 + j + 8), pv); }
{ __m256i w = LDU(t3 + j + 8); v3 = vshoup(v3, w, LDU(s3 + j + 8), pv); }
STU(p + j, u0);
STU(p + j + h, u1);
STU(p + j + 2 * h, u2);
STU(p + j + 3 * h, u3);
STU(p + j + 8, v0);
STU(p + j + h + 8, v1);
STU(p + j + 2 * h + 8, v2);
STU(p + j + 3 * h + 8, v3);
}
for (; j + 8 <= h; j += 8) {
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
{ __m256i w = LDU(t1 + j); u1 = vshoup(u1, w, LDU(s1 + j), pv); }
{ __m256i w = LDU(t2 + j); u2 = vshoup(u2, w, LDU(s2 + j), pv); }
{ __m256i w = LDU(t3 + j); u3 = vshoup(u3, w, LDU(s3 + j), pv); }
STU(p + j, u0);
STU(p + j + h, u1);
STU(p + j + 2 * h, u2);
STU(p + j + 3 * h, u3);
}
}
}
TGT static void stage4_dif_nt(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
u32 j = 0;
for (; j + 16 <= h; j += 16) {
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i y0 = LDU(p + j + 8);
__m256i y1 = LDU(p + j + h + 8);
__m256i y2 = LDU(p + j + 2 * h + 8);
__m256i y3 = LDU(p + j + 3 * h + 8);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
__m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i v0 = vaddm(A2, C2, pv), v2 = vsubm(A2, C2, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i F = vshoup(D2, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
__m256i v1 = vaddm(B2, F, pv), v3 = vsubm(B2, F, pv);
u1 = vshoup(u1, LDU(t1 + j), der_ws(LDU(t1 + j)), pv);
u2 = vshoup(u2, LDU(t2 + j), der_ws(LDU(t2 + j)), pv);
u3 = vshoup(u3, LDU(t3 + j), der_ws(LDU(t3 + j)), pv);
v1 = vshoup(v1, LDU(t1 + j + 8), der_ws(LDU(t1 + j + 8)), pv);
v2 = vshoup(v2, LDU(t2 + j + 8), der_ws(LDU(t2 + j + 8)), pv);
v3 = vshoup(v3, LDU(t3 + j + 8), der_ws(LDU(t3 + j + 8)), pv);
_mm256_stream_si256((__m256i *)(p + j), u0);
_mm256_stream_si256((__m256i *)(p + j + h), u1);
_mm256_stream_si256((__m256i *)(p + j + 2 * h), u2);
_mm256_stream_si256((__m256i *)(p + j + 3 * h), u3);
_mm256_stream_si256((__m256i *)(p + j + 8), v0);
_mm256_stream_si256((__m256i *)(p + j + h + 8), v1);
_mm256_stream_si256((__m256i *)(p + j + 2 * h + 8), v2);
_mm256_stream_si256((__m256i *)(p + j + 3 * h + 8), v3);
}
for (; j + 8 <= h; j += 8) {
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
u1 = vshoup(u1, LDU(t1 + j), der_ws(LDU(t1 + j)), pv);
u2 = vshoup(u2, LDU(t2 + j), der_ws(LDU(t2 + j)), pv);
u3 = vshoup(u3, LDU(t3 + j), der_ws(LDU(t3 + j)), pv);
_mm256_stream_si256((__m256i *)(p + j), u0);
_mm256_stream_si256((__m256i *)(p + j + h), u1);
_mm256_stream_si256((__m256i *)(p + j + 2 * h), u2);
_mm256_stream_si256((__m256i *)(p + j + 3 * h), u3);
}
}
}
TGT static void stage4_dif_h4(u32 *a, u32 nsub, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
if (nsub < 32) {
const __m128i pv4 = _mm256_castsi256_si128(pv), cI4 = _mm256_castsi256_si128(cI), cIs4 = _mm256_castsi256_si128(cIs);
__m128i W1 = _mm_loadu_si128((const __m128i *)t1), X1 = _mm_loadu_si128((const __m128i *)s1);
__m128i W2 = _mm_loadu_si128((const __m128i *)t2), X2 = _mm_loadu_si128((const __m128i *)s2);
__m128i W3 = _mm_loadu_si128((const __m128i *)t3), X3 = _mm_loadu_si128((const __m128i *)s3);
for (u32 i = 0; i < nsub; i += 16) {
u32 *p = a + i;
__m128i x0 = _mm_loadu_si128((const __m128i *)(p));
__m128i x1 = _mm_loadu_si128((const __m128i *)(p + 4));
__m128i x2 = _mm_loadu_si128((const __m128i *)(p + 8));
__m128i x3 = _mm_loadu_si128((const __m128i *)(p + 12));
__m128i A = add4(x0, x2, pv4), B = sub4(x0, x2, pv4);
__m128i C = add4(x1, x3, pv4), D = sub4(x1, x3, pv4);
__m128i u0 = add4(A, C, pv4), u2 = sub4(A, C, pv4);
__m128i E = sh4(D, cI4, cIs4, pv4);
__m128i u1 = add4(B, E, pv4), u3 = sub4(B, E, pv4);
u1 = sh4(u1, W1, X1, pv4); u2 = sh4(u2, W2, X2, pv4); u3 = sh4(u3, W3, X3, pv4);
_mm_storeu_si128((__m128i *)(p), u0);
_mm_storeu_si128((__m128i *)(p + 4), u1);
_mm_storeu_si128((__m128i *)(p + 8), u2);
_mm_storeu_si128((__m128i *)(p + 12), u3);
}
return;
}
__m256i W1 = bc128(_mm_loadu_si128((const __m128i *)t1));
__m256i W2 = bc128(_mm_loadu_si128((const __m128i *)t2));
__m256i W3 = bc128(_mm_loadu_si128((const __m128i *)t3));
__m256i X1 = bc128(_mm_loadu_si128((const __m128i *)s1));
__m256i X2 = bc128(_mm_loadu_si128((const __m128i *)s2));
__m256i X3 = bc128(_mm_loadu_si128((const __m128i *)s3));
for (u32 i = 0; i < nsub; i += 32) {
u32 *p = a + i;
__m256i x0 = ld2((const __m128i *)(p), (const __m128i *)(p + 16));
__m256i x1 = ld2((const __m128i *)(p + 4), (const __m128i *)(p + 20));
__m256i x2 = ld2((const __m128i *)(p + 8), (const __m128i *)(p + 24));
__m256i x3 = ld2((const __m128i *)(p + 12), (const __m128i *)(p + 28));
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
u1 = vshoup(u1, W1, X1, pv); u2 = vshoup(u2, W2, X2, pv); u3 = vshoup(u3, W3, X3, pv);
st2((__m128i *)(p), (__m128i *)(p + 16), u0);
st2((__m128i *)(p + 4), (__m128i *)(p + 20), u1);
st2((__m128i *)(p + 8), (__m128i *)(p + 24), u2);
st2((__m128i *)(p + 12), (__m128i *)(p + 28), u3);
}
}
// h==1 radix-4 (every stage twiddle at this level is w^0 = 1). 32 points per
// iteration: a 4x4 32-bit-lane transpose turns four 8-point shuffle groups
// (32 insns each) into one 67-instruction group.
TGT static void h1_rad4(u32 *a, u32 blk, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
const __m256i SORT = _mm256_setr_epi32(0, 4, 1, 5, 2, 6, 3, 7);
const __m256i RS = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
u32 i = 0;
for (; i + 32 <= blk; i += 32) {
u32 *p = a + i;
__m256i Y0 = LDU(p);
__m256i Y1 = LDU(p + 8);
__m256i Y2 = LDU(p + 16);
__m256i Y3 = LDU(p + 24);
__m256i T0 = _mm256_unpacklo_epi32(Y0, Y1);
__m256i T1 = _mm256_unpackhi_epi32(Y0, Y1);
__m256i T2 = _mm256_unpacklo_epi32(Y2, Y3);
__m256i T3 = _mm256_unpackhi_epi32(Y2, Y3);
__m256i Q0 = _mm256_permutevar8x32_epi32(T0, SORT);
__m256i Q1 = _mm256_permutevar8x32_epi32(T1, SORT);
__m256i Q2 = _mm256_permutevar8x32_epi32(T2, SORT);
__m256i Q3 = _mm256_permutevar8x32_epi32(T3, SORT);
__m256i X0 = _mm256_permute2x128_si256(Q0, Q2, 0x20);
__m256i X1 = _mm256_permute2x128_si256(Q0, Q2, 0x31);
__m256i X2 = _mm256_permute2x128_si256(Q1, Q3, 0x20);
__m256i X3 = _mm256_permute2x128_si256(Q1, Q3, 0x31);
__m256i A = vaddm(X0, X2, pv), B = vsubm(X0, X2, pv);
__m256i C = vaddm(X1, X3, pv), D = vsubm(X1, X3, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i U0 = vaddm(A, C, pv), U1 = vaddm(B, E, pv);
__m256i U2 = vsubm(A, C, pv), U3 = vsubm(B, E, pv);
__m256i PA = _mm256_unpacklo_epi32(U0, U1);
__m256i QA = _mm256_unpacklo_epi32(U2, U3);
__m256i RA = _mm256_unpackhi_epi32(U0, U1);
__m256i SA = _mm256_unpackhi_epi32(U2, U3);
__m256i W0 = _mm256_permute2x128_si256(PA, QA, 0x20);
__m256i W1 = _mm256_permute2x128_si256(PA, QA, 0x31);
__m256i W2 = _mm256_permute2x128_si256(RA, SA, 0x20);
__m256i W3 = _mm256_permute2x128_si256(RA, SA, 0x31);
STU(p, _mm256_permutevar8x32_epi32(W0, RS));
STU(p + 8, _mm256_permutevar8x32_epi32(W1, RS));
STU(p + 16, _mm256_permutevar8x32_epi32(W2, RS));
STU(p + 24, _mm256_permutevar8x32_epi32(W3, RS));
}
for (; i + 8 <= blk; i += 8) {
u32 *p = a + i;
__m256i x = LDU((const __m256i *)p);
__m256i t = _mm256_shuffle_epi32(x, _MM_SHUFFLE(1, 0, 3, 2));
__m256i s = vaddm(x, t, pv);
__m256i d = vsubm(x, t, pv);
__m256i Z = _mm256_shuffle_epi32(d, 0x00);
__m256i P = _mm256_blend_epi32(s, Z, 0xAA);
__m256i E = vshoup(d, cI, cIs, pv);
__m256i E2 = _mm256_shuffle_epi32(E, 0x55);
__m256i W = _mm256_shuffle_epi32(s, 0x55);
__m256i Q = _mm256_blend_epi32(W, E2, 0xAA);
__m256i sum = vaddm(P, Q, pv);
__m256i dif = vsubm(P, Q, pv);
STU((__m256i *)p, _mm256_blend_epi32(sum, dif, 0xCC));
}
}
TGT static void stage4_dif_h1(u32 *a, u32 blk, u32 ci, u32 cis) { h1_rad4(a, blk, ci, cis); }
TGT static void stage2_dif(u32 *a, u32 blk, const u32 *t, const u32 *s) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
u32 h = blk >> 1;
for (u32 j = 0; j < h; j += 8) {
__m256i x0 = LDU(a + j);
__m256i x1 = LDU(a + j + h);
__m256i u0 = vaddm(x0, x1, pv);
__m256i u1 = vshoup(vsubm(x0, x1, pv), LDU(t + j),
LDU(s + j), pv);
STU(a + j, u0);
STU(a + j + h, u1);
}
}
// ================= inverse (DIT) =================
TGT static void merge4_dit(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
for (u32 j = 0; j < h; j += 8) {
__builtin_prefetch(p + j + 768);
__builtin_prefetch(p + j + h + 256);
__builtin_prefetch(p + j + 2 * h + 256);
__builtin_prefetch(p + j + 3 * h + 256);
__builtin_prefetch(t1 + j + 768);
__builtin_prefetch(t2 + j + 768);
__builtin_prefetch(t3 + j + 768);
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
{ __m256i w = LDU(t1 + j); x1 = vshoup(x1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); x2 = vshoup(x2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); x3 = vshoup(x3, w, der_ws(w), pv); }
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i E = vshoup(D, cI, cIs, pv);
STU(p + j, vaddm(A, C, pv));
STU(p + j + h, vaddm(B, E, pv));
STU(p + j + 2 * h, vsubm(A, C, pv));
STU(p + j + 3 * h, vsubm(B, E, pv));
}
}
}
template<int H>
TGT static void merge4_dit_t(u32 *a, u32 blk, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const u32 h = (u32)H;
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
u32 j = 0;
for (; j + 16 <= h; j += 16) {
__builtin_prefetch(p + j + 768);
__builtin_prefetch(p + j + h + 256);
__builtin_prefetch(p + j + 2 * h + 256);
__builtin_prefetch(p + j + 3 * h + 256);
__builtin_prefetch(t1 + j + 512);
__builtin_prefetch(t2 + j + 512);
__builtin_prefetch(t3 + j + 512);
__builtin_prefetch(p + j + 776);
__builtin_prefetch(p + j + h + 264);
__builtin_prefetch(p + j + 2 * h + 264);
__builtin_prefetch(p + j + 3 * h + 264);
__builtin_prefetch(t1 + j + 512);
__builtin_prefetch(t2 + j + 512);
__builtin_prefetch(t3 + j + 512);
__m256i x0 = LDU(p + j), x1 = LDU(p + j + h), x2 = LDU(p + j + 2 * h), x3 = LDU(p + j + 3 * h);
__m256i y0 = LDU(p + j + 8), y1 = LDU(p + j + h + 8), y2 = LDU(p + j + 2 * h + 8), y3 = LDU(p + j + 3 * h + 8);
{ __m256i w = LDU(t1 + j); x1 = vshoup(x1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); x2 = vshoup(x2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); x3 = vshoup(x3, w, der_ws(w), pv); }
{ __m256i w = LDU(t1 + j + 8); y1 = vshoup(y1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j + 8); y2 = vshoup(y2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j + 8); y3 = vshoup(y3, w, der_ws(w), pv); }
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
__m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
__m256i E = vshoup(D, cI, cIs, pv), F = vshoup(D2, cI, cIs, pv);
STU(p + j, vaddm(A, C, pv));
STU(p + j + h, vaddm(B, E, pv));
STU(p + j + 2 * h, vsubm(A, C, pv));
STU(p + j + 3 * h, vsubm(B, E, pv));
STU(p + j + 8, vaddm(A2, C2, pv));
STU(p + j + h + 8, vaddm(B2, F, pv));
STU(p + j + 2 * h + 8, vsubm(A2, C2, pv));
STU(p + j + 3 * h + 8, vsubm(B2, F, pv));
}
for (; j < h; j += 8) {
__builtin_prefetch(p + j + 768);
__builtin_prefetch(p + j + h + 256);
__builtin_prefetch(p + j + 2 * h + 256);
__builtin_prefetch(p + j + 3 * h + 256);
__builtin_prefetch(t1 + j + 512);
__builtin_prefetch(t2 + j + 512);
__builtin_prefetch(t3 + j + 512);
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
{ __m256i w = LDU(t1 + j); x1 = vshoup(x1, w, der_ws(w), pv); }
{ __m256i w = LDU(t2 + j); x2 = vshoup(x2, w, der_ws(w), pv); }
{ __m256i w = LDU(t3 + j); x3 = vshoup(x3, w, der_ws(w), pv); }
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i E = vshoup(D, cI, cIs, pv);
STU(p + j, vaddm(A, C, pv));
STU(p + j + h, vaddm(B, E, pv));
STU(p + j + 2 * h, vsubm(A, C, pv));
STU(p + j + 3 * h, vsubm(B, E, pv));
}
}
}
TGT static void merge4_dit_s(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
for (u32 j = 0; j < h; j += 8) {
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
{ __m256i w = LDU(t1 + j); x1 = vshoup(x1, w, LDU(s1 + j), pv); }
{ __m256i w = LDU(t2 + j); x2 = vshoup(x2, w, LDU(s2 + j), pv); }
{ __m256i w = LDU(t3 + j); x3 = vshoup(x3, w, LDU(s3 + j), pv); }
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i E = vshoup(D, cI, cIs, pv);
STU(p + j, vaddm(A, C, pv));
STU(p + j + h, vaddm(B, E, pv));
STU(p + j + 2 * h, vsubm(A, C, pv));
STU(p + j + 3 * h, vsubm(B, E, pv));
}
}
}
TGT static void merge4_dit_nt(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * h) {
u32 *p = a + i;
for (u32 j = 0; j < h; j += 8) {
__m256i x0 = LDU(p + j);
__m256i x1 = LDU(p + j + h);
__m256i x2 = LDU(p + j + 2 * h);
__m256i x3 = LDU(p + j + 3 * h);
x1 = vshoup(x1, LDU(t1 + j), der_ws(LDU(t1 + j)), pv);
x2 = vshoup(x2, LDU(t2 + j), der_ws(LDU(t2 + j)), pv);
x3 = vshoup(x3, LDU(t3 + j), der_ws(LDU(t3 + j)), pv);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i E = vshoup(D, cI, cIs, pv);
_mm256_stream_si256((__m256i *)(p + j), vaddm(A, C, pv));
_mm256_stream_si256((__m256i *)(p + j + h), vaddm(B, E, pv));
_mm256_stream_si256((__m256i *)(p + j + 2 * h), vsubm(A, C, pv));
_mm256_stream_si256((__m256i *)(p + j + 3 * h), vsubm(B, E, pv));
}
}
}
TGT static void merge4_dit_h4(u32 *a, u32 nsub, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
if (nsub < 32) {
const __m128i pv4 = _mm256_castsi256_si128(pv), cI4 = _mm256_castsi256_si128(cI), cIs4 = _mm256_castsi256_si128(cIs);
__m128i W1 = _mm_loadu_si128((const __m128i *)t1), X1 = _mm_loadu_si128((const __m128i *)s1);
__m128i W2 = _mm_loadu_si128((const __m128i *)t2), X2 = _mm_loadu_si128((const __m128i *)s2);
__m128i W3 = _mm_loadu_si128((const __m128i *)t3), X3 = _mm_loadu_si128((const __m128i *)s3);
for (u32 i = 0; i < nsub; i += 16) {
u32 *p = a + i;
__m128i x0 = _mm_loadu_si128((const __m128i *)(p));
__m128i x1 = sh4(_mm_loadu_si128((const __m128i *)(p + 4)), W1, X1, pv4);
__m128i x2 = sh4(_mm_loadu_si128((const __m128i *)(p + 8)), W2, X2, pv4);
__m128i x3 = sh4(_mm_loadu_si128((const __m128i *)(p + 12)), W3, X3, pv4);
__m128i A = add4(x0, x2, pv4), B = sub4(x0, x2, pv4);
__m128i C = add4(x1, x3, pv4), D = sub4(x1, x3, pv4);
__m128i E = sh4(D, cI4, cIs4, pv4);
_mm_storeu_si128((__m128i *)(p), add4(A, C, pv4));
_mm_storeu_si128((__m128i *)(p + 4), add4(B, E, pv4));
_mm_storeu_si128((__m128i *)(p + 8), sub4(A, C, pv4));
_mm_storeu_si128((__m128i *)(p + 12), sub4(B, E, pv4));
}
return;
}
__m256i W1 = bc128(_mm_loadu_si128((const __m128i *)t1));
__m256i W2 = bc128(_mm_loadu_si128((const __m128i *)t2));
__m256i W3 = bc128(_mm_loadu_si128((const __m128i *)t3));
__m256i X1 = bc128(_mm_loadu_si128((const __m128i *)s1));
__m256i X2 = bc128(_mm_loadu_si128((const __m128i *)s2));
__m256i X3 = bc128(_mm_loadu_si128((const __m128i *)s3));
for (u32 i = 0; i < nsub; i += 32) {
u32 *p = a + i;
__m256i x0 = ld2((const __m128i *)(p), (const __m128i *)(p + 16));
__m256i x1 = ld2((const __m128i *)(p + 4), (const __m128i *)(p + 20));
__m256i x2 = ld2((const __m128i *)(p + 8), (const __m128i *)(p + 24));
__m256i x3 = ld2((const __m128i *)(p + 12), (const __m128i *)(p + 28));
x1 = vshoup(x1, W1, X1, pv); x2 = vshoup(x2, W2, X2, pv); x3 = vshoup(x3, W3, X3, pv);
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i E = vshoup(D, cI, cIs, pv);
st2((__m128i *)(p), (__m128i *)(p + 16), vaddm(A, C, pv));
st2((__m128i *)(p + 4), (__m128i *)(p + 20), vaddm(B, E, pv));
st2((__m128i *)(p + 8), (__m128i *)(p + 24), vsubm(A, C, pv));
st2((__m128i *)(p + 12), (__m128i *)(p + 28), vsubm(B, E, pv));
}
}
TGT static void merge4_dit_h1(u32 *a, u32 blk, u32 ci, u32 cis) { h1_rad4(a, blk, ci, cis); }
TGT static void merge2_dit(u32 *a, u32 blk, const u32 *t, const u32 *s) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
u32 h = blk >> 1;
for (u32 j = 0; j < h; j += 8) {
__m256i x0 = LDU(a + j);
__m256i x1 = vshoup(LDU(a + j + h),
LDU(t + j),
LDU(s + j), pv);
STU(a + j, vaddm(x0, x1, pv));
STU(a + j + h, vsubm(x0, x1, pv));
}
}
// ================= table generation =================
// dst[i] = v0*r^i ; sc[i] = floor(dst[i]*2^32/p) (within 1, sufficient for Shoup)
TGT static void gen_geo(u32 *dst, u32 *sc, u32 cnt, u32 v0, u32 r) {
if (!cnt) return;
const __m256i pv = _mm256_set1_epi32((int)g_P);
u32 seq[8], cur = v0 % g_P;
for (int i = 0; i < 8; i++) { seq[i] = cur; cur = mulmod(cur, r); }
u32 r8 = 1; for (int t = 0; t < 8; t++) r8 = mulmod(r8, r);
u32 rs8 = shoup_c(r8);
__m256i v = LDU((const __m256i *)seq);
__m256i r8v = _mm256_set1_epi32((int)r8), rs8v = _mm256_set1_epi32((int)rs8);
u32 i = 0;
const bool nt = NT_OK(dst, cnt) && NT_OK(sc, cnt);
if (cnt >= 32) {
__m256i w1 = vshoup(v, r8v, rs8v, pv);
__m256i w2 = vshoup(w1, r8v, rs8v, pv);
__m256i w3 = vshoup(w2, r8v, rs8v, pv);
u32 r32 = r8; for (int t = 0; t < 3; t++) r32 = mulmod(r32, r8);
const u32 rs32 = shoup_c(r32);
const __m256i r32v = _mm256_set1_epi32((int)r32), rs32v = _mm256_set1_epi32((int)rs32);
if (nt) {
for (; i + 32 <= cnt; i += 32) {
stu256_nt(dst + i, v); stu256_nt(sc + i, der_ws(v));
stu256_nt(dst + i + 8, w1); stu256_nt(sc + i + 8, der_ws(w1));
stu256_nt(dst + i + 16, w2); stu256_nt(sc + i + 16, der_ws(w2));
stu256_nt(dst + i + 24, w3); stu256_nt(sc + i + 24, der_ws(w3));
v = vshoup(v, r32v, rs32v, pv);
w1 = vshoup(w1, r32v, rs32v, pv);
w2 = vshoup(w2, r32v, rs32v, pv);
w3 = vshoup(w3, r32v, rs32v, pv);
}
} else {
for (; i + 32 <= cnt; i += 32) {
STU(dst + i, v); STU(sc + i, der_ws(v));
STU(dst + i + 8, w1); STU(sc + i + 8, der_ws(w1));
STU(dst + i + 16, w2); STU(sc + i + 16, der_ws(w2));
STU(dst + i + 24, w3); STU(sc + i + 24, der_ws(w3));
v = vshoup(v, r32v, rs32v, pv);
w1 = vshoup(w1, r32v, rs32v, pv);
w2 = vshoup(w2, r32v, rs32v, pv);
w3 = vshoup(w3, r32v, rs32v, pv);
}
}
}
if (nt) {
for (; i + 8 <= cnt; i += 8) { stu256_nt(dst + i, v); stu256_nt(sc + i, der_ws(v)); v = vshoup(v, r8v, rs8v, pv); }
} else {
for (; i + 8 <= cnt; i += 8) { STU(dst + i, v); STU(sc + i, der_ws(v)); v = vshoup(v, r8v, rs8v, pv); }
}
if (i < cnt) {
u32 x = v0 % g_P;
for (u32 t = 0; t < i; t++) x = mulmod(x, r);
for (; i < cnt; i++) { dst[i] = x; sc[i] = shoup_c(x); x = mulmod(x, r); }
}
}
TGT static void gen_geo_val(u32 *dst, u32 cnt, u32 v0, u32 r) {
if (!cnt) return;
const __m256i pv = _mm256_set1_epi32((int)g_P);
u32 seq[8], cur = v0 % g_P;
for (int i = 0; i < 8; i++) { seq[i] = cur; cur = mulmod(cur, r); }
u32 r8 = 1; for (int t = 0; t < 8; t++) r8 = mulmod(r8, r);
u32 rs8 = shoup_c(r8);
__m256i v = LDU((const __m256i *)seq);
__m256i r8v = _mm256_set1_epi32((int)r8), rs8v = _mm256_set1_epi32((int)rs8);
u32 i = 0;
if (NT_OK(dst, cnt)) { /* write-once stream: skip the RFO */
if (cnt >= 32) {
__m256i w1 = vshoup(v, r8v, rs8v, pv);
__m256i w2 = vshoup(w1, r8v, rs8v, pv);
__m256i w3 = vshoup(w2, r8v, rs8v, pv);
u32 r32 = r8; for (int t = 0; t < 3; t++) r32 = mulmod(r32, r8);
u32 rs32 = shoup_c(r32);
__m256i r32v = _mm256_set1_epi32((int)r32), rs32v = _mm256_set1_epi32((int)rs32);
for (; i + 32 <= cnt; i += 32) {
stu256_nt(dst + i, v); v = vshoup(v, r32v, rs32v, pv);
stu256_nt(dst + i + 8, w1); w1 = vshoup(w1, r32v, rs32v, pv);
stu256_nt(dst + i + 16, w2); w2 = vshoup(w2, r32v, rs32v, pv);
stu256_nt(dst + i + 24, w3); w3 = vshoup(w3, r32v, rs32v, pv);
}
}
for (; i + 8 <= cnt; i += 8) { stu256_nt(dst + i, v); v = vshoup(v, r8v, rs8v, pv); }
if (i < cnt) { u32 x = v0 % g_P; for (u32 t = 0; t < i; t++) x = mulmod(x, r); for (; i < cnt; i++) { dst[i] = x; x = mulmod(x, r); } }
return;
}
if (cnt >= 32) {
__m256i w1 = vshoup(v, r8v, rs8v, pv);
__m256i w2 = vshoup(w1, r8v, rs8v, pv);
__m256i w3 = vshoup(w2, r8v, rs8v, pv);
u32 r32 = r8; for (int t = 0; t < 3; t++) r32 = mulmod(r32, r8);
u32 rs32 = shoup_c(r32);
__m256i r32v = _mm256_set1_epi32((int)r32), rs32v = _mm256_set1_epi32((int)rs32);
for (; i + 32 <= cnt; i += 32) {
STU(dst + i, v); v = vshoup(v, r32v, rs32v, pv);
STU(dst + i + 8, w1); w1 = vshoup(w1, r32v, rs32v, pv);
STU(dst + i + 16, w2); w2 = vshoup(w2, r32v, rs32v, pv);
STU(dst + i + 24, w3); w3 = vshoup(w3, r32v, rs32v, pv);
}
}
for (; i + 8 <= cnt; i += 8) { STU(dst + i, v); v = vshoup(v, r8v, rs8v, pv); }
if (i < cnt) { u32 x = v0 % g_P; for (u32 t = 0; t < i; t++) x = mulmod(x, r); for (; i < cnt; i++) { dst[i] = x; x = mulmod(x, r); } }
}
static void build_tabs(Tabs &T, u32 *tab, u32 n, u32 Sb, int L, int odd, u32 root) {
u32 *p = tab;
u32 S = n;
for (int l = 0; l < L; l++) {
p = pad8(p);
u32 M = S >> 2;
// Same power-of-two stride problem, one level up: the three DRAM-level twiddle
// tables are M u32 apart with M a multiple of the 4 KB L1 set period, so the
// three table streams plus the four array streams of stage4_dif/merge4_dit all
// pile onto the same L1/L2 sets (the judge's L2 is 256 KB 4-WAY, 64 KB period).
// A 48 KB pad rotates the set index per table; the values themselves are
// unchanged (only the addresses move).
u32 Sl = M + 12288;
T.l1[l] = p; T.l2[l] = p + Sl; T.l3[l] = p + 2 * Sl;
T.m1[l] = p; T.m2[l] = p + Sl; T.m3[l] = p + 2 * Sl;
u32 base = powmod(root, n / S);
u32 b1 = base, b2 = mulmod(base, base), b3 = mulmod(b2, base);
gen_geo_val(p, M, 1, b1);
gen_geo_val(p + Sl, M, 1, b2);
gen_geo_val(p + 2 * Sl, M, 1, b3);
p += 3 * Sl;
S >>= 2;
}
T.b2 = 0; T.b2s = 0;
if (odd) {
p = pad8(p);
u32 M = Sb >> 1;
T.b2 = p; T.b2s = p + M;
gen_geo(p, p + M, M, 1, powmod(root, n / Sb));
p += 2 * M;
}
int s = 0;
u32 S2v = 4;
while (s < g_nb) {
p = pad8(p);
u32 M = S2v >> 2;
// The six base-case tables are M u32 each; at M = 4096 (16 KB) they would
// land on IDENTICAL L1/L2 sets (both spans are multiples of the 4 KB set
// period). One 64-byte line of gap per table rotates the set index, and
// keeps every start 32-byte aligned.
u32 St = M + 256;
T.h1[s] = p; T.h2[s] = p + St; T.h3[s] = p + 2 * St;
T.q1[s] = p + 3 * St; T.q2[s] = p + 4 * St; T.q3[s] = p + 5 * St;
u32 base = powmod(root, n / S2v);
u32 b1 = base, b2 = mulmod(base, base), b3 = mulmod(b2, base);
gen_geo(p, p + 3 * St, M, 1, b1);
gen_geo(p + St, p + 4 * St, M, 1, b2);
gen_geo(p + 2 * St, p + 5 * St, M, 1, b3);
p += 6 * St;
S2v <<= 2;
s++;
}
}
static u32 g_tabwords;
static u32 g_initP = 2013265921u;
static void ntt_init(u32 n) {
g_P = g_initP; g_pinv = 0;
{ u32 iv = 1; for (int q = 0; q < 5; q++) iv *= 2u - g_P * iv; g_pinv = 0u - iv;
u64 M = (u64)(~0ull) / g_P; u32 mh = (u32)(M >> 32); g_Msh = 0; while ((1u << g_Msh) < mh) g_Msh++; g_Ml = (u32)(M & 0xffffffffu); }
g_n = n;
int lg = 0; while ((1u << lg) < n) lg++;
g_log2n = lg;
const int BT = g_BT;
int L = 0;
while (lg - 2 * (L + 1) >= BT) L++;
g_L = L;
g_Sb = n >> (2 * L);
g_odd = ((lg - 2 * L) & 1) ? 1 : 0;
u32 S2 = g_Sb; if (g_odd) S2 >>= 1;
g_nb = 0; while (S2 >= 4) { S2 >>= 2; g_nb++; }
u32 w = 0;
{ u32 e = (g_P - 1) / n;
for (u32 g = 2; g < 400 && !w; g++) {
u32 c = powmod(g, e);
if (c != 1 && powmod(c, n >> 1) != 1) w = c; /* order exactly n (n = 2^k) */
} }
u32 wi = powmod(w, g_P - 2);
u32 iF = powmod(w, n >> 2), iI = powmod(wi, n >> 2);
g_ciF = iF; g_cisF = shoup_c(iF);
g_ciI = iI; g_cisI = shoup_c(iI);
{ u32 S=n; u32 tot=0; for (int l=0;l<L;l++){tot+=3*(S>>2);S>>=2;} u32 S2v=4; for(int s=0;s<g_nb;s++){tot+=6*(S2v>>2);S2v<<=2;} if(g_odd) tot+=2*(g_Sb>>1); g_tabwords=tot; }
build_tabs(g_TF, g_tabF, n, g_Sb, L, g_odd, w);
build_tabs(g_TI, g_tabI, n, g_Sb, L, g_odd, wi);
nt_fence();
}
TGT static void ntt_fwd_dram(u32 *a, const u32 *src = 0, u32 len = 0,
u32 prescale = 1u) {
u32 S = g_n;
int l0 = 0;
if (src && g_L >= 1 && len >= (g_n >> 2) && len <= (g_n >> 1)) {
stage0_dif_src(a, src, len, g_n >> 2, g_TF.l1[0], g_TF.l2[0], g_TF.l3[0], g_ciF, g_cisF, prescale);
S = g_n >> 2;
l0 = 1;
}
for (int l = l0; l < g_L; l++) {
u32 h = S >> 2;
for (u32 b = 0; b < g_n; b += S)
{ u32 *pp = a + b;
switch (h) {
case 131072: stage4_dif_t<131072>(pp, S, g_TF.l1[l], g_TF.l2[l], g_TF.l3[l], g_TF.m1[l], g_TF.m2[l], g_TF.m3[l], g_ciF, g_cisF); break;
case 32768: stage4_dif_t<32768>(pp, S, g_TF.l1[l], g_TF.l2[l], g_TF.l3[l], g_TF.m1[l], g_TF.m2[l], g_TF.m3[l], g_ciF, g_cisF); break;
case 8192: stage4_dif_t<8192>(pp, S, g_TF.l1[l], g_TF.l2[l], g_TF.l3[l], g_TF.m1[l], g_TF.m2[l], g_TF.m3[l], g_ciF, g_cisF); break;
case 2048: stage4_dif_t<2048>(pp, S, g_TF.l1[l], g_TF.l2[l], g_TF.l3[l], g_TF.m1[l], g_TF.m2[l], g_TF.m3[l], g_ciF, g_cisF); break;
default: stage4_dif(pp, S, h, g_TF.l1[l], g_TF.l2[l], g_TF.l3[l], g_TF.m1[l], g_TF.m2[l], g_TF.m3[l], g_ciF, g_cisF);
} }
S = h;
}
}
template<int H>
TGT static void stage4_dif_h(u32 *a, u32 blk, u32 _h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * H) {
u32 *p = a + i;
u32 j = 0;
for (; j + 16 <= H; j += 16) {
__m256i x0 = LDU((p + j));
__m256i x1 = LDU((p + j + H));
__m256i x2 = LDU((p + j + 2 * H));
__m256i x3 = LDU((p + j + 3 * H));
__m256i y0 = LDU((p + j + 8));
__m256i y1 = LDU((p + j + H + 8));
__m256i y2 = LDU((p + j + 2 * H + 8));
__m256i y3 = LDU((p + j + 3 * H + 8));
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
__m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i v0 = vaddm(A2, C2, pv), v2 = vsubm(A2, C2, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i F = vshoup(D2, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
__m256i v1 = vaddm(B2, F, pv), v3 = vsubm(B2, F, pv);
{ __m256i w = LDU((t1 + j)); u1 = vshoup(u1, w, LDU((s1 + j)), pv); }
{ __m256i w = LDU((t2 + j)); u2 = vshoup(u2, w, LDU((s2 + j)), pv); }
{ __m256i w = LDU((t3 + j)); u3 = vshoup(u3, w, LDU((s3 + j)), pv); }
{ __m256i w = LDU((t1 + j + 8)); v1 = vshoup(v1, w, LDU((s1 + j + 8)), pv); }
{ __m256i w = LDU((t2 + j + 8)); v2 = vshoup(v2, w, LDU((s2 + j + 8)), pv); }
{ __m256i w = LDU((t3 + j + 8)); v3 = vshoup(v3, w, LDU((s3 + j + 8)), pv); }
STU((p + j), u0);
STU((p + j + H), u1);
STU((p + j + 2 * H), u2);
STU((p + j + 3 * H), u3);
STU((p + j + 8), v0);
STU((p + j + H + 8), v1);
STU((p + j + 2 * H + 8), v2);
STU((p + j + 3 * H + 8), v3);
}
for (; j + 8 <= H; j += 8) {
__builtin_prefetch(p + j + 256);
__builtin_prefetch(p + j + H + 256);
__builtin_prefetch(p + j + 2 * H + 256);
__builtin_prefetch(p + j + 3 * H + 256);
__m256i x0 = LDU((p + j));
__m256i x1 = LDU((p + j + H));
__m256i x2 = LDU((p + j + 2 * H));
__m256i x3 = LDU((p + j + 3 * H));
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
__m256i E = vshoup(D, cI, cIs, pv);
__m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
{ __m256i w = LDU((t1 + j)); u1 = vshoup(u1, w, LDU((s1 + j)), pv); }
{ __m256i w = LDU((t2 + j)); u2 = vshoup(u2, w, LDU((s2 + j)), pv); }
{ __m256i w = LDU((t3 + j)); u3 = vshoup(u3, w, LDU((s3 + j)), pv); }
STU((p + j), u0);
STU((p + j + H), u1);
STU((p + j + 2 * H), u2);
STU((p + j + 3 * H), u3);
}
}
}
template<int H>
TGT static void merge4_dit_h(u32 *a, u32 blk, u32 _h, const u32 *t1, const u32 *t2, const u32 *t3,
const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
for (u32 i = 0; i < blk; i += 4 * H) {
u32 *p = a + i;
u32 j = 0;
for (; j + 16 <= H; j += 16) {
__m256i x0 = LDU((p + j)), x1 = LDU((p + j + H)), x2 = LDU((p + j + 2 * H)), x3 = LDU((p + j + 3 * H));
__m256i y0 = LDU((p + j + 8)), y1 = LDU((p + j + H + 8)), y2 = LDU((p + j + 2 * H + 8)), y3 = LDU((p + j + 3 * H + 8));
{ __m256i w = LDU((t1 + j)); x1 = vshoup(x1, w, LDU((s1 + j)), pv); }
{ __m256i w = LDU((t2 + j)); x2 = vshoup(x2, w, LDU((s2 + j)), pv); }
{ __m256i w = LDU((t3 + j)); x3 = vshoup(x3, w, LDU((s3 + j)), pv); }
{ __m256i w = LDU((t1 + j + 8)); y1 = vshoup(y1, w, LDU((s1 + j + 8)), pv); }
{ __m256i w = LDU((t2 + j + 8)); y2 = vshoup(y2, w, LDU((s2 + j + 8)), pv); }
{ __m256i w = LDU((t3 + j + 8)); y3 = vshoup(y3, w, LDU((s3 + j + 8)), pv); }
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
__m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
__m256i E = vshoup(D, cI, cIs, pv), F = vshoup(D2, cI, cIs, pv);
STU((p + j), vaddm(A, C, pv));
STU((p + j + H), vaddm(B, E, pv));
STU((p + j + 2 * H), vsubm(A, C, pv));
STU((p + j + 3 * H), vsubm(B, E, pv));
STU((p + j + 8), vaddm(A2, C2, pv));
STU((p + j + H + 8), vaddm(B2, F, pv));
STU((p + j + 2 * H + 8), vsubm(A2, C2, pv));
STU((p + j + 3 * H + 8), vsubm(B2, F, pv));
}
for (; j < H; j += 8) {
__m256i x0 = LDU((p + j));
__m256i x1 = LDU((p + j + H));
__m256i x2 = LDU((p + j + 2 * H));
__m256i x3 = LDU((p + j + 3 * H));
{ __m256i w = LDU((t1 + j)); x1 = vshoup(x1, w, LDU((s1 + j)), pv); }
{ __m256i w = LDU((t2 + j)); x2 = vshoup(x2, w, LDU((s2 + j)), pv); }
{ __m256i w = LDU((t3 + j)); x3 = vshoup(x3, w, LDU((s3 + j)), pv); }
__m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
__m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
__m256i E = vshoup(D, cI, cIs, pv);
STU((p + j), vaddm(A, C, pv));
STU((p + j + H), vaddm(B, E, pv));
STU((p + j + 2 * H), vsubm(A, C, pv));
STU((p + j + 3 * H), vsubm(B, E, pv));
}
}
}
TGT static void ntt_fwd_block(u32 *p) {
if (g_odd) stage2_dif(p, g_Sb, g_TF.b2, g_TF.b2s);
for (int s = g_nb - 1; s >= 0; s--) {
u32 S2 = 4u << (2 * s);
if (S2 >= 32) {
for (u32 i = 0; i < g_Sb; i += S2)
{ u32 *pp = p + i;
switch (S2) {
case 16384: stage4_dif_h<4096>(pp, 16384, 4096, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF); break;
case 4096: stage4_dif_h<1024>(pp, 4096, 1024, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF); break;
case 1024: stage4_dif_h<256>(pp, 1024, 256, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF); break;
case 256: stage4_dif_h<64>(pp, 256, 64, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF); break;
case 64: stage4_dif_h<16>(pp, 64, 16, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF); break;
default: stage4_dif_s(pp, S2, S2 >> 2, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF);
} }
} else if (S2 == 16) {
stage4_dif_h4(p, g_Sb, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF);
} else {
stage4_dif_h1(p, g_Sb, g_ciF, g_cisF);
}
}
}
TGT static void ntt_inv_block(u32 *p, const u32 *q, u32 nscale) {
if (q) {
const __m256i pv = _mm256_set1_epi32((int)g_P);
/* (a) vshoup(vmont(x,y), ninv*R) == x*y*ninv mod P: one Montgomery reduction
plus one Shoup multiply replaces two chained Montgomery reductions.
(b) 4 vectors per iteration: the loop is issue-latency limited, not
multiplier-port limited, so the extra independent chains pay. */
(void)nscale; // already incorporated in the second forward transform
u32 i = 0;
for (; i + 32 <= g_Sb; i += 32) {
__m256i x0 = LDU(p + i), y0 = LDU(q + i);
__m256i x1 = LDU(p + i + 8), y1 = LDU(q + i + 8);
__m256i x2 = LDU(p + i + 16), y2 = LDU(q + i + 16);
__m256i x3 = LDU(p + i + 24), y3 = LDU(q + i + 24);
STU(p + i, vmont(x0, y0, pv));
STU(p + i + 8, vmont(x1, y1, pv));
STU(p + i + 16, vmont(x2, y2, pv));
STU(p + i + 24, vmont(x3, y3, pv));
}
for (; i < g_Sb; i += 8) {
__m256i x = LDU(p + i);
__m256i y = LDU(q + i);
STU(p + i, vmont(x, y, pv));
}
}
for (int s = 0; s < g_nb; s++) {
u32 S = 4u << (2 * s);
if (S >= 32) {
for (u32 i = 0; i < g_Sb; i += S)
{ u32 *pp = p + i;
switch (S) {
case 16384: merge4_dit_h<4096>(pp, 16384, 4096, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI); break;
case 4096: merge4_dit_h<1024>(pp, 4096, 1024, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI); break;
case 1024: merge4_dit_h<256>(pp, 1024, 256, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI); break;
case 256: merge4_dit_h<64>(pp, 256, 64, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI); break;
case 64: merge4_dit_h<16>(pp, 64, 16, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI); break;
default: merge4_dit_s(pp, S, S >> 2, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI);
} }
} else if (S == 16) {
merge4_dit_h4(p, g_Sb, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI);
} else {
merge4_dit_h1(p, g_Sb, g_ciI, g_cisI);
}
}
if (g_odd) merge2_dit(p, g_Sb, g_TI.b2, g_TI.b2s);
}
TGT static void ntt_inv_dram(u32 *a) {
for (int l = g_L - 1; l >= 0; l--) {
u32 SB = g_n >> (2 * l);
u32 h = SB >> 2;
for (u32 b = 0; b < g_n; b += SB)
{ u32 *pp = a + b;
switch (h) {
case 131072: merge4_dit_t<131072>(pp, SB, g_TI.l1[l], g_TI.l2[l], g_TI.l3[l], g_TI.m1[l], g_TI.m2[l], g_TI.m3[l], g_ciI, g_cisI); break;
case 32768: merge4_dit_t<32768>(pp, SB, g_TI.l1[l], g_TI.l2[l], g_TI.l3[l], g_TI.m1[l], g_TI.m2[l], g_TI.m3[l], g_ciI, g_cisI); break;
case 8192: merge4_dit_t<8192>(pp, SB, g_TI.l1[l], g_TI.l2[l], g_TI.l3[l], g_TI.m1[l], g_TI.m2[l], g_TI.m3[l], g_ciI, g_cisI); break;
case 2048: merge4_dit_t<2048>(pp, SB, g_TI.l1[l], g_TI.l2[l], g_TI.l3[l], g_TI.m1[l], g_TI.m2[l], g_TI.m3[l], g_ciI, g_cisI); break;
default: merge4_dit(pp, SB, h, g_TI.l1[l], g_TI.l2[l], g_TI.l3[l], g_TI.m1[l], g_TI.m2[l], g_TI.m3[l], g_ciI, g_cisI);
} }
}
}
// ==================== 1004e7: decimal multiply, n = m = 1e7 digits ====================
struct DI {
unsigned long abi;
const char *s; unsigned long sn;
char *o; unsigned long ol; unsigned long os;
char *e; unsigned long el; unsigned long es;
const char *IB; unsigned long IBl;
char *OB; unsigned long OBl;
unsigned long tsc;
} __attribute__((packed));
static const u32 P1 = 1998585857u, P2 = 2013265921u, P3 = 2088763393u;
#define MAXLIMB 1001024
/* no RED arrays: the parse writes the reduced limbs straight into the NTT arrays */
static u32 RA[NTT_MAXN + 64] __attribute__((aligned(64))), RB[NTT_MAXN + 64] __attribute__((aligned(64))), RC[NTT_MAXN + 64] __attribute__((aligned(64)));
static u64 RO[MAXLIMB * 2 + 8];
static u32 PB_[3][NTT_MAXN] __attribute__((aligned(64))); /* one second-operand NTT array per prime (in-place parse) */
static const char DIG2[201] =
"00010203040506070809101112131415161718192021222324252627282930313233343536373839"
"40414243444546474849505152535455565758596061626364656667686970717273747576777879"
"8081828384858687888990919293949596979899";
// little-endian limbs: out[0] = least significant group of up to 6 digits
// 8 ASCII digits -> numeric value (SWAR, 4 multiplies)
// 8 ASCII digits -> numeric value (SWAR, 4 multiplies)
// AVX2 scan: end of a run of ASCII digits starting at `from`.
TGT static long scan_digits(const char *s, long from, long sn) {
long i = from;
const __m256i lo = _mm256_set1_epi8((char)('0' - 1));
const __m256i hi = _mm256_set1_epi8((char)('9' + 1));
for (; i + 32 <= sn; i += 32) {
__m256i v = LDU(s + i);
__m256i d = _mm256_and_si256(_mm256_cmpgt_epi8(v, lo), _mm256_cmpgt_epi8(hi, v));
unsigned m = (unsigned)_mm256_movemask_epi8(d);
if (m != 0xFFFFFFFFu) return i + __builtin_ctz(~m);
}
while (i < sn && s[i] >= '0' && s[i] <= '9') i++;
return i;
}
static inline u64 swar8(const char *s) {
u64 x;
memcpy(&x, s, 8);
x -= 0x3030303030303030ull;
x = (x * 10 + (x >> 8)) & 0x00FF00FF00FF00FFull;
x = (x * 100 + (x >> 16)) & 0x0000FFFF0000FFFFull;
return (x * 10000 + (x >> 32)) & 0xFFFFFFFFull;
}
static inline u32 ld2d(const char *s) {
u32 v = (u32)(unsigned char)s[0] - '0';
return (u32)((u32)(unsigned char)s[1] - '0') + v * 10u;
}
// base 10^10 limbs, little-endian; emits nlm limbs
TGT static inline __m256i red4(__m256i v, __m256i cV, __m256i pV, __m256i pm1V, __m256i m31V) {
__m256i q = _mm256_srli_epi64(v, 31);
__m256i s = _mm256_and_si256(v, m31V);
__m256i t = _mm256_add_epi64(_mm256_mul_epu32(q, cV), s);
__m256i ge = _mm256_cmpgt_epi64(t, pm1V);
return _mm256_sub_epi64(t, _mm256_and_si256(ge, pV));
}
static inline u32 redsc(u64 v, u32 p, u32 c) {
u64 q = v >> 31, lo = v & 0x7fffffffULL;
u64 t = q * (u64)c + lo;
return (u32)(t >= p ? t - p : t);
}
static u32 parse_limbs6(const char *s, long len, u32 *o1, u32 *o2, u32 *o3) {
const __m256i cV1 = _mm256_set1_epi64x((long long)(u64)(2147483648u - P1));
const __m256i cV2 = _mm256_set1_epi64x((long long)(u64)(2147483648u - P2));
const __m256i cV3 = _mm256_set1_epi64x((long long)(u64)(2147483648u - P3));
const __m256i pV1 = _mm256_set1_epi64x((long long)P1), pV2 = _mm256_set1_epi64x((long long)P2), pV3 = _mm256_set1_epi64x((long long)P3);
const __m256i qV1 = _mm256_set1_epi64x((long long)(P1 - 1)), qV2 = _mm256_set1_epi64x((long long)(P2 - 1)), qV3 = _mm256_set1_epi64x((long long)(P3 - 1));
const __m256i m31V = _mm256_set1_epi64x(0x7fffffffLL);
const __m256i shuf = _mm256_setr_epi32(0, 2, 4, 6, 1, 3, 5, 7);
#define ST3(v) do { u64 _v = (v); o1[k] = redsc(_v, P1, (u32)(2147483648u - P1)); o2[k] = redsc(_v, P2, (u32)(2147483648u - P2)); o3[k] = redsc(_v, P3, (u32)(2147483648u - P3)); } while (0)
long nlm = (len + 9) / 10;
if (nlm <= 0) { o1[0] = 0; o2[0] = 0; o3[0] = 0; return 1; }
int rem = (int)(len - (nlm - 1) * 10);
u32 k = (u32)nlm;
long i;
{ u64 v = 0; for (int j = 0; j < rem; j++) v = v * 10 + (u32)(s[j] - '0'); --k; ST3(v); i = rem; }
/* [C] peel up to 3 limbs so every vector-path store is 16-byte aligned (NT store) */
while (k >= 4u && ((k - 4u) & 3u) != 0u && i + 10 <= len) { --k; ST3(swar8(s + i) * 100ull + ld2d(s + i + 8)); i += 10; }
if (i + 40 <= len) {
const __m256i z0 = _mm256_set1_epi8((char)'0');
const __m256i p1010 = _mm256_set1_epi16(0x010A);
const __m256i p100_1 = _mm256_set1_epi32(0x00010064);
const __m128i z0h = _mm_set1_epi8((char)'0');
const __m128i p1010h = _mm_set1_epi16(0x010A);
const __m128i p100_1h = _mm_set1_epi32(0x00010064);
while (k >= 4 && i + 40 <= len) {
const char *q = s + i;
__m256i v0 = LDU((const __m256i *)q);
__m256i t0 = _mm256_maddubs_epi16(_mm256_sub_epi8(v0, z0), p1010);
/* keep the eight 4-digit groups in a register: a 32-byte store followed by
ten 4-byte loads from the same address is a store-forwarding hazard the
hardware cannot forward, so the loads stall on the store reaching L1. */
__m256i Gv = _mm256_madd_epi16(t0, p100_1);
__m128i v1 = _mm_loadu_si128((const __m128i *)(q + 24));
__m128i t1 = _mm_maddubs_epi16(_mm_sub_epi8(v1, z0h), p1010h);
__m128i Ghv = _mm_madd_epi16(t1, p100_1h);
u32 g0 = (u32)_mm256_extract_epi32(Gv, 0), g1 = (u32)_mm256_extract_epi32(Gv, 1);
u32 g2 = (u32)_mm256_extract_epi32(Gv, 2), g3 = (u32)_mm256_extract_epi32(Gv, 3);
u32 g4 = (u32)_mm256_extract_epi32(Gv, 4), g5 = (u32)_mm256_extract_epi32(Gv, 5);
u32 g6 = (u32)_mm256_extract_epi32(Gv, 6), g7 = (u32)_mm256_extract_epi32(Gv, 7);
u32 g8 = (u32)_mm_extract_epi32(Ghv, 2), g9 = (u32)_mm_extract_epi32(Ghv, 3);
u64 l0 = (u64)(g7 % 100u) * 100000000ull + (u64)g8 * 10000ull + (u64)g9;
u64 l1 = (u64)g5 * 1000000ull + (u64)g6 * 100ull + (u32)(g7 / 100u);
u64 l2 = (u64)(g2 % 100u) * 100000000ull + (u64)g3 * 10000ull + (u64)g4;
u64 l3 = (u64)g0 * 1000000ull + (u64)g1 * 100ull + (u32)(g2 / 100u);
__m256i vv = _mm256_set_epi64x((long long)l3, (long long)l2, (long long)l1, (long long)l0);
_mm_stream_si128((__m128i *)(o1 + k - 4), _mm256_castsi256_si128(_mm256_permutevar8x32_epi32(red4(vv, cV1, pV1, qV1, m31V), shuf)));
_mm_stream_si128((__m128i *)(o2 + k - 4), _mm256_castsi256_si128(_mm256_permutevar8x32_epi32(red4(vv, cV2, pV2, qV2, m31V), shuf)));
_mm_stream_si128((__m128i *)(o3 + k - 4), _mm256_castsi256_si128(_mm256_permutevar8x32_epi32(red4(vv, cV3, pV3, qV3, m31V), shuf)));
k -= 4; i += 40;
}
}
while (i + 10 <= len) { --k; ST3(swar8(s + i) * 100ull + ld2d(s + i + 8)); i += 10; }
while (k > 0) { u64 v = 0; for (long j = i; j < len; j++) v = v * 10 + (u64)(s[j] - '0'); --k; ST3(v); i = len; }
return (u32)nlm;
}
static inline u32 mulconst2(u32 a, u32 w, u64 ws) {
u64 q = ((u64)a * ws) >> 32;
u64 r = (u64)a * w - q * P2;
return (u32)(r >= P2 ? r - P2 : r);
}
static inline u32 mulconst3(u32 a, u32 w, u64 ws) {
u64 q = ((u64)a * ws) >> 32;
u64 r = (u64)a * w - q * P3;
return (u32)(r >= P3 ? r - P3 : r);
}
static inline u32 mulmodp(u32 a, u32 b, u32 p) { return (u32)((u64)a * b % p); }
static u32 powmodp(u32 a, u32 e, u32 p) { u32 r = 1; a %= p; while (e) { if (e & 1) r = mulmodp(r, a, p); a = mulmodp(a, a, p); e >>= 1; } return r; }
#define RED1E(I) { u64 v = in[z+(I)]; u64 q = (v * mn) >> 40; u64 r = v - q * p; \
u64 t = r + p; r = (r > v) ? t : r; t = r - p; r = (r >= p) ? t : r; \
out[z+(I)] = (u32)r; }
// Exact reduction of v < 2^34 mod p (2^30 < p < 2^31) via 2^31 = (2^31-p) mod p.
// v = q*2^31 + s, q = v>>31 <= 4 ; t = q*(2^31-p) + s < 2p ; r = t - (t>=p ? p : 0).
static void redn(const u64 *in, u32 *out, u32 n, u32 p, u32 mn) {
u32 z = 0;
const u64 c64 = (u64)2147483648u - p;
const __m256i cV = _mm256_set1_epi64x((long long)c64);
const __m256i pV = _mm256_set1_epi64x((long long)p);
const __m256i pm1V = _mm256_set1_epi64x((long long)(p - 1));
const __m256i m31V = _mm256_set1_epi64x(0x7fffffffLL);
const __m256i shuf = _mm256_setr_epi32(0, 2, 4, 6, 1, 3, 5, 7);
for (; z + 8 <= n; z += 8) {
__m256i v0 = LDU(in + z);
__m256i v1 = LDU(in + z + 4);
__m256i t0 = _mm256_permutevar8x32_epi32(red4(v0, cV, pV, pm1V, m31V), shuf);
__m256i t1 = _mm256_permutevar8x32_epi32(red4(v1, cV, pV, pm1V, m31V), shuf);
_mm_storeu_si128((__m128i *)(out + z), _mm256_castsi256_si128(t0));
_mm_storeu_si128((__m128i *)(out + z + 4), _mm256_castsi256_si128(t1));
}
for (; z < n; z++) {
u64 v = in[z];
u64 t = ((v >> 31) * c64) + (v & 0x7fffffff);
out[z] = (u32)(t >= p ? t - p : t);
}
}
__attribute__((target("avx2"))) static void mul_run_job(DI *d) {
const char *s = d->s;
long sn = (long)d->sn;
long l1 = scan_digits(s, 0, sn);
long b0 = l1;
while (b0 < sn && (s[b0] < '0' || s[b0] > '9')) b0++;
// The second operand's digit run ends where the trailing non-digits begin, so it
// is reachable in O(trailing non-digits) instead of a second 10 MB forward pass.
// (Valid because the second operand is the last token of this problem's input.)
long e2 = sn;
while (e2 > b0 && (s[e2 - 1] < '0' || s[e2 - 1] > '9')) e2--;
long l2 = e2 - b0;
if (l1 <= 0 || l2 <= 0) return;
u32 *RES[3]; RES[0] = RA; RES[1] = RB; RES[2] = RC;
/* P2: the reduced limbs are written straight into the NTT array of their own prime,
so stage0_dif_src runs in place and the RED round trip (24 MB written + 24 MB read)
is gone. Each prime owns its own arrays, so nothing is overwritten early. */
u32 na = parse_limbs6(s, l1, RES[0], RES[1], RES[2]);
u32 nb = parse_limbs6(s + b0, l2, PB_[0], PB_[1], PB_[2]);
nt_fence();
u32 need = na + nb;
u32 n = 1; while (n < need) n <<= 1;
const u32 PR[3] = {P1, P2, P3};
for (int pr = 0; pr < 3; pr++) {
g_initP = PR[pr];
ntt_init(n);
u32 ninv = powmodp(n % g_P, g_P - 2, g_P);
u32 *A = RES[pr], *B = PB_[pr];
/* stage0 reads exactly one element past the parsed run (index `len`); everything
beyond it is zero-filled by the kernel's own guards. */
A[na] = 0;
B[nb] = 0;
ntt_fwd_dram(A, A, na);
const u32 prescale = mulmodp(ninv, (u32)(((u64)1 << 32) % g_P), g_P);
ntt_fwd_dram(B, B, nb, prescale);
for (u32 b = 0; b < g_n; b += g_Sb) { /* keep A's and B's block hot across the three phases */
ntt_fwd_block(A + b);
ntt_fwd_block(B + b);
ntt_inv_block(A + b, B + b, ninv);
}
ntt_inv_dram(A);
nt_fence();
}
// ---------------- 3-prime CRT, base 10^10, u64-only carry ----------------
const u64 p1 = P1, p2 = P2, p3 = P3;
const u64 M2 = p1 * p2; // < 2^63
u32 i12 = powmodp((u32)(p1 % p2), P2 - 2, P2);
u32 Kinv2 = (u32)((u64)i12 % p2 * (((u64)1 << 32) % p2) % p2);
u32 pinv2 = 0; { u32 iv = 1; for (int q = 0; q < 5; q++) iv *= 2u - P2 * iv; pinv2 = 0u - iv; }
u32 iM2 = powmodp((u32)(M2 % p3), P3 - 2, P3);
u32 iP2 = powmodp(P2, P3 - 2, P3);
const u64 CS3 = ((u64)iM2 << 32) / P3, DS3 = ((u64)iP2 << 32) / P3;
const u64 CS2 = ((u64)i12 << 32) / P2; /* Shoup constant for inv(p1) mod P2 */
u32 Kinv3 = (u32)((u64)iM2 % p3 * (((u64)1 << 32) % p3) % p3);
u32 pinv3 = 0; { u32 iv = 1; for (int q = 0; q < 5; q++) iv *= 2u - P3 * iv; pinv3 = 0u - iv; }
const u64 B10 = 10000000000ull;
const u64 Q64 = 1844674407ull; // floor(2^64 / 10^10)
const u64 R64 = 3709551616ull; // 2^64 mod 10^10
const u64 R34 = 7179869184ull; // 2^34 mod 10^10
u64 carry = 0;
u64 *res = RO;
u32 j = 0, top = 0;
#define CRTEX(J, cLo, cHi) do { \
u64 H34 = (cLo) >> 34, L34 = (cLo) & 0x3FFFFFFFFull; \
u64 S = (cHi) * R64 + H34 * R34 + L34 + carry; \
res[J] = S % B10; \
carry = (cHi) * Q64 + H34 + S / B10; \
} while (0)
/* ---- stage 1: the 32-bit Garner half, 8 limbs per ymm ----
dd2 / u2 / dr31 / q31 / q32 / u3 are all plain 32-bit modular arithmetic, so
they vectorise lane-for-lane. Results go to a small L1-resident scratch that
is filled and drained per chunk, so the extra traffic never leaves L1. */
static u32 gU2[E7R_CH + 8] __attribute__((aligned(64)));
static u32 gU3[E7R_CH + 8] __attribute__((aligned(64)));
{
const __m256i P2v = _mm256_set1_epi32((int)P2);
const __m256i P3v = _mm256_set1_epi32((int)P3);
const __m256i W12 = _mm256_set1_epi32((int)i12);
const __m256i WM2 = _mm256_set1_epi32((int)iM2);
const __m256i WP2 = _mm256_set1_epi32((int)iP2);
const __m256i S12 = _mm256_set1_epi32((int)(u32)CS2);
const __m256i SM2 = _mm256_set1_epi32((int)(u32)CS3);
const __m256i SP2 = _mm256_set1_epi32((int)(u32)DS3);
for (u32 c0 = 0; c0 < need; c0 += E7R_CH) {
u32 n = need - c0; if (n > E7R_CH) n = E7R_CH;
u32 k = 0;
for (; k + 8 <= n; k += 8) {
__m256i r1 = LDU(RA + c0 + k);
__m256i r2 = LDU(RB + c0 + k);
__m256i r3 = LDU(RC + c0 + k);
__m256i t2 = _mm256_sub_epi32(r2, r1);
__m256i dd2 = _mm256_min_epu32(t2, _mm256_add_epi32(t2, P2v));
__m256i u2 = vmulcv(dd2, W12, S12, P2v);
__m256i t3 = _mm256_sub_epi32(r3, r1);
__m256i dr31 = _mm256_min_epu32(t3, _mm256_add_epi32(t3, P3v));
__m256i q31 = vmulcv(dr31, WM2, SM2, P3v);
__m256i q32 = vmulcv(u2, WP2, SP2, P3v);
__m256i d3 = _mm256_sub_epi32(q31, q32);
__m256i u3 = _mm256_min_epu32(d3, _mm256_add_epi32(d3, P3v));
STU(gU2 + k, u2);
STU(gU3 + k, u3);
}
for (; k < n; k++) {
u32 r1 = RA[c0 + k], r2 = RB[c0 + k], r3 = RC[c0 + k];
u32 dd2 = r2 - r1; if (dd2 >= P2) dd2 += P2;
u32 u2 = mulconst2(dd2, i12, CS2);
u32 dr31 = r3 - r1; if (dr31 >= P3) dr31 += P3;
u32 q31 = mulconst3(dr31, iM2, CS3);
u32 q32 = mulconst3(u2, iP2, DS3);
u32 u3 = q31 - q32; if (u3 >= P3) u3 += P3;
gU2[k] = u2; gU3[k] = u3;
}
/* ---- stage 2: 128-bit combine + carry-propagating digit extraction ----
(unchanged arithmetic; the carry chain is inherently serial) */
u32 w = 0;
for (; w + 4 <= n; w += 4) {
#define CRTSTEP(W) do { \
u32 _r1 = RA[c0 + (W)]; \
u64 _c12 = (u64)_r1 + p1 * gU2[W]; \
unsigned long long _cHi; \
u64 _cLo = _mulx_u64((unsigned long long)M2, (unsigned long long)gU3[W], &_cHi); \
u64 _old = _cLo; \
_cLo += _c12; \
_cHi += (_cLo < _old); \
CRTEX(c0 + (W), _cLo, _cHi); \
} while (0)
CRTSTEP(w + 0);
CRTSTEP(w + 1);
CRTSTEP(w + 2);
CRTSTEP(w + 3);
#undef CRTSTEP
}
for (; w < n; w++) {
u32 r1 = RA[c0 + w];
u64 c12 = (u64)r1 + p1 * gU2[w];
unsigned long long cHi;
u64 cLo = _mulx_u64((unsigned long long)M2, (unsigned long long)gU3[w], &cHi);
u64 oldLo = cLo;
cLo += c12;
cHi += (cLo < oldLo);
CRTEX(c0 + w, cLo, cHi);
}
}
j = need;
}
#undef CRTEX
top = need; while (top > 0 && res[top-1] == 0) top--;
while (carry) { res[top++] = carry % B10; carry /= B10; }
if (top == 0) top = 1;
// ---------------- output: 10 digits per limb ----------------
char *o = d->o;
char *olim = d->o + d->ol - 64;
{
u64 v = res[top - 1];
char tmp[24]; int k = 0;
if (v == 0) tmp[k++] = '0';
while (v) { tmp[k++] = (char)('0' + (v % 10)); v /= 10; }
while (k) { if (o > olim) { d->os = (unsigned long)(o - d->o); return; } *o++ = tmp[--k]; }
}
static u32 t2d[100], t4d[10000];
for (u32 j = 0; j < 100; j++) t2d[j] = (u32)('0' + j / 10) | ((u32)('0' + j % 10) << 8);
for (u32 j = 0; j < 10000; j++)
t4d[j] = (u32)('0' + j / 1000) | ((u32)('0' + (j / 100) % 10) << 8)
| ((u32)('0' + (j / 10) % 10) << 16) | ((u32)('0' + j % 10) << 24);
for (int i = (int)top - 2; i >= 0; i--) {
u64 v = res[i];
if (o > olim) { d->os = (unsigned long)(o - d->o); return; }
u64 h10 = v / 100000000ull; /* one u64 division, the rest 32-bit */
u32 lo = (u32)(v - h10 * 100000000ull);
u32 B4 = lo / 10000u;
u32 C4 = lo - B4 * 10000u;
u64 w = (u64)t4d[B4] | ((u64)t4d[C4] << 32);
memcpy(o + 2, &w, 8);
memcpy(o, &t2d[h10], 2);
o += 10;
}
if (o <= olim) *o++ = '\n';
d->os = (unsigned long)(o - d->o);
}
#ifndef NO_LOCAL_TEST
extern "C" void __libc_start_main(void *m, int argc, char **argv) {
(void)m;
unsigned long *p = (unsigned long *)(argv + argc + 1);
while (*p) p++;
p++;
DI *d = 0;
for (int i = 0; i < 32 && p[0]; i++, p += 2)
if (p[0] == 0x6b637564UL) { d = (DI *)p[1]; break; }
if (d) mul_run_job(d);
__asm__ volatile("syscall" ::"a"(60), "D"(0) : "rcx", "r11", "memory");
for (;;);
}
int main() { return 0; }
#endif
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 103.141 ms | 98 MB + 816 KB | Accepted | Score: 100 | 显示更多 |