// Copyright (C) 2026 Kiyotsugu Arai // SPDX-License-Identifier: LGPL-3.0-or-later // PrimeNtt.hpp // Fast multiplication via prime-modular NTT + CRT // Run NTT with three 64-bit NTT-friendly primes // and recover the result via the Chinese Remainder Theorem. // Alternative to Schönhage-Strassen (mod B^F+1) where // pointwise becomes O(1) (64-bit modular multiplication). #pragma once #ifndef SANGI_FORCEINLINE #ifdef _MSC_VER #define SANGI_FORCEINLINE __forceinline #else #define SANGI_FORCEINLINE __attribute__((always_inline)) inline #endif #endif #include #include #include #include #include #include #ifdef _MSC_VER #include #endif #include // AVX2 intrinsics must be included in global scope // (including them inside a namespace causes a GCC error) #if defined(__AVX2__) #include #endif namespace sangi { namespace prime_ntt { // Cache-blocking threshold (block size fitting in L1 32KB) // 512 elements = 4KB data + ~4KB roots = 8KB, with ample headroom in L1 static constexpr size_t NTT_CACHE_BLOCK = 512; // ================================================================ // NTT-friendly prime definitions // ================================================================ // Condition: p = k * 2^s + 1 (larger s permits longer NTT) // p < 2^63 (so 2p < 2^64 under unsigned addition) // Adopt primes proven in FLINT/GMP // p1 = 2^62 - 2^16 + 1 = 4611686018427387905 // = 1 * 2^62 + ... (primality verified separately) // -> Reference: primes used in FLINT's n_mulmod_preinv family // Established NTT-friendly primes (used by FLINT/NTL etc.): // Of form p = k * 2^s + 1 with s >= 47 // Verified NTT-friendly primes (p = k * 2^47 + 1) // Maximum NTT length: 2^47 (= 140 trillion points, effectively unlimited) // p1*p2*p3 ~ 2^163 -> CRT recovers exactly for n < 2^29 (~500M limbs, ~10G digits) // For all primes 15 | (p-1) -> supports mixed-radix NTT (N = {3,5}*2^k) // Selected from verified primes in the BS table constexpr uint64_t P1 = 0x000F'0000'0000'0001ULL; // 30 * 2^47 + 1, g=19 (30=2×3×5) constexpr uint64_t P2 = 0x00AC'8000'0000'0001ULL; // 345 * 2^47 + 1, g=13 (345=3×5×23) constexpr uint64_t P3 = 0x00D2'0000'0000'0001ULL; // 420 * 2^47 + 1, g=17 (420=4×3×5×7) constexpr uint64_t P4 = 0x0032'8000'0000'0001ULL; // 101 * 2^47 + 1, g=3 constexpr uint64_t P5 = 0x0037'8000'0000'0001ULL; // 111 * 2^47 + 1, g=11 constexpr uint64_t G1 = 19; // Primitive root for P1 constexpr uint64_t G2 = 13; // Primitive root for P2 constexpr uint64_t G3 = 17; // Primitive root for P3 constexpr uint64_t G4 = 3; // Primitive root for P4 constexpr uint64_t G5 = 11; // Primitive root for P5 constexpr int NTT_MAX_S = 47; // Maximum NTT length = 2^47 // ================================================================ // mod-p arithmetic primitives // ================================================================ // (a + b) mod p — assumes a, b < p inline uint64_t mod_add(uint64_t a, uint64_t b, uint64_t p) { uint64_t sum = a + b; // If sum >= p (including overflow), subtract p // Since a, b < p < 2^63, sum < 2^64 (no overflow) return (sum >= p) ? (sum - p) : sum; } // (a - b) mod p — assumes a, b < p inline uint64_t mod_sub(uint64_t a, uint64_t b, uint64_t p) { // If a < b, compute a - b + p (underflow correction) return (a >= b) ? (a - b) : (a - b + p); } // (a * b) mod p — assumes a, b < p // 128-bit product -> Barrett reduction inline uint64_t mod_mul(uint64_t a, uint64_t b, uint64_t p) { #if defined(_MSC_VER) && defined(_M_X64) uint64_t hi; uint64_t lo = _umul128(a, b, &hi); uint64_t rem; _udiv128(hi, lo, p, &rem); return rem; #elif defined(__GNUC__) || defined(__clang__) __uint128_t prod = static_cast<__uint128_t>(a) * b; return static_cast(prod % p); #else // Fallback: slow __uint128_t prod = static_cast<__uint128_t>(a) * b; return static_cast(prod % p); #endif } // base^exp mod p (fast modular exponentiation) inline uint64_t mod_pow(uint64_t base, uint64_t exp, uint64_t p) { uint64_t result = 1; base %= p; while (exp > 0) { if (exp & 1) result = mod_mul(result, base, p); base = mod_mul(base, base, p); exp >>= 1; } return result; } // a^(-1) mod p (Fermat's little theorem: a^(p-2) mod p) inline uint64_t mod_inv(uint64_t a, uint64_t p) { return mod_pow(a, p - 2, p); } // ================================================================ // Primality verification and primitive-root search (used during initialization) // ================================================================ // Miller-Rabin primality test (deterministic, 64-bit) inline bool is_prime_64(uint64_t n) { if (n < 2) return false; if (n < 4) return true; if (n % 2 == 0) return false; // n-1 = d * 2^r uint64_t d = n - 1; int r = 0; while ((d & 1) == 0) { d >>= 1; r++; } // Deterministic 64-bit witness set const uint64_t witnesses[] = {2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37}; for (uint64_t a : witnesses) { if (a >= n) continue; uint64_t x = mod_pow(a, d, n); if (x == 1 || x == n - 1) continue; bool composite = true; for (int i = 0; i < r - 1; i++) { x = mod_mul(x, x, n); if (x == n - 1) { composite = false; break; } } if (composite) return false; } return true; } // Search for primes of the form k * 2^s + 1 // min_s: minimum s (guarantees NTT length 2^s) // start_k: starting k for the search (odd) inline uint64_t find_ntt_prime(int min_s, uint64_t start_k = 1) { for (uint64_t k = start_k; k < (1ULL << (63 - min_s)); k += 2) { uint64_t p = k * (1ULL << min_s) + 1; if (p >= (1ULL << 63)) break; // Guarantees 2p < 2^64 if (is_prime_64(p)) return p; } return 0; // Not found } // Search for a primitive root (factorization of p-1 required) // When p-1 = 2^s * k, the prime factors of p-1 are 2 and the factors of k inline uint64_t find_primitive_root(uint64_t p) { // Enumerate the prime factors of p-1 uint64_t pm1 = p - 1; std::vector factors; // 2 is always a factor factors.push_back(2); // Factor the odd part k uint64_t k = pm1; while ((k & 1) == 0) k >>= 1; uint64_t temp = k; for (uint64_t f = 3; f * f <= temp; f += 2) { if (temp % f == 0) { factors.push_back(f); while (temp % f == 0) temp /= f; } } if (temp > 1) factors.push_back(temp); // g^((p-1)/q) != 1 (mod p) for all prime factors q of p-1 for (uint64_t g = 2; g < p; ++g) { bool is_root = true; for (uint64_t q : factors) { if (mod_pow(g, pm1 / q, p) == 1) { is_root = false; break; } } if (is_root) return g; } return 0; } // ================================================================ // Montgomery multiplication (for NTT acceleration) // ================================================================ // R = 2^64, p' = -p^(-1) mod R // REDC(T) = (T + (T_lo * p') * p) >> 64 // Done in 2 x MULX + ADD (_udiv128 35-45 cycles -> ~8 cycles) // Compute p^(-1) mod 2^64 by Newton's method (p is odd) inline uint64_t compute_p_inv_neg(uint64_t p) { // Solve p * x = 1 (mod 2^64) by Newton's method and negate uint64_t x = 1; for (int i = 0; i < 6; i++) // Converges to 64 bits in 6 iterations x *= 2 - p * x; return ~x + 1; // -p^(-1) mod 2^64 } // Montgomery multiplication: REDC(a * b) = a*b*R^(-1) mod p // a, b may be in Montgomery form (a*R mod p, b*R mod p) or as ordinary values inline uint64_t mont_mul(uint64_t a, uint64_t b, uint64_t p, uint64_t p_inv_neg) { #if defined(_MSC_VER) && defined(_M_X64) uint64_t T_hi; uint64_t T_lo = _umul128(a, b, &T_hi); // m = T_lo * (-p^(-1)) mod 2^64 uint64_t m = T_lo * p_inv_neg; // m * p uint64_t mp_hi; uint64_t mp_lo = _umul128(m, p, &mp_hi); // Upper 64 bits of T + m*p (the lower bits are structurally 0) uint64_t carry = (T_lo + mp_lo < T_lo) ? 1 : 0; uint64_t t = T_hi + mp_hi + carry; return (t >= p) ? (t - p) : t; #elif defined(__GNUC__) || defined(__clang__) unsigned __int128 T = (unsigned __int128)a * b; uint64_t T_lo = (uint64_t)T; uint64_t m = T_lo * p_inv_neg; unsigned __int128 mp = (unsigned __int128)m * p; unsigned __int128 sum = T + mp; uint64_t t = (uint64_t)(sum >> 64); return (t >= p) ? (t - p) : t; #endif } // Conversion to Montgomery form: to_mont(a) = a * R mod p = REDC(a * R^2) inline uint64_t to_mont(uint64_t a, uint64_t p, uint64_t p_inv_neg, uint64_t r2_mod_p) { return mont_mul(a, r2_mod_p, p, p_inv_neg); } // Conversion from Montgomery form: from_mont(a_mont) = REDC(a_mont * 1) inline uint64_t from_mont(uint64_t a_mont, uint64_t p, uint64_t p_inv_neg) { return mont_mul(a_mont, 1, p, p_inv_neg); } // ================================================================ // NTT prime parameter struct // ================================================================ struct NttPrime { uint64_t p; // Prime uint64_t g; // Primitive root int s; // The s in p-1 = k * 2^s + ... (maximum NTT length = 2^s) uint64_t inv_2; // 2^(-1) mod p // Montgomery constants uint64_t p_inv_neg; // -p^(-1) mod 2^64 uint64_t r2_mod_p; // R^2 mod p (R = 2^64) // Primitive N-th root for N-point NTT: omega = g^((p-1)/N) mod p uint64_t nth_root(size_t N) const { return mod_pow(g, (p - 1) / N, p); } // N^(-1) mod p uint64_t inv_n(size_t N) const { return mod_inv(N, p); } }; // ================================================================ // Root-table (twiddle factor) cache // ================================================================ struct NttRoots { std::vector roots; // forward: omega^i std::vector inv_roots; // inverse: omega^(-i) uint64_t inv_n; // N^(-1) mod p size_t N; // NTT length uint64_t p; // Prime void build(const NttPrime& prime, size_t ntt_len) { N = ntt_len; p = prime.p; inv_n = prime.inv_n(N); uint64_t omega = prime.nth_root(N); uint64_t omega_inv = mod_inv(omega, p); roots.resize(N / 2); inv_roots.resize(N / 2); roots[0] = 1; inv_roots[0] = 1; for (size_t i = 1; i < N / 2; ++i) { roots[i] = mod_mul(roots[i - 1], omega, p); inv_roots[i] = mod_mul(inv_roots[i - 1], omega_inv, p); } } }; // ================================================================ // NTT forward (DIF: Decimation In Frequency) // ================================================================ inline void forward_ntt(uint64_t* data, size_t N, uint64_t p, const uint64_t* roots) { for (size_t len = N; len >= 2; len >>= 1) { size_t half = len >> 1; size_t step = N / len; // root index step for (size_t start = 0; start < N; start += len) { for (size_t j = 0; j < half; ++j) { uint64_t u = data[start + j]; uint64_t v = data[start + j + half]; // DIF butterfly: (u, v) → (u+v, (u-v)*ω) data[start + j] = mod_add(u, v, p); uint64_t diff = mod_sub(u, v, p); data[start + j + half] = mod_mul(diff, roots[j * step], p); } } } } // ================================================================ // NTT inverse (DIT: Decimation In Time) // ================================================================ inline void inverse_ntt(uint64_t* data, size_t N, uint64_t p, const uint64_t* inv_roots, uint64_t inv_n_val) { for (size_t len = 2; len <= N; len <<= 1) { size_t half = len >> 1; size_t step = N / len; for (size_t start = 0; start < N; start += len) { for (size_t j = 0; j < half; ++j) { uint64_t u = data[start + j]; uint64_t v = mod_mul(data[start + j + half], inv_roots[j * step], p); // DIT butterfly: (u, v·ω⁻¹) → (u+v, u-v) data[start + j] = mod_add(u, v, p); data[start + j + half] = mod_sub(u, v, p); } } } // Normalize by N^(-1) for (size_t i = 0; i < N; ++i) { data[i] = mod_mul(data[i], inv_n_val, p); } } // ================================================================ // Pointwise operations // ================================================================ inline void pointwise_mul(uint64_t* a, const uint64_t* b, size_t N, uint64_t p) { for (size_t i = 0; i < N; ++i) { a[i] = mod_mul(a[i], b[i], p); } } inline void pointwise_sqr(uint64_t* a, size_t N, uint64_t p) { for (size_t i = 0; i < N; ++i) { a[i] = mod_mul(a[i], a[i], p); } } // ================================================================ // Root table and NTT in Montgomery form // ================================================================ struct NttRootsMont { std::vector roots; // forward: omega^i (Montgomery form) std::vector inv_roots; // inverse: omega^(-i) (Montgomery form) uint64_t inv_n; // N^(-1) mod p (normal form — reused for from_mont) size_t N; uint64_t p; uint64_t p_inv_neg; void build(const NttPrime& prime, size_t ntt_len) { N = ntt_len; p = prime.p; p_inv_neg = prime.p_inv_neg; inv_n = prime.inv_n(N); // Normal form (for final scaling in the inverse transform) uint64_t omega = prime.nth_root(N); uint64_t omega_inv = mod_inv(omega, p); // Convert to Montgomery form uint64_t omega_mont = to_mont(omega, p, p_inv_neg, prime.r2_mod_p); uint64_t omega_inv_mont = to_mont(omega_inv, p, p_inv_neg, prime.r2_mod_p); roots.resize(N / 2); inv_roots.resize(N / 2); // roots[0] = Montgomery form of 1 = R mod p uint64_t one_mont = to_mont(1, p, p_inv_neg, prime.r2_mod_p); roots[0] = one_mont; inv_roots[0] = one_mont; for (size_t i = 1; i < N / 2; ++i) { roots[i] = mont_mul(roots[i - 1], omega_mont, p, p_inv_neg); inv_roots[i] = mont_mul(inv_roots[i - 1], omega_inv_mont, p, p_inv_neg); } } }; // Montgomery-form NTT forward (DIF) // data is in Montgomery form; roots are also in Montgomery form inline void forward_ntt_mont(uint64_t* data, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* roots) { for (size_t len = N; len >= 2; len >>= 1) { size_t half = len >> 1; size_t step = N / len; for (size_t start = 0; start < N; start += len) { for (size_t j = 0; j < half; ++j) { uint64_t u = data[start + j]; uint64_t v = data[start + j + half]; data[start + j] = mod_add(u, v, p); uint64_t diff = mod_sub(u, v, p); data[start + j + half] = mont_mul(diff, roots[j * step], p, p_inv_neg); } } } } // Montgomery-form NTT inverse (DIT) // At the end, apply from_mont and inv_n scaling together // inv_n_normal: N^(-1) mod p (normal form) inline void inverse_ntt_mont(uint64_t* data, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* inv_roots, uint64_t inv_n_normal) { for (size_t len = 2; len <= N; len <<= 1) { size_t half = len >> 1; size_t step = N / len; for (size_t start = 0; start < N; start += len) { for (size_t j = 0; j < half; ++j) { uint64_t u = data[start + j]; uint64_t v = mont_mul(data[start + j + half], inv_roots[j * step], p, p_inv_neg); data[start + j] = mod_add(u, v, p); data[start + j + half] = mod_sub(u, v, p); } } } // N^(-1) scaling combined with Montgomery -> normal-form conversion // mont_mul(data[i]_mont, inv_n_normal) = data_i * inv_n mod p for (size_t i = 0; i < N; ++i) { data[i] = mont_mul(data[i], inv_n_normal, p, p_inv_neg); } } // Pointwise multiplication in Montgomery form inline void pointwise_mul_mont(uint64_t* a, const uint64_t* b, size_t N, uint64_t p, uint64_t p_inv_neg) { for (size_t i = 0; i < N; ++i) { a[i] = mont_mul(a[i], b[i], p, p_inv_neg); } } inline void pointwise_sqr_mont(uint64_t* a, size_t N, uint64_t p, uint64_t p_inv_neg) { for (size_t i = 0; i < N; ++i) { a[i] = mont_mul(a[i], a[i], p, p_inv_neg); } } // YC-2: pointwise multiply-accumulate: r[i] += a[i] * b[i] mod p (scalar) inline void pointwise_mac_mont(uint64_t* r, const uint64_t* a, const uint64_t* b, size_t N, uint64_t p, uint64_t p_inv_neg) { for (size_t i = 0; i < N; ++i) { r[i] = mod_add(r[i], mont_mul(a[i], b[i], p, p_inv_neg), p); } } // ================================================================ // AVX2 fast NTT (Montgomery form) // ================================================================ // Accelerate NTT butterflies with 4-wide Montgomery multiplication via _mm256_mul_epu32. // Data remains in Montgomery form. Per-layer root tables // eliminate strided access. #if defined(__AVX2__) // --- AVX2 64-bit multiplication helpers --- // 4-lane 64x64 -> 128-bit multiplication (lo and hi together) // Compute and combine 32-bit partial products with four _mm256_mul_epu32. // Returns correct results for arbitrary uint64_t input (< 2^64). inline void avx2_mul_full_epu64(__m256i a, __m256i b, __m256i& lo, __m256i& hi) { const __m256i mask32 = _mm256_set1_epi64x(0xFFFFFFFF); __m256i a_hi = _mm256_srli_epi64(a, 32); __m256i b_hi = _mm256_srli_epi64(b, 32); __m256i ll = _mm256_mul_epu32(a, b); // a_lo * b_lo __m256i lh = _mm256_mul_epu32(a, b_hi); // a_lo * b_hi __m256i hl = _mm256_mul_epu32(a_hi, b); // a_hi * b_lo __m256i hh = _mm256_mul_epu32(a_hi, b_hi); // a_hi * b_hi // lo = ll + (lh + hl) << 32 [mod 2^64] __m256i cross = _mm256_add_epi64(lh, hl); lo = _mm256_add_epi64(ll, _mm256_slli_epi64(cross, 32)); // hi: propagate carries exactly // mid = lh + (ll >> 32), no carry propagation (lh < 2^64, ll>>32 < 2^32) __m256i ll_hi = _mm256_srli_epi64(ll, 32); __m256i mid = _mm256_add_epi64(lh, ll_hi); __m256i mid_lo = _mm256_and_si256(mid, mask32); __m256i mid_hi = _mm256_srli_epi64(mid, 32); // mid_lo + hl: mid_lo < 2^32, hl < 2^64 → sum < 2^64 + 2^32 < 2^64 (∵ hl ≤ (2^32-1)^2 = 2^64-2^33+1, mid_lo ≤ 2^32-1) __m256i t = _mm256_add_epi64(mid_lo, hl); __m256i t_hi = _mm256_srli_epi64(t, 32); hi = _mm256_add_epi64(hh, _mm256_add_epi64(mid_hi, t_hi)); } // low 64 bits of a*b (3 mul_epu32) inline __m256i avx2_mullo_epu64(__m256i a, __m256i b) { __m256i a_hi = _mm256_srli_epi64(a, 32); __m256i b_hi = _mm256_srli_epi64(b, 32); __m256i ll = _mm256_mul_epu32(a, b); __m256i lh = _mm256_mul_epu32(a, b_hi); __m256i hl = _mm256_mul_epu32(a_hi, b); __m256i cross = _mm256_add_epi64(lh, hl); return _mm256_add_epi64(ll, _mm256_slli_epi64(cross, 32)); } // high 64 bits of a*b (4 mul_epu32) inline __m256i avx2_mulhi_epu64(__m256i a, __m256i b) { const __m256i mask32 = _mm256_set1_epi64x(0xFFFFFFFF); __m256i a_hi = _mm256_srli_epi64(a, 32); __m256i b_hi = _mm256_srli_epi64(b, 32); __m256i ll = _mm256_mul_epu32(a, b); __m256i lh = _mm256_mul_epu32(a, b_hi); __m256i hl = _mm256_mul_epu32(a_hi, b); __m256i hh = _mm256_mul_epu32(a_hi, b_hi); __m256i ll_hi = _mm256_srli_epi64(ll, 32); __m256i mid = _mm256_add_epi64(lh, ll_hi); __m256i mid_lo = _mm256_and_si256(mid, mask32); __m256i mid_hi = _mm256_srli_epi64(mid, 32); __m256i t = _mm256_add_epi64(mid_lo, hl); __m256i t_hi = _mm256_srli_epi64(t, 32); return _mm256_add_epi64(hh, _mm256_add_epi64(mid_hi, t_hi)); } // 4-wide Montgomery multiplication: REDC(a * b) mod p // Input: a, b in Montgomery form (< p) // p_vec = broadcast(p), pin_vec = broadcast(-p^{-1} mod 2^64) // Property used: T_lo + m*p_lo = 0 (mod 2^64) -> carry = (T_lo != 0) inline __m256i avx2_mont_mul(__m256i a, __m256i b, __m256i p_vec, __m256i pin_vec) { // Step 1: T = a * b (128-bit), get lo and hi together __m256i T_lo, T_hi; avx2_mul_full_epu64(a, b, T_lo, T_hi); // 4 VPMULUDQ // Step 2: m = T_lo * (-p^{-1}) mod 2^64 __m256i m = avx2_mullo_epu64(T_lo, pin_vec); // 3 VPMULUDQ // Step 3: mp_hi = mulhi(m, p) __m256i mp_hi = avx2_mulhi_epu64(m, p_vec); // 4 VPMULUDQ // Step 4: carry = (T_lo != 0) ? 1 : 0 __m256i eq_zero = _mm256_cmpeq_epi64(T_lo, _mm256_setzero_si256()); __m256i carry = _mm256_andnot_si256(eq_zero, _mm256_set1_epi64x(1)); // Step 5: t = T_hi + mp_hi + carry __m256i t = _mm256_add_epi64(_mm256_add_epi64(T_hi, mp_hi), carry); // Step 6: if t >= p: t -= p (decide via sign bit) __m256i t_sub_p = _mm256_sub_epi64(t, p_vec); // t < p -> t_sub_p is huge (bit 63 = 1, signed negative) __m256i neg_mask = _mm256_cmpgt_epi64(_mm256_setzero_si256(), t_sub_p); return _mm256_blendv_epi8(t_sub_p, t, neg_mask); } // Total: 11 VPMULUDQ + ~10 add/cmp/blend per 4 elements // --- Per-layer root table --- // Lay out roots of each NTT layer contiguously to enable sequential AVX2 loads. struct NttRootsLayered { std::vector fwd; // forward roots (contiguous per layer) std::vector inv; // inverse roots (contiguous per layer) std::vector fwd_offset; // per-layer offsets within fwd std::vector inv_offset; // per-layer offsets within inv uint64_t inv_n; // N^{-1} mod p (normal form) size_t N = 0; uint64_t p = 0; uint64_t p_inv_neg = 0; void build(const NttPrime& prime, size_t ntt_len) { N = ntt_len; p = prime.p; p_inv_neg = prime.p_inv_neg; inv_n = prime.inv_n(N); uint64_t omega = prime.nth_root(N); uint64_t omega_inv = mod_inv(omega, p); // Convert to Montgomery form uint64_t omega_mont = to_mont(omega, p, p_inv_neg, prime.r2_mod_p); uint64_t omega_inv_mont = to_mont(omega_inv, p, p_inv_neg, prime.r2_mod_p); // Build the ordinary root table std::vector all_roots(N / 2), all_inv(N / 2); uint64_t one_mont = to_mont(1, p, p_inv_neg, prime.r2_mod_p); all_roots[0] = one_mont; all_inv[0] = one_mont; for (size_t i = 1; i < N / 2; ++i) { all_roots[i] = mont_mul(all_roots[i - 1], omega_mont, p, p_inv_neg); all_inv[i] = mont_mul(all_inv[i - 1], omega_inv_mont, p, p_inv_neg); } // Lay out contiguously per layer // Forward DIF: layer l → len = N >> l, half = len/2, step = 1 << l // roots needed: all_roots[j * step] for j = 0..half-1 int num_layers = 0; for (size_t l = N; l >= 2; l >>= 1) ++num_layers; fwd.resize(N - 1); // Sum across layers: N/2 + N/4 + ... + 1 = N - 1 inv.resize(N - 1); fwd_offset.resize(num_layers); inv_offset.resize(num_layers); size_t off = 0; int layer = 0; // Forward (DIF): len = N, N/2, ..., 2 for (size_t len = N; len >= 2; len >>= 1) { size_t half = len >> 1; size_t step = N / len; fwd_offset[layer] = off; for (size_t j = 0; j < half; ++j) fwd[off + j] = all_roots[j * step]; off += half; ++layer; } off = 0; layer = 0; // Inverse (DIT): len = 2, 4, ..., N for (size_t len = 2; len <= N; len <<= 1) { size_t half = len >> 1; size_t step = N / len; inv_offset[layer] = off; for (size_t j = 0; j < half; ++j) inv[off + j] = all_inv[j * step]; off += half; ++layer; } } }; // ================================================================ // Mixed-radix NTT support (N = 3 * M, M = 2^k) // ================================================================ // For all primes 3 | (p-1) -> primitive cube root omega_3 exists // Radix-3 DFT optimized to 2 multiplications: // c₊ = (ω₃+ω₃²)·inv2, c₋ = (ω₃-ω₃²)·inv2 // y₀ = x₀ + s, y₁ = x₀ + c₊·s + c₋·d, y₂ = x₀ + c₊·s - c₋·d // where s = x₁+x₂, d = x₁-x₂ struct MixedRadixConstants { uint64_t c_plus; // (ω₃+ω₃²)·inv2 mod p (Montgomery) uint64_t c_minus; // (ω₃-ω₃²)·inv2 mod p (Montgomery) uint64_t inv3; // 3⁻¹ mod p (Montgomery) std::vector twiddle1; // ω_N^j (j=0..M-1, Montgomery) std::vector twiddle2; // ω_N^{2j} (Montgomery) std::vector inv_twiddle1; // ω_N^{-j} (Montgomery) std::vector inv_twiddle2; // ω_N^{-2j} (Montgomery) size_t N = 0; uint64_t p = 0; void build(const NttPrime& prime, size_t ntt_len) { N = ntt_len; p = prime.p; size_t M = N / 3; uint64_t pinv = prime.p_inv_neg; uint64_t r2 = prime.r2_mod_p; // ω₃ = g^((p-1)/3) — primitive cube root of unity uint64_t w3 = mod_pow(prime.g, (prime.p - 1) / 3, prime.p); uint64_t w3sq = mod_mul(w3, w3, prime.p); // Compute c_plus, c_minus (in normal form, then convert to Montgomery) uint64_t sum_w = mod_add(w3, w3sq, prime.p); uint64_t diff_w = mod_sub(w3, w3sq, prime.p); uint64_t cp = mod_mul(sum_w, prime.inv_2, prime.p); uint64_t cm = mod_mul(diff_w, prime.inv_2, prime.p); c_plus = to_mont(cp, p, pinv, r2); c_minus = to_mont(cm, p, pinv, r2); inv3 = to_mont(mod_inv(3, p), p, pinv, r2); // Twiddle factors: ω_N^j (Montgomery) uint64_t wN = mod_pow(prime.g, (prime.p - 1) / N, prime.p); uint64_t wN_inv = mod_inv(wN, prime.p); uint64_t wN_mont = to_mont(wN, p, pinv, r2); uint64_t wN_inv_mont = to_mont(wN_inv, p, pinv, r2); uint64_t wN2_mont = mont_mul(wN_mont, wN_mont, p, pinv); uint64_t wN_inv2_mont = mont_mul(wN_inv_mont, wN_inv_mont, p, pinv); uint64_t one_mont = to_mont(1, p, pinv, r2); twiddle1.resize(M); twiddle2.resize(M); inv_twiddle1.resize(M); inv_twiddle2.resize(M); twiddle1[0] = twiddle2[0] = inv_twiddle1[0] = inv_twiddle2[0] = one_mont; for (size_t j = 1; j < M; ++j) { twiddle1[j] = mont_mul(twiddle1[j-1], wN_mont, p, pinv); twiddle2[j] = mont_mul(twiddle2[j-1], wN2_mont, p, pinv); inv_twiddle1[j] = mont_mul(inv_twiddle1[j-1], wN_inv_mont, p, pinv); inv_twiddle2[j] = mont_mul(inv_twiddle2[j-1], wN_inv2_mont, p, pinv); } } }; // --- Radix-3 forward DFT (3-point DFT over M columns, 2 mont_mul per column) --- inline void radix3_dft_forward(uint64_t* data, size_t M, uint64_t c_plus, uint64_t c_minus, uint64_t p, uint64_t pinv) { for (size_t j = 0; j < M; ++j) { uint64_t x0 = data[j], x1 = data[j + M], x2 = data[j + 2*M]; uint64_t s = mod_add(x1, x2, p); uint64_t d = mod_sub(x1, x2, p); uint64_t ps = mont_mul(s, c_plus, p, pinv); uint64_t pd = mont_mul(d, c_minus, p, pinv); data[j] = mod_add(x0, s, p); data[j + M] = mod_add(x0, mod_add(ps, pd, p), p); data[j + 2*M] = mod_add(x0, mod_sub(ps, pd, p), p); } } // --- Radix-3 inverse DFT with 1/3 scaling --- // Inverse DFT: negate c_minus (swap omega_3 and omega_3^2) inline void radix3_dft_inverse_scaled(uint64_t* data, size_t M, uint64_t c_plus, uint64_t c_minus, uint64_t inv3, uint64_t p, uint64_t pinv) { uint64_t c_minus_neg = (c_minus == 0) ? 0 : (p - c_minus); // -c₋ mod p for (size_t j = 0; j < M; ++j) { uint64_t y0 = data[j], y1 = data[j + M], y2 = data[j + 2*M]; uint64_t s = mod_add(y1, y2, p); uint64_t d = mod_sub(y1, y2, p); uint64_t ps = mont_mul(s, c_plus, p, pinv); uint64_t pd = mont_mul(d, c_minus_neg, p, pinv); data[j] = mont_mul(mod_add(y0, s, p), inv3, p, pinv); data[j + M] = mont_mul(mod_add(y0, mod_add(ps, pd, p), p), inv3, p, pinv); data[j + 2*M] = mont_mul(mod_add(y0, mod_sub(ps, pd, p), p), inv3, p, pinv); } } // --- Apply twiddle factor --- inline void apply_twiddles(uint64_t* data, size_t M, const uint64_t* tw1, const uint64_t* tw2, uint64_t p, uint64_t pinv) { // Row 0: no twiddle. Row 1: ×ω_N^j. Row 2: ×ω_N^{2j}. for (size_t j = 0; j < M; ++j) { data[M + j] = mont_mul(data[M + j], tw1[j], p, pinv); data[2*M + j] = mont_mul(data[2*M + j], tw2[j], p, pinv); } } #ifdef __AVX2__ // --- AVX2 Radix-3 forward DFT --- inline void radix3_dft_forward_avx2(uint64_t* data, size_t M, uint64_t c_plus_val, uint64_t c_minus_val, uint64_t p, uint64_t pinv) { const __m256i pv = _mm256_set1_epi64x(p); const __m256i pnv = _mm256_set1_epi64x(pinv); const __m256i cpv = _mm256_set1_epi64x(c_plus_val); const __m256i cmv = _mm256_set1_epi64x(c_minus_val); const __m256i zero = _mm256_setzero_si256(); size_t j = 0; for (; j + 4 <= M; j += 4) { __m256i x0 = _mm256_loadu_si256((__m256i*)(data + j)); __m256i x1 = _mm256_loadu_si256((__m256i*)(data + j + M)); __m256i x2 = _mm256_loadu_si256((__m256i*)(data + j + 2*M)); // s = x1 + x2 __m256i s = _mm256_add_epi64(x1, x2); __m256i tmp = _mm256_sub_epi64(s, pv); s = _mm256_blendv_epi8(tmp, s, _mm256_cmpgt_epi64(zero, tmp)); // d = x1 - x2 __m256i d = _mm256_sub_epi64(x1, x2); d = _mm256_blendv_epi8(d, _mm256_add_epi64(d, pv), _mm256_cmpgt_epi64(zero, d)); // ps, pd __m256i ps = avx2_mont_mul(s, cpv, pv, pnv); __m256i pd = avx2_mont_mul(d, cmv, pv, pnv); // y0 = x0 + s __m256i y0 = _mm256_add_epi64(x0, s); tmp = _mm256_sub_epi64(y0, pv); y0 = _mm256_blendv_epi8(tmp, y0, _mm256_cmpgt_epi64(zero, tmp)); // t1 = ps + pd, t2 = ps - pd __m256i t1 = _mm256_add_epi64(ps, pd); tmp = _mm256_sub_epi64(t1, pv); t1 = _mm256_blendv_epi8(tmp, t1, _mm256_cmpgt_epi64(zero, tmp)); __m256i t2 = _mm256_sub_epi64(ps, pd); t2 = _mm256_blendv_epi8(t2, _mm256_add_epi64(t2, pv), _mm256_cmpgt_epi64(zero, t2)); // y1 = x0 + t1, y2 = x0 + t2 __m256i y1 = _mm256_add_epi64(x0, t1); tmp = _mm256_sub_epi64(y1, pv); y1 = _mm256_blendv_epi8(tmp, y1, _mm256_cmpgt_epi64(zero, tmp)); __m256i y2 = _mm256_add_epi64(x0, t2); tmp = _mm256_sub_epi64(y2, pv); y2 = _mm256_blendv_epi8(tmp, y2, _mm256_cmpgt_epi64(zero, tmp)); _mm256_storeu_si256((__m256i*)(data + j), y0); _mm256_storeu_si256((__m256i*)(data + j + M), y1); _mm256_storeu_si256((__m256i*)(data + j + 2*M), y2); } for (; j < M; ++j) { uint64_t x0 = data[j], x1 = data[j+M], x2 = data[j+2*M]; uint64_t sv = mod_add(x1, x2, p), dv = mod_sub(x1, x2, p); uint64_t psv = mont_mul(sv, c_plus_val, p, pinv); uint64_t pdv = mont_mul(dv, c_minus_val, p, pinv); data[j] = mod_add(x0, sv, p); data[j+M] = mod_add(x0, mod_add(psv, pdv, p), p); data[j+2*M] = mod_add(x0, mod_sub(psv, pdv, p), p); } } // --- AVX2 Radix-3 inverse DFT with 1/3 scaling --- inline void radix3_dft_inverse_scaled_avx2(uint64_t* data, size_t M, uint64_t c_plus_val, uint64_t c_minus_val, uint64_t inv3_val, uint64_t p, uint64_t pinv) { uint64_t cm_neg = (c_minus_val == 0) ? 0 : (p - c_minus_val); const __m256i pv = _mm256_set1_epi64x(p); const __m256i pnv = _mm256_set1_epi64x(pinv); const __m256i cpv = _mm256_set1_epi64x(c_plus_val); const __m256i cmv = _mm256_set1_epi64x(cm_neg); const __m256i i3v = _mm256_set1_epi64x(inv3_val); const __m256i zero = _mm256_setzero_si256(); size_t j = 0; for (; j + 4 <= M; j += 4) { __m256i y0 = _mm256_loadu_si256((__m256i*)(data + j)); __m256i y1 = _mm256_loadu_si256((__m256i*)(data + j + M)); __m256i y2 = _mm256_loadu_si256((__m256i*)(data + j + 2*M)); __m256i s = _mm256_add_epi64(y1, y2); __m256i tmp = _mm256_sub_epi64(s, pv); s = _mm256_blendv_epi8(tmp, s, _mm256_cmpgt_epi64(zero, tmp)); __m256i d = _mm256_sub_epi64(y1, y2); d = _mm256_blendv_epi8(d, _mm256_add_epi64(d, pv), _mm256_cmpgt_epi64(zero, d)); __m256i ps = avx2_mont_mul(s, cpv, pv, pnv); __m256i pd = avx2_mont_mul(d, cmv, pv, pnv); // z0 = y0+s, z1 = y0+ps+pd, z2 = y0+ps-pd __m256i z0 = _mm256_add_epi64(y0, s); tmp = _mm256_sub_epi64(z0, pv); z0 = _mm256_blendv_epi8(tmp, z0, _mm256_cmpgt_epi64(zero, tmp)); __m256i t1 = _mm256_add_epi64(ps, pd); tmp = _mm256_sub_epi64(t1, pv); t1 = _mm256_blendv_epi8(tmp, t1, _mm256_cmpgt_epi64(zero, tmp)); __m256i t2 = _mm256_sub_epi64(ps, pd); t2 = _mm256_blendv_epi8(t2, _mm256_add_epi64(t2, pv), _mm256_cmpgt_epi64(zero, t2)); __m256i z1 = _mm256_add_epi64(y0, t1); tmp = _mm256_sub_epi64(z1, pv); z1 = _mm256_blendv_epi8(tmp, z1, _mm256_cmpgt_epi64(zero, tmp)); __m256i z2 = _mm256_add_epi64(y0, t2); tmp = _mm256_sub_epi64(z2, pv); z2 = _mm256_blendv_epi8(tmp, z2, _mm256_cmpgt_epi64(zero, tmp)); // scale by inv3 _mm256_storeu_si256((__m256i*)(data + j), avx2_mont_mul(z0, i3v, pv, pnv)); _mm256_storeu_si256((__m256i*)(data + j + M), avx2_mont_mul(z1, i3v, pv, pnv)); _mm256_storeu_si256((__m256i*)(data + j + 2*M), avx2_mont_mul(z2, i3v, pv, pnv)); } uint64_t cm_neg_s = cm_neg; for (; j < M; ++j) { uint64_t y0v = data[j], y1v = data[j+M], y2v = data[j+2*M]; uint64_t sv = mod_add(y1v, y2v, p), dv = mod_sub(y1v, y2v, p); uint64_t psv = mont_mul(sv, c_plus_val, p, pinv); uint64_t pdv = mont_mul(dv, cm_neg_s, p, pinv); data[j] = mont_mul(mod_add(y0v, sv, p), inv3_val, p, pinv); data[j+M] = mont_mul(mod_add(y0v, mod_add(psv, pdv, p), p), inv3_val, p, pinv); data[j+2*M] = mont_mul(mod_add(y0v, mod_sub(psv, pdv, p), p), inv3_val, p, pinv); } } // --- Apply AVX2 twiddle --- inline void apply_twiddles_avx2(uint64_t* data, size_t M, const uint64_t* tw1, const uint64_t* tw2, uint64_t p, uint64_t pinv) { const __m256i pv = _mm256_set1_epi64x(p); const __m256i pnv = _mm256_set1_epi64x(pinv); size_t j = 0; for (; j + 4 <= M; j += 4) { __m256i d1 = _mm256_loadu_si256((__m256i*)(data + M + j)); __m256i w1 = _mm256_loadu_si256((__m256i*)(tw1 + j)); _mm256_storeu_si256((__m256i*)(data + M + j), avx2_mont_mul(d1, w1, pv, pnv)); __m256i d2 = _mm256_loadu_si256((__m256i*)(data + 2*M + j)); __m256i w2 = _mm256_loadu_si256((__m256i*)(tw2 + j)); _mm256_storeu_si256((__m256i*)(data + 2*M + j), avx2_mont_mul(d2, w2, pv, pnv)); } for (; j < M; ++j) { data[M + j] = mont_mul(data[M + j], tw1[j], p, pinv); data[2*M + j] = mont_mul(data[2*M + j], tw2[j], p, pinv); } } // ================================================================ // Radix-5 mixed-radix NTT (N = 5M, M = 2^k) // ================================================================ // Compute 5-point DFT with 8 mont_mul: // a' = (w5+w5⁴)/2, b' = (w5²+w5³)/2 // c' = (w5-w5⁴)/2, e' = (w5²-w5³)/2 struct MixedRadix5Constants { uint64_t a_prime, b_prime, c_prime, e_prime; // Montgomery uint64_t a_plus_b; // (a'+b') mod p — for Karatsuba uint64_t c_plus_e; // (c'+e') mod p — for Karatsuba uint64_t inv5; // 5⁻¹ mod p (Montgomery) std::vector twiddle[4]; // ω_N^{rj} r=1..4 (Montgomery) std::vector inv_twiddle[4]; // ω_N^{-rj} size_t N = 0; uint64_t p = 0; void build(const NttPrime& prime, size_t ntt_len) { N = ntt_len; p = prime.p; size_t M = N / 5; uint64_t pinv = prime.p_inv_neg; uint64_t r2 = prime.r2_mod_p; uint64_t w5 = mod_pow(prime.g, (prime.p - 1) / 5, prime.p); uint64_t w5_2 = mod_mul(w5, w5, prime.p); uint64_t w5_3 = mod_mul(w5_2, w5, prime.p); uint64_t w5_4 = mod_mul(w5_3, w5, prime.p); uint64_t inv2 = prime.inv_2; a_prime = to_mont(mod_mul(mod_add(w5, w5_4, prime.p), inv2, prime.p), p, pinv, r2); b_prime = to_mont(mod_mul(mod_add(w5_2, w5_3, prime.p), inv2, prime.p), p, pinv, r2); c_prime = to_mont(mod_mul(mod_sub(w5, w5_4, prime.p), inv2, prime.p), p, pinv, r2); e_prime = to_mont(mod_mul(mod_sub(w5_2, w5_3, prime.p), inv2, prime.p), p, pinv, r2); // Karatsuba constants (mod_add in Montgomery form is correct — linear) a_plus_b = mod_add(a_prime, b_prime, p); c_plus_e = mod_add(c_prime, e_prime, p); inv5 = to_mont(mod_inv(5, p), p, pinv, r2); uint64_t wN = mod_pow(prime.g, (prime.p - 1) / N, prime.p); uint64_t wN_inv = mod_inv(wN, prime.p); uint64_t one_mont = to_mont(1, p, pinv, r2); // twiddle[r-1][j] = ω_N^{rj} for r=1..4, j=0..M-1 uint64_t wN_r_mont[4], wN_inv_r_mont[4]; uint64_t wr = wN, wri = wN_inv; for (int r = 0; r < 4; ++r) { wN_r_mont[r] = to_mont(wr, p, pinv, r2); wN_inv_r_mont[r] = to_mont(wri, p, pinv, r2); wr = mod_mul(wr, wN, prime.p); wri = mod_mul(wri, wN_inv, prime.p); twiddle[r].resize(M); inv_twiddle[r].resize(M); twiddle[r][0] = inv_twiddle[r][0] = one_mont; } for (size_t j = 1; j < M; ++j) { for (int r = 0; r < 4; ++r) { twiddle[r][j] = mont_mul(twiddle[r][j-1], wN_r_mont[r], p, pinv); inv_twiddle[r][j] = mont_mul(inv_twiddle[r][j-1], wN_inv_r_mont[r], p, pinv); } } } }; // --- Radix-5 forward DFT (M columns, 6 mont_mul per column — Karatsuba) --- inline void radix5_dft_forward(uint64_t* data, size_t M, uint64_t ap, uint64_t bp, uint64_t cp, uint64_t ep, uint64_t a_plus_b, uint64_t c_plus_e, uint64_t p, uint64_t pinv) { for (size_t j = 0; j < M; ++j) { uint64_t x0=data[j], x1=data[j+M], x2=data[j+2*M], x3=data[j+3*M], x4=data[j+4*M]; uint64_t s1 = mod_add(x1, x4, p), d1 = mod_sub(x1, x4, p); uint64_t s2 = mod_add(x2, x3, p), d2 = mod_sub(x2, x3, p); // Karatsuba: 6 mul instead of 8 uint64_t P1 = mont_mul(s1, ap, p, pinv); uint64_t P2 = mont_mul(s2, bp, p, pinv); uint64_t P3 = mont_mul(mod_add(s1, s2, p), a_plus_b, p, pinv); uint64_t t14 = mod_add(P1, P2, p); uint64_t t23 = mod_sub(P3, t14, p); // (a+b)(s1+s2) - a·s1 - b·s2 = b·s1 + a·s2 uint64_t Q1 = mont_mul(d1, cp, p, pinv); uint64_t Q2 = mont_mul(d2, ep, p, pinv); uint64_t Q3 = mont_mul(mod_sub(d1, d2, p), c_plus_e, p, pinv); uint64_t u14 = mod_add(Q1, Q2, p); uint64_t u23 = mod_add(mod_sub(Q3, Q1, p), Q2, p); // (c+e)(d1-d2) - c·d1 + e·d2 = e·d1 - c·d2 data[j] = mod_add(x0, mod_add(s1, s2, p), p); data[j+M] = mod_add(x0, mod_add(t14, u14, p), p); data[j+4*M] = mod_add(x0, mod_sub(t14, u14, p), p); data[j+2*M] = mod_add(x0, mod_add(t23, u23, p), p); data[j+3*M] = mod_add(x0, mod_sub(t23, u23, p), p); } } // --- Radix-5 inverse DFT with 1/5 scaling (6 mul + 5 inv5 scaling) --- inline void radix5_dft_inverse_scaled(uint64_t* data, size_t M, uint64_t ap, uint64_t bp, uint64_t cp, uint64_t ep, uint64_t a_plus_b, uint64_t c_plus_e, uint64_t inv5, uint64_t p, uint64_t pinv) { uint64_t cp_neg = cp ? (p - cp) : 0; uint64_t ep_neg = ep ? (p - ep) : 0; uint64_t ce_neg = c_plus_e ? (p - c_plus_e) : 0; // -(c+e) mod p for (size_t j = 0; j < M; ++j) { uint64_t y0=data[j], y1=data[j+M], y2=data[j+2*M], y3=data[j+3*M], y4=data[j+4*M]; uint64_t s1 = mod_add(y1, y4, p), d1 = mod_sub(y1, y4, p); uint64_t s2 = mod_add(y2, y3, p), d2 = mod_sub(y2, y3, p); // t terms: same Karatsuba (a,b unchanged in inverse) uint64_t P1 = mont_mul(s1, ap, p, pinv); uint64_t P2 = mont_mul(s2, bp, p, pinv); uint64_t P3 = mont_mul(mod_add(s1, s2, p), a_plus_b, p, pinv); uint64_t t14 = mod_add(P1, P2, p); uint64_t t23 = mod_sub(P3, t14, p); // u terms: c→-c, e→-e, so (c+e)→-(c+e) uint64_t Q1 = mont_mul(d1, cp_neg, p, pinv); uint64_t Q2 = mont_mul(d2, ep_neg, p, pinv); uint64_t Q3 = mont_mul(mod_sub(d1, d2, p), ce_neg, p, pinv); uint64_t u14 = mod_add(Q1, Q2, p); uint64_t u23 = mod_add(mod_sub(Q3, Q1, p), Q2, p); data[j] = mont_mul(mod_add(y0, mod_add(s1, s2, p), p), inv5, p, pinv); data[j+M] = mont_mul(mod_add(y0, mod_add(t14, u14, p), p), inv5, p, pinv); data[j+4*M] = mont_mul(mod_add(y0, mod_sub(t14, u14, p), p), inv5, p, pinv); data[j+2*M] = mont_mul(mod_add(y0, mod_add(t23, u23, p), p), inv5, p, pinv); data[j+3*M] = mont_mul(mod_add(y0, mod_sub(t23, u23, p), p), inv5, p, pinv); } } // --- Radix-5 twiddle (rows 1-4) --- inline void apply_twiddles5(uint64_t* data, size_t M, const uint64_t* tw[4], uint64_t p, uint64_t pinv) { for (size_t j = 0; j < M; ++j) { for (int r = 0; r < 4; ++r) data[(r+1)*M + j] = mont_mul(data[(r+1)*M + j], tw[r][j], p, pinv); } } // --- AVX2 Radix-5 forward DFT (6 mul per column — Karatsuba) --- inline void radix5_dft_forward_avx2(uint64_t* data, size_t M, uint64_t ap, uint64_t bp, uint64_t cp, uint64_t ep, uint64_t a_plus_b, uint64_t c_plus_e, uint64_t p, uint64_t pinv) { const __m256i pv = _mm256_set1_epi64x(p), pnv = _mm256_set1_epi64x(pinv); const __m256i av = _mm256_set1_epi64x(ap), bv = _mm256_set1_epi64x(bp); const __m256i cv = _mm256_set1_epi64x(cp), ev = _mm256_set1_epi64x(ep); const __m256i abv = _mm256_set1_epi64x(a_plus_b); const __m256i cev = _mm256_set1_epi64x(c_plus_e); const __m256i zero = _mm256_setzero_si256(); auto madd = [&](__m256i a, __m256i b) -> __m256i { __m256i s = _mm256_add_epi64(a, b); __m256i t = _mm256_sub_epi64(s, pv); return _mm256_blendv_epi8(t, s, _mm256_cmpgt_epi64(zero, t)); }; auto msub = [&](__m256i a, __m256i b) -> __m256i { __m256i d = _mm256_sub_epi64(a, b); return _mm256_blendv_epi8(d, _mm256_add_epi64(d, pv), _mm256_cmpgt_epi64(zero, d)); }; auto mmul = [&](__m256i a, __m256i b) -> __m256i { return avx2_mont_mul(a, b, pv, pnv); }; size_t j = 0; for (; j + 4 <= M; j += 4) { __m256i x0=_mm256_loadu_si256((__m256i*)(data+j)); __m256i x1=_mm256_loadu_si256((__m256i*)(data+j+M)); __m256i x2=_mm256_loadu_si256((__m256i*)(data+j+2*M)); __m256i x3=_mm256_loadu_si256((__m256i*)(data+j+3*M)); __m256i x4=_mm256_loadu_si256((__m256i*)(data+j+4*M)); __m256i s1=madd(x1,x4), d1=msub(x1,x4), s2=madd(x2,x3), d2=msub(x2,x3); // Karatsuba: 6 mmul instead of 8 __m256i P1=mmul(s1,av), P2=mmul(s2,bv), P3=mmul(madd(s1,s2),abv); __m256i t14=madd(P1,P2), t23=msub(P3,t14); __m256i Q1=mmul(d1,cv), Q2=mmul(d2,ev), Q3=mmul(msub(d1,d2),cev); __m256i u14=madd(Q1,Q2), u23=madd(msub(Q3,Q1),Q2); _mm256_storeu_si256((__m256i*)(data+j), madd(x0, madd(s1,s2))); _mm256_storeu_si256((__m256i*)(data+j+M), madd(x0, madd(t14,u14))); _mm256_storeu_si256((__m256i*)(data+j+4*M), madd(x0, msub(t14,u14))); _mm256_storeu_si256((__m256i*)(data+j+2*M), madd(x0, madd(t23,u23))); _mm256_storeu_si256((__m256i*)(data+j+3*M), madd(x0, msub(t23,u23))); } for (; j < M; ++j) { uint64_t x0=data[j],x1=data[j+M],x2=data[j+2*M],x3=data[j+3*M],x4=data[j+4*M]; uint64_t s1=mod_add(x1,x4,p),d1=mod_sub(x1,x4,p),s2=mod_add(x2,x3,p),d2=mod_sub(x2,x3,p); uint64_t P1v=mont_mul(s1,ap,p,pinv), P2v=mont_mul(s2,bp,p,pinv); uint64_t P3v=mont_mul(mod_add(s1,s2,p),a_plus_b,p,pinv); uint64_t t14v=mod_add(P1v,P2v,p), t23v=mod_sub(P3v,t14v,p); uint64_t Q1v=mont_mul(d1,cp,p,pinv), Q2v=mont_mul(d2,ep,p,pinv); uint64_t Q3v=mont_mul(mod_sub(d1,d2,p),c_plus_e,p,pinv); uint64_t u14v=mod_add(Q1v,Q2v,p), u23v=mod_add(mod_sub(Q3v,Q1v,p),Q2v,p); data[j]=mod_add(x0,mod_add(s1,s2,p),p); data[j+M]=mod_add(x0,mod_add(t14v,u14v,p),p); data[j+4*M]=mod_add(x0,mod_sub(t14v,u14v,p),p); data[j+2*M]=mod_add(x0,mod_add(t23v,u23v,p),p); data[j+3*M]=mod_add(x0,mod_sub(t23v,u23v,p),p); } } // --- AVX2 Radix-5 inverse DFT with 1/5 scaling (6 mul + 5 inv5 — Karatsuba) --- inline void radix5_dft_inverse_scaled_avx2(uint64_t* data, size_t M, uint64_t ap, uint64_t bp, uint64_t cp, uint64_t ep, uint64_t a_plus_b, uint64_t c_plus_e, uint64_t inv5v, uint64_t p, uint64_t pinv) { uint64_t cpn = cp ? (p-cp) : 0, epn = ep ? (p-ep) : 0; uint64_t cen = c_plus_e ? (p-c_plus_e) : 0; const __m256i pv=_mm256_set1_epi64x(p), pnv=_mm256_set1_epi64x(pinv); const __m256i av=_mm256_set1_epi64x(ap),bvv=_mm256_set1_epi64x(bp); const __m256i cvv=_mm256_set1_epi64x(cpn),evv=_mm256_set1_epi64x(epn); const __m256i abv=_mm256_set1_epi64x(a_plus_b); const __m256i cenv=_mm256_set1_epi64x(cen); const __m256i i5=_mm256_set1_epi64x(inv5v); const __m256i zero=_mm256_setzero_si256(); auto madd=[&](__m256i a,__m256i b){__m256i s=_mm256_add_epi64(a,b);__m256i t=_mm256_sub_epi64(s,pv);return _mm256_blendv_epi8(t,s,_mm256_cmpgt_epi64(zero,t));}; auto msub=[&](__m256i a,__m256i b){__m256i d=_mm256_sub_epi64(a,b);return _mm256_blendv_epi8(d,_mm256_add_epi64(d,pv),_mm256_cmpgt_epi64(zero,d));}; auto mmul=[&](__m256i a,__m256i b){return avx2_mont_mul(a,b,pv,pnv);}; size_t j=0; for(;j+4<=M;j+=4){ __m256i y0=_mm256_loadu_si256((__m256i*)(data+j)); __m256i y1=_mm256_loadu_si256((__m256i*)(data+j+M)); __m256i y2=_mm256_loadu_si256((__m256i*)(data+j+2*M)); __m256i y3=_mm256_loadu_si256((__m256i*)(data+j+3*M)); __m256i y4=_mm256_loadu_si256((__m256i*)(data+j+4*M)); __m256i s1=madd(y1,y4),d1=msub(y1,y4),s2=madd(y2,y3),d2=msub(y2,y3); // Karatsuba: t terms (a,b unchanged) __m256i P1=mmul(s1,av), P2=mmul(s2,bvv), P3=mmul(madd(s1,s2),abv); __m256i t14=madd(P1,P2), t23=msub(P3,t14); // Karatsuba: u terms (c_neg, e_neg, ce_neg) __m256i Q1=mmul(d1,cvv), Q2=mmul(d2,evv), Q3=mmul(msub(d1,d2),cenv); __m256i u14=madd(Q1,Q2), u23=madd(msub(Q3,Q1),Q2); _mm256_storeu_si256((__m256i*)(data+j),mmul(madd(y0,madd(s1,s2)),i5)); _mm256_storeu_si256((__m256i*)(data+j+M), mmul(madd(y0,madd(t14,u14)),i5)); _mm256_storeu_si256((__m256i*)(data+j+4*M),mmul(madd(y0,msub(t14,u14)),i5)); _mm256_storeu_si256((__m256i*)(data+j+2*M),mmul(madd(y0,madd(t23,u23)),i5)); _mm256_storeu_si256((__m256i*)(data+j+3*M),mmul(madd(y0,msub(t23,u23)),i5)); } for(;j block) — standard layer-by-layer int layer = 0; for (size_t len = N; len > block; len >>= 1) { size_t half = len >> 1; const uint64_t* roots = layer_roots + layer_offsets[layer]; for (size_t start = 0; start < N; start += len) dif_butterfly_avx2(data, start, half, roots, p_vec, pin_vec, p, p_inv_neg); ++layer; } // Phase 2: bottom layers (len <= block) — per-block processing within L1 int bottom_start = layer; for (size_t blk = 0; blk < N; blk += block) { int cur_layer = bottom_start; for (size_t len = block; len >= 2; len >>= 1) { size_t half = len >> 1; const uint64_t* roots = layer_roots + layer_offsets[cur_layer]; for (size_t start = blk; start < blk + block; start += len) dif_butterfly_avx2(data, start, half, roots, p_vec, pin_vec, p, p_inv_neg); ++cur_layer; } } } // --- AVX2 DIT butterfly (one group) --- SANGI_FORCEINLINE void dit_butterfly_avx2(uint64_t* data, size_t start, size_t half, const uint64_t* roots, __m256i p_vec, __m256i pin_vec, uint64_t p, uint64_t p_inv_neg) { uint64_t* d0 = data + start; uint64_t* d1 = data + start + half; size_t j = 0; for (; j + 4 <= half; j += 4) { __m256i u = _mm256_loadu_si256((__m256i*)(d0 + j)); __m256i v_raw = _mm256_loadu_si256((__m256i*)(d1 + j)); __m256i w = _mm256_loadu_si256((__m256i*)(roots + j)); // DIT: v = mont_mul(v_raw, ω⁻¹), then (u+v, u-v) __m256i v = avx2_mont_mul(v_raw, w, p_vec, pin_vec); __m256i sum = _mm256_add_epi64(u, v); __m256i sum_sub = _mm256_sub_epi64(sum, p_vec); __m256i sm = _mm256_cmpgt_epi64(_mm256_setzero_si256(), sum_sub); sum = _mm256_blendv_epi8(sum_sub, sum, sm); __m256i diff = _mm256_sub_epi64(u, v); __m256i diff_add = _mm256_add_epi64(diff, p_vec); __m256i dm = _mm256_cmpgt_epi64(_mm256_setzero_si256(), diff); diff = _mm256_blendv_epi8(diff, diff_add, dm); _mm256_storeu_si256((__m256i*)(d0 + j), sum); _mm256_storeu_si256((__m256i*)(d1 + j), diff); } for (; j < half; ++j) { uint64_t u = d0[j]; uint64_t v = mont_mul(d1[j], roots[j], p, p_inv_neg); d0[j] = mod_add(u, v, p); d1[j] = mod_sub(u, v, p); } } // --- AVX2 NTT inverse (DIT, Montgomery, cache-blocked) --- inline void inverse_ntt_mont_avx2(uint64_t* data, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* layer_roots, const size_t* layer_offsets, uint64_t inv_n_normal) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); size_t block = (N <= NTT_CACHE_BLOCK) ? N : NTT_CACHE_BLOCK; // Phase 1: bottom layers (len <= block) — per-block processing within L1 int num_bottom_layers = 0; for (size_t l = 2; l <= block; l <<= 1) ++num_bottom_layers; for (size_t blk = 0; blk < N; blk += block) { int cur_layer = 0; for (size_t len = 2; len <= block; len <<= 1) { size_t half = len >> 1; const uint64_t* roots = layer_roots + layer_offsets[cur_layer]; for (size_t start = blk; start < blk + block; start += len) dit_butterfly_avx2(data, start, half, roots, p_vec, pin_vec, p, p_inv_neg); ++cur_layer; } } // Phase 2: top layers (len > block) — standard layer-by-layer int layer = num_bottom_layers; for (size_t len = block * 2; len <= N; len <<= 1) { size_t half = len >> 1; const uint64_t* roots = layer_roots + layer_offsets[layer]; for (size_t start = 0; start < N; start += len) dit_butterfly_avx2(data, start, half, roots, p_vec, pin_vec, p, p_inv_neg); ++layer; } // N^{-1} scaling + from_mont (AVX2) __m256i inv_n_vec = _mm256_set1_epi64x(inv_n_normal); size_t i = 0; for (; i + 4 <= N; i += 4) { __m256i d = _mm256_loadu_si256((__m256i*)(data + i)); d = avx2_mont_mul(d, inv_n_vec, p_vec, pin_vec); _mm256_storeu_si256((__m256i*)(data + i), d); } for (; i < N; ++i) { data[i] = mont_mul(data[i], inv_n_normal, p, p_inv_neg); } } // --- AVX2 pointwise multiplication --- inline void pointwise_mul_mont_avx2(uint64_t* a, const uint64_t* b, size_t N, uint64_t p, uint64_t p_inv_neg) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); size_t i = 0; for (; i + 4 <= N; i += 4) { __m256i va = _mm256_loadu_si256((__m256i*)(a + i)); __m256i vb = _mm256_loadu_si256((__m256i*)(b + i)); va = avx2_mont_mul(va, vb, p_vec, pin_vec); _mm256_storeu_si256((__m256i*)(a + i), va); } for (; i < N; ++i) { a[i] = mont_mul(a[i], b[i], p, p_inv_neg); } } inline void pointwise_sqr_mont_avx2(uint64_t* a, size_t N, uint64_t p, uint64_t p_inv_neg) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); size_t i = 0; for (; i + 4 <= N; i += 4) { __m256i va = _mm256_loadu_si256((__m256i*)(a + i)); va = avx2_mont_mul(va, va, p_vec, pin_vec); _mm256_storeu_si256((__m256i*)(a + i), va); } for (; i < N; ++i) { a[i] = mont_mul(a[i], a[i], p, p_inv_neg); } } // YC-2: pointwise multiply-accumulate: r[i] += a[i] * b[i] mod p (AVX2) inline void pointwise_mac_mont_avx2(uint64_t* r, const uint64_t* a, const uint64_t* b, size_t N, uint64_t p, uint64_t p_inv_neg) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); size_t i = 0; for (; i + 4 <= N; i += 4) { __m256i va = _mm256_loadu_si256((__m256i*)(a + i)); __m256i vb = _mm256_loadu_si256((__m256i*)(b + i)); __m256i vr = _mm256_loadu_si256((__m256i*)(r + i)); __m256i prod = avx2_mont_mul(va, vb, p_vec, pin_vec); // mod_add: sum = vr + prod, if sum >= p then sum -= p __m256i sum = _mm256_add_epi64(vr, prod); __m256i sub = _mm256_sub_epi64(sum, p_vec); __m256i mask = _mm256_cmpgt_epi64(_mm256_setzero_si256(), sub); sum = _mm256_blendv_epi8(sub, sum, mask); _mm256_storeu_si256((__m256i*)(r + i), sum); } for (; i < N; ++i) { r[i] = mod_add(r[i], mont_mul(a[i], b[i], p, p_inv_neg), p); } } #endif // __AVX2__ // ================================================================ // CRT reconstruction (Garner's algorithm) // ================================================================ // Recover the true coefficients from the NTT results across 3 primes, // and convert to a limb sequence (uint64_t). // r1[i] mod p1, r2[i] mod p2, r3[i] mod p3 → result limbs // // Garner: // v1 = r1[i] // v2 = (r2[i] - v1) * inv_p1_mod_p2 mod p2 // v3 = (r3[i] - v1 - v2*p1) * inv_p1p2_mod_p3 mod p3 // c[i] = v1 + v2*p1 + v3*p1*p2 (up to ~189 bits) struct CrtConstants { uint64_t p1, p2, p3; uint64_t inv_p1_mod_p2; // p1^(-1) mod p2 uint64_t inv_p1p2_mod_p3; // (p1*p2)^(-1) mod p3 uint64_t p1_mod_p3; // p1 mod p3 void init(uint64_t _p1, uint64_t _p2, uint64_t _p3) { p1 = _p1; p2 = _p2; p3 = _p3; inv_p1_mod_p2 = mod_inv(p1 % p2, p2); // p1*p2 mod p3: (p1 mod p3) * (p2 mod p3) mod p3 p1_mod_p3 = p1 % p3; uint64_t p1p2_mod_p3 = mod_mul(p1_mod_p3, p2 % p3, p3); inv_p1p2_mod_p3 = mod_inv(p1p2_mod_p3, p3); } }; // Reconstruct one coefficient via CRT and return a 3-word (192-bit) value // result[0] = low 64bit, result[1] = mid 64bit, result[2] = high 64bit inline void crt_single(uint64_t result[3], uint64_t r1, uint64_t r2, uint64_t r3, const CrtConstants& crt) { // v1 = r1 uint64_t v1 = r1; // v2 = (r2 - v1) * inv_p1_mod_p2 mod p2 uint64_t t = mod_sub(r2, v1 % crt.p2, crt.p2); uint64_t v2 = mod_mul(t, crt.inv_p1_mod_p2, crt.p2); // v3 = (r3 - v1 - v2*p1_mod_p3) * inv_p1p2_mod_p3 mod p3 uint64_t v2p1_mod_p3 = mod_mul(v2, crt.p1_mod_p3, crt.p3); uint64_t s = mod_sub(r3, v1 % crt.p3, crt.p3); s = mod_sub(s, v2p1_mod_p3, crt.p3); uint64_t v3 = mod_mul(s, crt.inv_p1p2_mod_p3, crt.p3); // c = v1 + v2*p1 + v3*p1*p2 // v2*p1: at most 63-bit x 63-bit = 126-bit -> 2 words #if defined(_MSC_VER) && defined(_M_X64) uint64_t hi1; uint64_t lo1 = _umul128(v2, crt.p1, &hi1); #elif defined(__GNUC__) || defined(__clang__) __uint128_t prod1 = static_cast<__uint128_t>(v2) * crt.p1; uint64_t lo1 = static_cast(prod1); uint64_t hi1 = static_cast(prod1 >> 64); #endif // v1 + v2*p1 (128bit) uint64_t c0 = v1 + lo1; uint64_t c1 = hi1 + (c0 < v1 ? 1 : 0); // v3*p1*p2: v3 (63bit) × p1 (63bit) = 126bit → × p2 = 189bit // where v3 < p3 < 2^63, p1 < 2^63, p2 < 2^63 // v3*p1: 126bit → 2 words #if defined(_MSC_VER) && defined(_M_X64) uint64_t hi2; uint64_t lo2 = _umul128(v3, crt.p1, &hi2); // lo2:hi2 = v3*p1 (126bit) // × p2: (lo2:hi2) × p2 = 189bit → 3 words uint64_t hi3; uint64_t lo3 = _umul128(lo2, crt.p2, &hi3); uint64_t hi4; uint64_t lo4 = _umul128(hi2, crt.p2, &hi4); #elif defined(__GNUC__) || defined(__clang__) __uint128_t prod2 = static_cast<__uint128_t>(v3) * crt.p1; uint64_t lo2 = static_cast(prod2); uint64_t hi2 = static_cast(prod2 >> 64); __uint128_t prod3 = static_cast<__uint128_t>(lo2) * crt.p2; uint64_t lo3 = static_cast(prod3); uint64_t hi3 = static_cast(prod3 >> 64); __uint128_t prod4 = static_cast<__uint128_t>(hi2) * crt.p2; uint64_t lo4 = static_cast(prod4); uint64_t hi4 = static_cast(prod4 >> 64); #endif // v3*p1*p2 = lo3 + (hi3 + lo4) << 64 + (hi4) << 128 uint64_t mid = hi3 + lo4; uint64_t carry_mid = (mid < hi3) ? 1 : 0; uint64_t top = hi4 + carry_mid; // c += v3*p1*p2 c0 += lo3; uint64_t carry0 = (c0 < lo3) ? 1 : 0; c1 += mid + carry0; uint64_t carry1 = (c1 < mid || (carry0 && c1 == mid)) ? 1 : 0; uint64_t c2 = top + carry1; result[0] = c0; result[1] = c1; result[2] = c2; } // CRT reconstruction + carry propagation produces the final result inline void crt_recompose(uint64_t* rp, size_t result_limbs, const uint64_t* r1, const uint64_t* r2, const uint64_t* r3, size_t N, const CrtConstants& crt) { // Zero-clear the result std::memset(rp, 0, result_limbs * sizeof(uint64_t)); uint64_t carry_lo = 0, carry_hi = 0; for (size_t i = 0; i < N && i < result_limbs; ++i) { uint64_t coeff[3]; crt_single(coeff, r1[i], r2[i], r3[i], crt); // Add coeff[0:2] + carry uint64_t sum0 = coeff[0] + carry_lo; uint64_t c0 = (sum0 < coeff[0]) ? 1 : 0; uint64_t sum1 = coeff[1] + carry_hi + c0; uint64_t c1 = (sum1 < coeff[1] || (c0 && sum1 <= coeff[1])) ? 1 : 0; uint64_t sum2 = coeff[2] + c1; rp[i] = sum0; carry_lo = sum1; carry_hi = sum2; } // Write out the remaining carry if (carry_lo != 0 && N < result_limbs) { rp[N] = carry_lo; if (carry_hi != 0 && N + 1 < result_limbs) { rp[N + 1] = carry_hi; } } } // ================================================================ // Global NTT parameters (initialized on first use) // ================================================================ struct PrimeNttContext { NttPrime primes[3]; CrtConstants crt; bool initialized = false; void init() { const uint64_t ps[] = { P1, P2, P3 }; const uint64_t gs[] = { G1, G2, G3 }; for (int k = 0; k < 3; ++k) { uint64_t p = ps[k]; uint64_t pin = compute_p_inv_neg(p); // R^2 mod p: R = 2^64, compute via mod_pow(2, 128, p) // 2^128 mod p = (2^64 mod p)^2 mod p uint64_t r_mod_p = mod_pow(2, 64, p); uint64_t r2 = mod_mul(r_mod_p, r_mod_p, p); primes[k] = { p, gs[k], NTT_MAX_S, mod_inv(2, p), pin, r2 }; } crt.init(P1, P2, P3); initialized = true; } }; inline PrimeNttContext& getPrimeNttContext() { static PrimeNttContext ctx; if (!ctx.initialized) { ctx.init(); } return ctx; } // ================================================================ // NTT length selection // ================================================================ inline size_t next_power_of_2(size_t n) { size_t p = 1; while (p < n) p <<= 1; return p; } inline bool is_power_of_2(size_t n) { return n > 0 && (n & (n - 1)) == 0; } // Return the smallest size in {2^a, 3*2^a, 5*2^a} (for 5-smooth NTT) // Example: rn=2050 -> 2560 (5*512), rn=1026 -> 1280 (5*256) inline size_t next_smooth_size(size_t n) { if (n <= 1) return 1; size_t best = next_power_of_2(n); // 2^a size_t c3 = 3 * next_power_of_2((n + 2) / 3); // 3 × 2^a if (c3 >= n && c3 < best) best = c3; size_t c5 = 5 * next_power_of_2((n + 4) / 5); // 5 × 2^a if (c5 >= n && c5 < best) best = c5; return best; } // Thread-pool parallelization threshold (3-prime parallelism when NTT length >= this) static constexpr size_t PRIME_NTT_PARALLEL_THRESHOLD = 4096; // YC-10: super-parallel NTT threshold (intra-NTT multithreading when NTT length >= this) // 128K limbs ~ 2.5M digits. Beyond this, each prime's NTT is split into T threads static constexpr size_t MULTI_THREAD_NTT_THRESHOLD = 131072; // ================================================================ // Per-prime NTT pipeline (Montgomery version, callable from threads) // ================================================================ // Multiplication: da[N], db[N] -> result overwrites da[N] (returned in normal form) inline void ntt_mul_pipeline_mont(uint64_t* da, uint64_t* db, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* fwd_roots, const uint64_t* inv_roots, uint64_t inv_n_normal) { forward_ntt_mont(da, N, p, p_inv_neg, fwd_roots); forward_ntt_mont(db, N, p, p_inv_neg, fwd_roots); pointwise_mul_mont(da, db, N, p, p_inv_neg); inverse_ntt_mont(da, N, p, p_inv_neg, inv_roots, inv_n_normal); } // Squaring: da[N] -> result overwrites da[N] (returned in normal form) inline void ntt_sqr_pipeline_mont(uint64_t* da, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* fwd_roots, const uint64_t* inv_roots, uint64_t inv_n_normal) { forward_ntt_mont(da, N, p, p_inv_neg, fwd_roots); pointwise_sqr_mont(da, N, p, p_inv_neg); inverse_ntt_mont(da, N, p, p_inv_neg, inv_roots, inv_n_normal); } #ifdef __AVX2__ // AVX2 pipeline (uses layered root table) inline void ntt_mul_pipeline_avx2(uint64_t* da, uint64_t* db, size_t N, uint64_t p, uint64_t p_inv_neg, const NttRootsLayered& roots) { forward_ntt_mont_avx2(da, N, p, p_inv_neg, roots.fwd.data(), roots.fwd_offset.data()); forward_ntt_mont_avx2(db, N, p, p_inv_neg, roots.fwd.data(), roots.fwd_offset.data()); pointwise_mul_mont_avx2(da, db, N, p, p_inv_neg); inverse_ntt_mont_avx2(da, N, p, p_inv_neg, roots.inv.data(), roots.inv_offset.data(), roots.inv_n); } inline void ntt_sqr_pipeline_avx2(uint64_t* da, size_t N, uint64_t p, uint64_t p_inv_neg, const NttRootsLayered& roots) { forward_ntt_mont_avx2(da, N, p, p_inv_neg, roots.fwd.data(), roots.fwd_offset.data()); pointwise_sqr_mont_avx2(da, N, p, p_inv_neg); inverse_ntt_mont_avx2(da, N, p, p_inv_neg, roots.inv.data(), roots.inv_offset.data(), roots.inv_n); } // ================================================================ // Mixed-radix NTT pipeline (N = 3M, M = 2^k) // ================================================================ // Forward: radix3_dft → twiddle → 3×M-point NTT // Inverse: 3×M-point INTT → inv_twiddle → radix3_idft(×1/3) // The M-point NTT scales by 1/M via inv_n -> overall 1/(3M) = 1/N inline void mr_ntt_mul_pipeline_avx2(uint64_t* da, uint64_t* db, size_t N, uint64_t p, uint64_t pinv, const NttRootsLayered& sub_roots, const MixedRadixConstants& mr) { size_t M = N / 3; const uint64_t* fwd = sub_roots.fwd.data(); const size_t* fwd_off = sub_roots.fwd_offset.data(); const uint64_t* inv = sub_roots.inv.data(); const size_t* inv_off = sub_roots.inv_offset.data(); uint64_t inv_n_M = sub_roots.inv_n; // Forward da: radix3 → twiddle → 3 sub-NTTs radix3_dft_forward_avx2(da, M, mr.c_plus, mr.c_minus, p, pinv); apply_twiddles_avx2(da, M, mr.twiddle1.data(), mr.twiddle2.data(), p, pinv); forward_ntt_mont_avx2(da, M, p, pinv, fwd, fwd_off); forward_ntt_mont_avx2(da + M, M, p, pinv, fwd, fwd_off); forward_ntt_mont_avx2(da + 2*M, M, p, pinv, fwd, fwd_off); // Forward db radix3_dft_forward_avx2(db, M, mr.c_plus, mr.c_minus, p, pinv); apply_twiddles_avx2(db, M, mr.twiddle1.data(), mr.twiddle2.data(), p, pinv); forward_ntt_mont_avx2(db, M, p, pinv, fwd, fwd_off); forward_ntt_mont_avx2(db + M, M, p, pinv, fwd, fwd_off); forward_ntt_mont_avx2(db + 2*M, M, p, pinv, fwd, fwd_off); // Pointwise multiply (N elements) pointwise_mul_mont_avx2(da, db, N, p, pinv); // Inverse: 3 sub-INTTs → inv_twiddle → radix3_idft(×1/3) inverse_ntt_mont_avx2(da, M, p, pinv, inv, inv_off, inv_n_M); inverse_ntt_mont_avx2(da + M, M, p, pinv, inv, inv_off, inv_n_M); inverse_ntt_mont_avx2(da + 2*M, M, p, pinv, inv, inv_off, inv_n_M); apply_twiddles_avx2(da, M, mr.inv_twiddle1.data(), mr.inv_twiddle2.data(), p, pinv); radix3_dft_inverse_scaled_avx2(da, M, mr.c_plus, mr.c_minus, mr.inv3, p, pinv); } inline void mr_ntt_sqr_pipeline_avx2(uint64_t* da, size_t N, uint64_t p, uint64_t pinv, const NttRootsLayered& sub_roots, const MixedRadixConstants& mr) { size_t M = N / 3; const uint64_t* fwd = sub_roots.fwd.data(); const size_t* fwd_off = sub_roots.fwd_offset.data(); const uint64_t* inv = sub_roots.inv.data(); const size_t* inv_off = sub_roots.inv_offset.data(); uint64_t inv_n_M = sub_roots.inv_n; radix3_dft_forward_avx2(da, M, mr.c_plus, mr.c_minus, p, pinv); apply_twiddles_avx2(da, M, mr.twiddle1.data(), mr.twiddle2.data(), p, pinv); forward_ntt_mont_avx2(da, M, p, pinv, fwd, fwd_off); forward_ntt_mont_avx2(da + M, M, p, pinv, fwd, fwd_off); forward_ntt_mont_avx2(da + 2*M, M, p, pinv, fwd, fwd_off); pointwise_sqr_mont_avx2(da, N, p, pinv); inverse_ntt_mont_avx2(da, M, p, pinv, inv, inv_off, inv_n_M); inverse_ntt_mont_avx2(da + M, M, p, pinv, inv, inv_off, inv_n_M); inverse_ntt_mont_avx2(da + 2*M, M, p, pinv, inv, inv_off, inv_n_M); apply_twiddles_avx2(da, M, mr.inv_twiddle1.data(), mr.inv_twiddle2.data(), p, pinv); radix3_dft_inverse_scaled_avx2(da, M, mr.c_plus, mr.c_minus, mr.inv3, p, pinv); } // ================================================================ // Radix-5 pipeline (N = 5M, M = 2^k) // ================================================================ inline void mr5_ntt_mul_pipeline_avx2(uint64_t* da, uint64_t* db, size_t N, uint64_t p, uint64_t pinv, const NttRootsLayered& sub_roots, const MixedRadix5Constants& mr) { size_t M = N / 5; const uint64_t* fwd = sub_roots.fwd.data(); const size_t* fwd_off = sub_roots.fwd_offset.data(); const uint64_t* inv = sub_roots.inv.data(); const size_t* inv_off = sub_roots.inv_offset.data(); uint64_t inv_n_M = sub_roots.inv_n; const uint64_t* tw[4] = {mr.twiddle[0].data(), mr.twiddle[1].data(), mr.twiddle[2].data(), mr.twiddle[3].data()}; const uint64_t* itw[4] = {mr.inv_twiddle[0].data(), mr.inv_twiddle[1].data(), mr.inv_twiddle[2].data(), mr.inv_twiddle[3].data()}; // Forward da radix5_dft_forward_avx2(da, M, mr.a_prime, mr.b_prime, mr.c_prime, mr.e_prime, mr.a_plus_b, mr.c_plus_e, p, pinv); apply_twiddles5_avx2(da, M, tw, p, pinv); for (int i = 0; i < 5; ++i) forward_ntt_mont_avx2(da + i*M, M, p, pinv, fwd, fwd_off); // Forward db radix5_dft_forward_avx2(db, M, mr.a_prime, mr.b_prime, mr.c_prime, mr.e_prime, mr.a_plus_b, mr.c_plus_e, p, pinv); apply_twiddles5_avx2(db, M, tw, p, pinv); for (int i = 0; i < 5; ++i) forward_ntt_mont_avx2(db + i*M, M, p, pinv, fwd, fwd_off); // Pointwise pointwise_mul_mont_avx2(da, db, N, p, pinv); // Inverse for (int i = 0; i < 5; ++i) inverse_ntt_mont_avx2(da + i*M, M, p, pinv, inv, inv_off, inv_n_M); apply_twiddles5_avx2(da, M, itw, p, pinv); radix5_dft_inverse_scaled_avx2(da, M, mr.a_prime, mr.b_prime, mr.c_prime, mr.e_prime, mr.a_plus_b, mr.c_plus_e, mr.inv5, p, pinv); } inline void mr5_ntt_sqr_pipeline_avx2(uint64_t* da, size_t N, uint64_t p, uint64_t pinv, const NttRootsLayered& sub_roots, const MixedRadix5Constants& mr) { size_t M = N / 5; const uint64_t* fwd = sub_roots.fwd.data(); const size_t* fwd_off = sub_roots.fwd_offset.data(); const uint64_t* inv = sub_roots.inv.data(); const size_t* inv_off = sub_roots.inv_offset.data(); uint64_t inv_n_M = sub_roots.inv_n; const uint64_t* tw[4] = {mr.twiddle[0].data(), mr.twiddle[1].data(), mr.twiddle[2].data(), mr.twiddle[3].data()}; const uint64_t* itw[4] = {mr.inv_twiddle[0].data(), mr.inv_twiddle[1].data(), mr.inv_twiddle[2].data(), mr.inv_twiddle[3].data()}; radix5_dft_forward_avx2(da, M, mr.a_prime, mr.b_prime, mr.c_prime, mr.e_prime, mr.a_plus_b, mr.c_plus_e, p, pinv); apply_twiddles5_avx2(da, M, tw, p, pinv); for (int i = 0; i < 5; ++i) forward_ntt_mont_avx2(da + i*M, M, p, pinv, fwd, fwd_off); pointwise_sqr_mont_avx2(da, N, p, pinv); for (int i = 0; i < 5; ++i) inverse_ntt_mont_avx2(da + i*M, M, p, pinv, inv, inv_off, inv_n_M); apply_twiddles5_avx2(da, M, itw, p, pinv); radix5_dft_inverse_scaled_avx2(da, M, mr.a_prime, mr.b_prime, mr.c_prime, mr.e_prime, mr.a_plus_b, mr.c_plus_e, mr.inv5, p, pinv); } #endif // ================================================================ // YC-10: super-parallel NTT (intra-NTT multithreading) // ================================================================ // Split each prime's NTT into T threads to exploit many-core CPUs. // Layout: 3 primes x T threads/prime = 3T threads (T ~ 21 on 64 cores) // Phase 1 (top layers): split each stage's butterflies across T threads with a barrier // Phase 2 (bottom layers): split into T chunks per cache block (single barrier) // Parallel range helper: split body(start, end) across T threads template inline void ntt_parallel_range(size_t total, int T, F&& body) { if (T <= 1 || total == 0) { body(static_cast(0), total); return; } auto& pool = sangi::threadPool(); size_t chunk = (total + T - 1) / T; std::future futures[128]; int nf = 0; for (int t = 0; t < T - 1; ++t) { size_t s = static_cast(t) * chunk; size_t e = std::min(s + chunk, total); if (s >= total) break; futures[nf++] = pool.submit([s, e, &body]{ body(s, e); }); } size_t last_s = static_cast(T - 1) * chunk; if (last_s < total) body(last_s, total); for (int i = 0; i < nf; ++i) futures[i].get(); } // --- Scalar multi-threaded NTT --- // Forward NTT (DIF) — T-thread parallel inline void forward_ntt_mont_mt(uint64_t* data, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* roots, int T) { for (size_t len = N; len >= 2; len >>= 1) { size_t half = len >> 1; size_t step = N / len; size_t num_groups = N / len; if (num_groups >= static_cast(T)) { // Number of groups >= T: split per group ntt_parallel_range(num_groups, T, [=](size_t g_start, size_t g_end) { for (size_t g = g_start; g < g_end; ++g) { size_t base = g * len; for (size_t j = 0; j < half; ++j) { uint64_t u = data[base + j]; uint64_t v = data[base + j + half]; data[base + j] = mod_add(u, v, p); data[base + j + half] = mont_mul( mod_sub(u, v, p), roots[j * step], p, p_inv_neg); } } }); } else { // Number of groups < T: split the j loop into T (each group independent) ntt_parallel_range(half, T, [=](size_t j_start, size_t j_end) { for (size_t g = 0; g < num_groups; ++g) { size_t base = g * len; for (size_t j = j_start; j < j_end; ++j) { uint64_t u = data[base + j]; uint64_t v = data[base + j + half]; data[base + j] = mod_add(u, v, p); data[base + j + half] = mont_mul( mod_sub(u, v, p), roots[j * step], p, p_inv_neg); } } }); } } } // Inverse NTT (DIT) — T-thread parallel inline void inverse_ntt_mont_mt(uint64_t* data, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* inv_roots, uint64_t inv_n_normal, int T) { for (size_t len = 2; len <= N; len <<= 1) { size_t half = len >> 1; size_t step = N / len; size_t num_groups = N / len; if (num_groups >= static_cast(T)) { ntt_parallel_range(num_groups, T, [=](size_t g_start, size_t g_end) { for (size_t g = g_start; g < g_end; ++g) { size_t base = g * len; for (size_t j = 0; j < half; ++j) { uint64_t u = data[base + j]; uint64_t v = mont_mul(data[base + j + half], inv_roots[j * step], p, p_inv_neg); data[base + j] = mod_add(u, v, p); data[base + j + half] = mod_sub(u, v, p); } } }); } else { ntt_parallel_range(half, T, [=](size_t j_start, size_t j_end) { for (size_t g = 0; g < num_groups; ++g) { size_t base = g * len; for (size_t j = j_start; j < j_end; ++j) { uint64_t u = data[base + j]; uint64_t v = mont_mul(data[base + j + half], inv_roots[j * step], p, p_inv_neg); data[base + j] = mod_add(u, v, p); data[base + j + half] = mod_sub(u, v, p); } } }); } } // N^{-1} scaling (parallel) ntt_parallel_range(N, T, [=](size_t start, size_t end) { for (size_t i = start; i < end; ++i) data[i] = mont_mul(data[i], inv_n_normal, p, p_inv_neg); }); } // Multi-threaded pointwise multiply (auto-select AVX2/scalar) inline void pointwise_mul_mont_mt(uint64_t* a, const uint64_t* b, size_t N, uint64_t p, uint64_t p_inv_neg, int T) { ntt_parallel_range(N, T, [=](size_t start, size_t end) { #ifdef __AVX2__ const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); size_t i = start; for (; i + 4 <= end; i += 4) { __m256i va = _mm256_loadu_si256((__m256i*)(a + i)); __m256i vb = _mm256_loadu_si256((__m256i*)(b + i)); va = avx2_mont_mul(va, vb, p_vec, pin_vec); _mm256_storeu_si256((__m256i*)(a + i), va); } for (; i < end; ++i) a[i] = mont_mul(a[i], b[i], p, p_inv_neg); #else for (size_t i = start; i < end; ++i) a[i] = mont_mul(a[i], b[i], p, p_inv_neg); #endif }); } // Multi-threaded pointwise square inline void pointwise_sqr_mont_mt(uint64_t* a, size_t N, uint64_t p, uint64_t p_inv_neg, int T) { ntt_parallel_range(N, T, [=](size_t start, size_t end) { #ifdef __AVX2__ const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); size_t i = start; for (; i + 4 <= end; i += 4) { __m256i va = _mm256_loadu_si256((__m256i*)(a + i)); va = avx2_mont_mul(va, va, p_vec, pin_vec); _mm256_storeu_si256((__m256i*)(a + i), va); } for (; i < end; ++i) a[i] = mont_mul(a[i], a[i], p, p_inv_neg); #else for (size_t i = start; i < end; ++i) a[i] = mont_mul(a[i], a[i], p, p_inv_neg); #endif }); } // Scalar multi-threaded pipeline: mul inline void ntt_mul_pipeline_mt(uint64_t* da, uint64_t* db, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* fwd_roots, const uint64_t* inv_roots, uint64_t inv_n_normal, int T) { forward_ntt_mont_mt(da, N, p, p_inv_neg, fwd_roots, T); forward_ntt_mont_mt(db, N, p, p_inv_neg, fwd_roots, T); pointwise_mul_mont_mt(da, db, N, p, p_inv_neg, T); inverse_ntt_mont_mt(da, N, p, p_inv_neg, inv_roots, inv_n_normal, T); } // Scalar multi-threaded pipeline: sqr inline void ntt_sqr_pipeline_mt(uint64_t* da, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* fwd_roots, const uint64_t* inv_roots, uint64_t inv_n_normal, int T) { forward_ntt_mont_mt(da, N, p, p_inv_neg, fwd_roots, T); pointwise_sqr_mont_mt(da, N, p, p_inv_neg, T); inverse_ntt_mont_mt(da, N, p, p_inv_neg, inv_roots, inv_n_normal, T); } #ifdef __AVX2__ // --- AVX2 multi-threaded NTT --- // AVX2 forward NTT (DIF, cache-blocked) — T-thread parallel inline void forward_ntt_mont_avx2_mt(uint64_t* data, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* layer_roots, const size_t* layer_offsets, int T) { size_t block = (N <= NTT_CACHE_BLOCK) ? N : NTT_CACHE_BLOCK; // Phase 1: top layers (len > block) — per-stage barrier int layer = 0; for (size_t len = N; len > block; len >>= 1) { size_t half = len >> 1; size_t num_groups = N / len; const uint64_t* roots = layer_roots + layer_offsets[layer]; if (num_groups >= static_cast(T)) { // Split per group ntt_parallel_range(num_groups, T, [=](size_t g_start, size_t g_end) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); for (size_t g = g_start; g < g_end; ++g) dif_butterfly_avx2(data, g * len, half, roots, p_vec, pin_vec, p, p_inv_neg); }); } else { // Split the j loop (split each group's half-butterflies into T) ntt_parallel_range(half, T, [=](size_t j_start, size_t j_end) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); for (size_t g = 0; g < num_groups; ++g) { size_t start = g * len; uint64_t* d0 = data + start; uint64_t* d1 = data + start + half; size_t j = j_start; for (; j + 4 <= j_end; j += 4) { __m256i u = _mm256_loadu_si256((__m256i*)(d0 + j)); __m256i v = _mm256_loadu_si256((__m256i*)(d1 + j)); __m256i w = _mm256_loadu_si256((__m256i*)(roots + j)); __m256i sum = _mm256_add_epi64(u, v); __m256i sub = _mm256_sub_epi64(sum, p_vec); __m256i sm = _mm256_cmpgt_epi64(_mm256_setzero_si256(), sub); sum = _mm256_blendv_epi8(sub, sum, sm); __m256i diff = _mm256_sub_epi64(u, v); __m256i da = _mm256_add_epi64(diff, p_vec); __m256i dm = _mm256_cmpgt_epi64(_mm256_setzero_si256(), diff); diff = _mm256_blendv_epi8(diff, da, dm); __m256i tw = avx2_mont_mul(diff, w, p_vec, pin_vec); _mm256_storeu_si256((__m256i*)(d0 + j), sum); _mm256_storeu_si256((__m256i*)(d1 + j), tw); } for (; j < j_end; ++j) { uint64_t uv = d0[j], vv = d1[j]; d0[j] = mod_add(uv, vv, p); d1[j] = mont_mul(mod_sub(uv, vv, p), roots[j], p, p_inv_neg); } } }); } ++layer; } // Phase 2: bottom layers (len <= block) — block-parallel (single barrier) int bottom_start = layer; size_t num_blocks = N / block; ntt_parallel_range(num_blocks, T, [=](size_t blk_start, size_t blk_end) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); for (size_t bi = blk_start; bi < blk_end; ++bi) { size_t blk = bi * block; int cur_layer = bottom_start; for (size_t len = block; len >= 2; len >>= 1) { size_t half = len >> 1; const uint64_t* roots = layer_roots + layer_offsets[cur_layer]; for (size_t start = blk; start < blk + block; start += len) dif_butterfly_avx2(data, start, half, roots, p_vec, pin_vec, p, p_inv_neg); ++cur_layer; } } }); } // AVX2 inverse NTT (DIT, cache-blocked) — T-thread parallel inline void inverse_ntt_mont_avx2_mt(uint64_t* data, size_t N, uint64_t p, uint64_t p_inv_neg, const uint64_t* layer_roots, const size_t* layer_offsets, uint64_t inv_n_normal, int T) { size_t block = (N <= NTT_CACHE_BLOCK) ? N : NTT_CACHE_BLOCK; // Phase 1: bottom layers (len <= block) — block-parallel (single barrier) int num_bottom_layers = 0; for (size_t l = 2; l <= block; l <<= 1) ++num_bottom_layers; size_t num_blocks = N / block; ntt_parallel_range(num_blocks, T, [=](size_t blk_start, size_t blk_end) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); for (size_t bi = blk_start; bi < blk_end; ++bi) { size_t blk = bi * block; int cur_layer = 0; for (size_t len = 2; len <= block; len <<= 1) { size_t half = len >> 1; const uint64_t* roots = layer_roots + layer_offsets[cur_layer]; for (size_t start = blk; start < blk + block; start += len) dit_butterfly_avx2(data, start, half, roots, p_vec, pin_vec, p, p_inv_neg); ++cur_layer; } } }); // Phase 2: top layers (len > block) — per-stage barrier int layer = num_bottom_layers; for (size_t len = block * 2; len <= N; len <<= 1) { size_t half = len >> 1; size_t num_groups = N / len; const uint64_t* roots = layer_roots + layer_offsets[layer]; if (num_groups >= static_cast(T)) { ntt_parallel_range(num_groups, T, [=](size_t g_start, size_t g_end) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); for (size_t g = g_start; g < g_end; ++g) dit_butterfly_avx2(data, g * len, half, roots, p_vec, pin_vec, p, p_inv_neg); }); } else { ntt_parallel_range(half, T, [=](size_t j_start, size_t j_end) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); for (size_t g = 0; g < num_groups; ++g) { size_t start = g * len; uint64_t* d0 = data + start; uint64_t* d1 = data + start + half; size_t j = j_start; for (; j + 4 <= j_end; j += 4) { __m256i u = _mm256_loadu_si256((__m256i*)(d0 + j)); __m256i v_raw = _mm256_loadu_si256((__m256i*)(d1 + j)); __m256i w = _mm256_loadu_si256((__m256i*)(roots + j)); __m256i v = avx2_mont_mul(v_raw, w, p_vec, pin_vec); __m256i sum = _mm256_add_epi64(u, v); __m256i sub = _mm256_sub_epi64(sum, p_vec); __m256i sm = _mm256_cmpgt_epi64(_mm256_setzero_si256(), sub); sum = _mm256_blendv_epi8(sub, sum, sm); __m256i diff = _mm256_sub_epi64(u, v); __m256i da = _mm256_add_epi64(diff, p_vec); __m256i dm = _mm256_cmpgt_epi64(_mm256_setzero_si256(), diff); diff = _mm256_blendv_epi8(diff, da, dm); _mm256_storeu_si256((__m256i*)(d0 + j), sum); _mm256_storeu_si256((__m256i*)(d1 + j), diff); } for (; j < j_end; ++j) { uint64_t uv = d0[j]; uint64_t vv = mont_mul(d1[j], roots[j], p, p_inv_neg); d0[j] = mod_add(uv, vv, p); d1[j] = mod_sub(uv, vv, p); } } }); } ++layer; } // N^{-1} scaling (parallel, AVX2) ntt_parallel_range(N, T, [=](size_t start, size_t end) { const __m256i p_vec = _mm256_set1_epi64x(p); const __m256i pin_vec = _mm256_set1_epi64x(p_inv_neg); __m256i inv_n_vec = _mm256_set1_epi64x(inv_n_normal); size_t i = start; for (; i + 4 <= end; i += 4) { __m256i d = _mm256_loadu_si256((__m256i*)(data + i)); d = avx2_mont_mul(d, inv_n_vec, p_vec, pin_vec); _mm256_storeu_si256((__m256i*)(data + i), d); } for (; i < end; ++i) data[i] = mont_mul(data[i], inv_n_normal, p, p_inv_neg); }); } // AVX2 multi-threaded pipeline: mul inline void ntt_mul_pipeline_avx2_mt(uint64_t* da, uint64_t* db, size_t N, uint64_t p, uint64_t p_inv_neg, const NttRootsLayered& roots, int T) { forward_ntt_mont_avx2_mt(da, N, p, p_inv_neg, roots.fwd.data(), roots.fwd_offset.data(), T); forward_ntt_mont_avx2_mt(db, N, p, p_inv_neg, roots.fwd.data(), roots.fwd_offset.data(), T); pointwise_mul_mont_mt(da, db, N, p, p_inv_neg, T); inverse_ntt_mont_avx2_mt(da, N, p, p_inv_neg, roots.inv.data(), roots.inv_offset.data(), roots.inv_n, T); } // AVX2 multi-threaded pipeline: sqr inline void ntt_sqr_pipeline_avx2_mt(uint64_t* da, size_t N, uint64_t p, uint64_t p_inv_neg, const NttRootsLayered& roots, int T) { forward_ntt_mont_avx2_mt(da, N, p, p_inv_neg, roots.fwd.data(), roots.fwd_offset.data(), T); pointwise_sqr_mont_mt(da, N, p, p_inv_neg, T); inverse_ntt_mont_avx2_mt(da, N, p, p_inv_neg, roots.inv.data(), roots.inv_offset.data(), roots.inv_n, T); } #endif // ================================================================ // YC-2: Fused Multiply-Add (rp = a*b + c*d) // ================================================================ // Forward NTT on 4 inputs -> pointwise MAC -> one inverse NTT + CRT // Compared with two normal multiplications + addition, saves one inverse NTT, one CRT, and one big-number addition // rp[0..rn-1] = ap[0..an-1]*bp[0..bn-1] + cp[0..cn-1]*dp[0..dn-1] inline void mul_add_prime_ntt(uint64_t* rp, size_t rn, const uint64_t* ap, size_t an, const uint64_t* bp, size_t bn, const uint64_t* cp, size_t cn, const uint64_t* dp, size_t dn) { auto& ctx = getPrimeNttContext(); // Set NTT length to the larger of the two products size_t prod1_n = an + bn; size_t prod2_n = cn + dn; size_t max_prod = std::max(prod1_n, prod2_n); size_t N = next_power_of_2(max_prod); // 4 inputs x 3 primes = 12N words thread_local std::vector work; size_t total = 12 * N; if (work.size() < total) work.resize(total); uint64_t* da[3]; // NTT of a uint64_t* db[3]; // NTT of b uint64_t* dc[3]; // NTT of c uint64_t* dd[3]; // NTT of d for (int k = 0; k < 3; ++k) { da[k] = work.data() + k * 4 * N; db[k] = da[k] + N; dc[k] = db[k] + N; dd[k] = dc[k] + N; } // Decompose input: mod p -> Montgomery form for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p; uint64_t pin = ctx.primes[k].p_inv_neg; uint64_t r2 = ctx.primes[k].r2_mod_p; // a for (size_t i = 0; i < an; ++i) da[k][i] = to_mont(ap[i] % p, p, pin, r2); for (size_t i = an; i < N; ++i) da[k][i] = 0; // b for (size_t i = 0; i < bn; ++i) db[k][i] = to_mont(bp[i] % p, p, pin, r2); for (size_t i = bn; i < N; ++i) db[k][i] = 0; // c for (size_t i = 0; i < cn; ++i) dc[k][i] = to_mont(cp[i] % p, p, pin, r2); for (size_t i = cn; i < N; ++i) dc[k][i] = 0; // d for (size_t i = 0; i < dn; ++i) dd[k][i] = to_mont(dp[i] % p, p, pin, r2); for (size_t i = dn; i < N; ++i) dd[k][i] = 0; } #ifdef __AVX2__ thread_local NttRootsLayered roots_avx[3]; for (int k = 0; k < 3; ++k) { if (roots_avx[k].N != N || roots_avx[k].p != ctx.primes[k].p) { roots_avx[k].build(ctx.primes[k], N); } } bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (parallel) { const NttRootsLayered* rptr[3] = {&roots_avx[0], &roots_avx[1], &roots_avx[2]}; // Process primes 0, 1 on workers and prime 2 on the main thread auto f0 = sangi::threadPool().submit([&, N, rptr]{ uint64_t p = ctx.primes[0].p, pin = ctx.primes[0].p_inv_neg; const auto& r = *rptr[0]; forward_ntt_mont_avx2(da[0], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(db[0], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dc[0], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dd[0], N, p, pin, r.fwd.data(), r.fwd_offset.data()); pointwise_mul_mont_avx2(da[0], db[0], N, p, pin); pointwise_mac_mont_avx2(da[0], dc[0], dd[0], N, p, pin); inverse_ntt_mont_avx2(da[0], N, p, pin, r.inv.data(), r.inv_offset.data(), r.inv_n); }); auto f1 = sangi::threadPool().submit([&, N, rptr]{ uint64_t p = ctx.primes[1].p, pin = ctx.primes[1].p_inv_neg; const auto& r = *rptr[1]; forward_ntt_mont_avx2(da[1], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(db[1], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dc[1], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dd[1], N, p, pin, r.fwd.data(), r.fwd_offset.data()); pointwise_mul_mont_avx2(da[1], db[1], N, p, pin); pointwise_mac_mont_avx2(da[1], dc[1], dd[1], N, p, pin); inverse_ntt_mont_avx2(da[1], N, p, pin, r.inv.data(), r.inv_offset.data(), r.inv_n); }); { uint64_t p = ctx.primes[2].p, pin = ctx.primes[2].p_inv_neg; const auto& r = *rptr[2]; forward_ntt_mont_avx2(da[2], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(db[2], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dc[2], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dd[2], N, p, pin, r.fwd.data(), r.fwd_offset.data()); pointwise_mul_mont_avx2(da[2], db[2], N, p, pin); pointwise_mac_mont_avx2(da[2], dc[2], dd[2], N, p, pin); inverse_ntt_mont_avx2(da[2], N, p, pin, r.inv.data(), r.inv_offset.data(), r.inv_n); } f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p, pin = ctx.primes[k].p_inv_neg; const auto& r = roots_avx[k]; forward_ntt_mont_avx2(da[k], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(db[k], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dc[k], N, p, pin, r.fwd.data(), r.fwd_offset.data()); forward_ntt_mont_avx2(dd[k], N, p, pin, r.fwd.data(), r.fwd_offset.data()); pointwise_mul_mont_avx2(da[k], db[k], N, p, pin); pointwise_mac_mont_avx2(da[k], dc[k], dd[k], N, p, pin); inverse_ntt_mont_avx2(da[k], N, p, pin, r.inv.data(), r.inv_offset.data(), r.inv_n); } } #else thread_local NttRootsMont roots[3]; for (int k = 0; k < 3; ++k) { if (roots[k].N != N || roots[k].p != ctx.primes[k].p) { roots[k].build(ctx.primes[k], N); } } bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (parallel) { auto f0 = sangi::threadPool().submit([&, N]{ uint64_t p = ctx.primes[0].p, pin = ctx.primes[0].p_inv_neg; forward_ntt_mont(da[0], N, p, pin, roots[0].roots.data()); forward_ntt_mont(db[0], N, p, pin, roots[0].roots.data()); forward_ntt_mont(dc[0], N, p, pin, roots[0].roots.data()); forward_ntt_mont(dd[0], N, p, pin, roots[0].roots.data()); pointwise_mul_mont(da[0], db[0], N, p, pin); pointwise_mac_mont(da[0], dc[0], dd[0], N, p, pin); inverse_ntt_mont(da[0], N, p, pin, roots[0].inv_roots.data(), roots[0].inv_n); }); auto f1 = sangi::threadPool().submit([&, N]{ uint64_t p = ctx.primes[1].p, pin = ctx.primes[1].p_inv_neg; forward_ntt_mont(da[1], N, p, pin, roots[1].roots.data()); forward_ntt_mont(db[1], N, p, pin, roots[1].roots.data()); forward_ntt_mont(dc[1], N, p, pin, roots[1].roots.data()); forward_ntt_mont(dd[1], N, p, pin, roots[1].roots.data()); pointwise_mul_mont(da[1], db[1], N, p, pin); pointwise_mac_mont(da[1], dc[1], dd[1], N, p, pin); inverse_ntt_mont(da[1], N, p, pin, roots[1].inv_roots.data(), roots[1].inv_n); }); { uint64_t p = ctx.primes[2].p, pin = ctx.primes[2].p_inv_neg; forward_ntt_mont(da[2], N, p, pin, roots[2].roots.data()); forward_ntt_mont(db[2], N, p, pin, roots[2].roots.data()); forward_ntt_mont(dc[2], N, p, pin, roots[2].roots.data()); forward_ntt_mont(dd[2], N, p, pin, roots[2].roots.data()); pointwise_mul_mont(da[2], db[2], N, p, pin); pointwise_mac_mont(da[2], dc[2], dd[2], N, p, pin); inverse_ntt_mont(da[2], N, p, pin, roots[2].inv_roots.data(), roots[2].inv_n); } f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p, pin = ctx.primes[k].p_inv_neg; forward_ntt_mont(da[k], N, p, pin, roots[k].roots.data()); forward_ntt_mont(db[k], N, p, pin, roots[k].roots.data()); forward_ntt_mont(dc[k], N, p, pin, roots[k].roots.data()); forward_ntt_mont(dd[k], N, p, pin, roots[k].roots.data()); pointwise_mul_mont(da[k], db[k], N, p, pin); pointwise_mac_mont(da[k], dc[k], dd[k], N, p, pin); inverse_ntt_mont(da[k], N, p, pin, roots[k].inv_roots.data(), roots[k].inv_n); } } #endif // CRT reconstruction crt_recompose(rp, rn, da[0], da[1], da[2], max_prod, ctx.crt); } // ================================================================ // Main multiplication entry point // ================================================================ // rp[0..an+bn-1] = ap[0..an-1] × bp[0..bn-1] inline void mul_prime_ntt(uint64_t* rp, const uint64_t* ap, size_t an, const uint64_t* bp, size_t bn) { if (an < bn) { std::swap(ap, bp); std::swap(an, bn); } auto& ctx = getPrimeNttContext(); size_t rn = an + bn; size_t N = next_smooth_size(rn); // 2^a, 3×2^a, or 5×2^a // Radix selection: 0=power-of-2, 3=radix-3, 5=radix-5 int radix = 0; if (N % 5 == 0 && is_power_of_2(N / 5)) radix = 5; else if (N % 3 == 0 && is_power_of_2(N / 3)) radix = 3; // Allocate NTT data arrays for the 3 primes thread_local std::vector work; size_t total = 6 * N; if (work.size() < total) work.resize(total); uint64_t* da[3]; uint64_t* db[3]; for (int k = 0; k < 3; ++k) { da[k] = work.data() + k * 2 * N; db[k] = da[k] + N; } // Decompose input: a[i] mod p -> convert to Montgomery form for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p; uint64_t pin = ctx.primes[k].p_inv_neg; uint64_t r2 = ctx.primes[k].r2_mod_p; for (size_t i = 0; i < an; ++i) da[k][i] = to_mont(ap[i] % p, p, pin, r2); for (size_t i = an; i < N; ++i) da[k][i] = 0; for (size_t i = 0; i < bn; ++i) db[k][i] = to_mont(bp[i] % p, p, pin, r2); for (size_t i = bn; i < N; ++i) db[k][i] = 0; } #ifdef __AVX2__ if (radix == 5) { // Radix-5 path (N = 5M) size_t M = N / 5; thread_local NttRootsLayered sub_roots_avx[3]; thread_local MixedRadix5Constants mr5_const[3]; for (int k = 0; k < 3; ++k) { if (sub_roots_avx[k].N != M || sub_roots_avx[k].p != ctx.primes[k].p) sub_roots_avx[k].build(ctx.primes[k], M); if (mr5_const[k].N != N || mr5_const[k].p != ctx.primes[k].p) mr5_const[k].build(ctx.primes[k], N); } // Always parallelize { const NttRootsLayered* srp[3] = {&sub_roots_avx[0], &sub_roots_avx[1], &sub_roots_avx[2]}; const MixedRadix5Constants* mrp[3] = {&mr5_const[0], &mr5_const[1], &mr5_const[2]}; auto f0 = sangi::threadPool().submit([&, N, srp, mrp]{ mr5_ntt_mul_pipeline_avx2(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *srp[0], *mrp[0]); }); auto f1 = sangi::threadPool().submit([&, N, srp, mrp]{ mr5_ntt_mul_pipeline_avx2(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *srp[1], *mrp[1]); }); mr5_ntt_mul_pipeline_avx2(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *srp[2], *mrp[2]); f0.get(); f1.get(); } } else if (radix == 3) { // Radix-3 path (N = 3M) size_t M = N / 3; thread_local NttRootsLayered sub_roots_avx[3]; thread_local MixedRadixConstants mr_const[3]; for (int k = 0; k < 3; ++k) { if (sub_roots_avx[k].N != M || sub_roots_avx[k].p != ctx.primes[k].p) sub_roots_avx[k].build(ctx.primes[k], M); if (mr_const[k].N != N || mr_const[k].p != ctx.primes[k].p) mr_const[k].build(ctx.primes[k], N); } // Mixed-radix: 9 sub-NTTs/prime -> always parallelize (enough work when M >= 256) { const NttRootsLayered* srp[3] = {&sub_roots_avx[0], &sub_roots_avx[1], &sub_roots_avx[2]}; const MixedRadixConstants* mrp[3] = {&mr_const[0], &mr_const[1], &mr_const[2]}; auto f0 = sangi::threadPool().submit([&, N, srp, mrp]{ mr_ntt_mul_pipeline_avx2(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *srp[0], *mrp[0]); }); auto f1 = sangi::threadPool().submit([&, N, srp, mrp]{ mr_ntt_mul_pipeline_avx2(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *srp[1], *mrp[1]); }); mr_ntt_mul_pipeline_avx2(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *srp[2], *mrp[2]); f0.get(); f1.get(); } } else { // Existing power-of-2 path thread_local NttRootsLayered roots_avx[3]; for (int k = 0; k < 3; ++k) { if (roots_avx[k].N != N || roots_avx[k].p != ctx.primes[k].p) roots_avx[k].build(ctx.primes[k], N); } bool use_mt = (N >= MULTI_THREAD_NTT_THRESHOLD); bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (use_mt) { unsigned hw = std::thread::hardware_concurrency(); int T = std::max(2, static_cast(hw > 3 ? (hw - 1) / 3 : 1)); const NttRootsLayered* rpp[3] = {&roots_avx[0], &roots_avx[1], &roots_avx[2]}; auto f0 = sangi::threadPool().submit([&, N, T, rpp]{ ntt_mul_pipeline_avx2_mt(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *rpp[0], T); }); auto f1 = sangi::threadPool().submit([&, N, T, rpp]{ ntt_mul_pipeline_avx2_mt(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *rpp[1], T); }); ntt_mul_pipeline_avx2_mt(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *rpp[2], T); f0.get(); f1.get(); } else if (parallel) { const NttRootsLayered* rpp[3] = {&roots_avx[0], &roots_avx[1], &roots_avx[2]}; auto f0 = sangi::threadPool().submit([&, N, rpp]{ ntt_mul_pipeline_avx2(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *rpp[0]); }); auto f1 = sangi::threadPool().submit([&, N, rpp]{ ntt_mul_pipeline_avx2(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *rpp[1]); }); ntt_mul_pipeline_avx2(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *rpp[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { ntt_mul_pipeline_avx2(da[k], db[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, roots_avx[k]); } } } #else // Scalar Montgomery NTT (mixed-radix is AVX2-only -> fall back to power-of-2) if (radix != 0) N = next_power_of_2(rn); thread_local NttRootsMont roots[3]; for (int k = 0; k < 3; ++k) { if (roots[k].N != N || roots[k].p != ctx.primes[k].p) roots[k].build(ctx.primes[k], N); } const uint64_t* fwd_r[3], *inv_r[3]; uint64_t inv_n_vals[3]; for (int k = 0; k < 3; ++k) { fwd_r[k] = roots[k].roots.data(); inv_r[k] = roots[k].inv_roots.data(); inv_n_vals[k] = roots[k].inv_n; } bool use_mt = (N >= MULTI_THREAD_NTT_THRESHOLD); bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (use_mt) { unsigned hw = std::thread::hardware_concurrency(); int T = std::max(2, static_cast(hw > 3 ? (hw - 1) / 3 : 1)); auto f0 = sangi::threadPool().submit([&, N, T]{ ntt_mul_pipeline_mt(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_r[0], inv_r[0], inv_n_vals[0], T); }); auto f1 = sangi::threadPool().submit([&, N, T]{ ntt_mul_pipeline_mt(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_r[1], inv_r[1], inv_n_vals[1], T); }); ntt_mul_pipeline_mt(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_r[2], inv_r[2], inv_n_vals[2], T); f0.get(); f1.get(); } else if (parallel) { auto f0 = sangi::threadPool().submit([&, N]{ ntt_mul_pipeline_mont(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_r[0], inv_r[0], inv_n_vals[0]); }); auto f1 = sangi::threadPool().submit([&, N]{ ntt_mul_pipeline_mont(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_r[1], inv_r[1], inv_n_vals[1]); }); ntt_mul_pipeline_mont(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_r[2], inv_r[2], inv_n_vals[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { ntt_mul_pipeline_mont(da[k], db[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, fwd_r[k], inv_r[k], inv_n_vals[k]); } } #endif // CRT reconstruction + carry propagation crt_recompose(rp, rn, da[0], da[1], da[2], rn, ctx.crt); } // ================================================================ // YC-1d: middle product // ================================================================ // // Of the full product c[0..an+bn-2] = ap[0..an-1] x bp[0..bn-1], // write the middle portion c[bn-1..an-1] (an-bn+1 limbs total) to rp. // // Mathematical justification: // Performing cyclic convolution with NTT length N = next_pow2(an), wrap-around at position k // comes from terms with i+j = k+N (i < an, j < bn). When k >= bn-1, k+N >= bn-1+an > an+bn-2, // and since i+j <= an-1+bn-1 = an+bn-2 < k+N, no wrap-around terms exist. // Hence positions bn-1..an-1 are obtained exactly. // // Use: residual computation in Newton division — compute the residual 1 - bx (with bx ~ 1) // using a shorter NTT length (an) instead of the full product (NTT length an+bn). // // Precondition: an >= bn >= 1 // Output: rp[0..an-bn] = c[bn-1..an-1] (an-bn+1 limbs) inline void middle_product_prime_ntt(uint64_t* rp, size_t rn, const uint64_t* ap, size_t an, const uint64_t* bp, size_t bn) { // rn = an - bn + 1 (guaranteed by caller) auto& ctx = getPrimeNttContext(); // NTT length: next_pow2(an) — a full product would need next_pow2(an+bn), but // for the middle product, an suffices (per the justification above) size_t N = next_power_of_2(an); // Workspace for 3 primes: da[N] + db[N] x 3 = 6N thread_local std::vector work_mp; size_t total = 6 * N; if (work_mp.size() < total) work_mp.resize(total); uint64_t* da[3]; uint64_t* db[3]; for (int k = 0; k < 3; ++k) { da[k] = work_mp.data() + k * 2 * N; db[k] = da[k] + N; } // Decompose input: a -> Montgomery form (as-is) // b -> reverse and convert to Montgomery form // Via b_rev[j] = b[bn-1-j], position k of cyclic convolution corresponds to // position k + (bn-1) of the full product for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p; uint64_t pin = ctx.primes[k].p_inv_neg; uint64_t r2 = ctx.primes[k].r2_mod_p; for (size_t i = 0; i < an; ++i) da[k][i] = to_mont(ap[i] % p, p, pin, r2); for (size_t i = an; i < N; ++i) da[k][i] = 0; // Place b in reversed order for (size_t i = 0; i < bn; ++i) db[k][i] = to_mont(bp[bn - 1 - i] % p, p, pin, r2); for (size_t i = bn; i < N; ++i) db[k][i] = 0; } #ifdef __AVX2__ thread_local NttRootsLayered roots_avx_mp[3]; for (int k = 0; k < 3; ++k) { if (roots_avx_mp[k].N != N || roots_avx_mp[k].p != ctx.primes[k].p) { roots_avx_mp[k].build(ctx.primes[k], N); } } bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (parallel) { const NttRootsLayered* rp_roots[3] = {&roots_avx_mp[0], &roots_avx_mp[1], &roots_avx_mp[2]}; auto f0 = sangi::threadPool().submit([&, N, rp_roots]{ ntt_mul_pipeline_avx2(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *rp_roots[0]); }); auto f1 = sangi::threadPool().submit([&, N, rp_roots]{ ntt_mul_pipeline_avx2(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *rp_roots[1]); }); ntt_mul_pipeline_avx2(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *rp_roots[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { ntt_mul_pipeline_avx2(da[k], db[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, roots_avx_mp[k]); } } #else thread_local NttRootsMont roots_mp[3]; for (int k = 0; k < 3; ++k) { if (roots_mp[k].N != N || roots_mp[k].p != ctx.primes[k].p) { roots_mp[k].build(ctx.primes[k], N); } } const uint64_t* fwd_r[3], *inv_r[3]; uint64_t inv_n_vals[3]; for (int k = 0; k < 3; ++k) { fwd_r[k] = roots_mp[k].roots.data(); inv_r[k] = roots_mp[k].inv_roots.data(); inv_n_vals[k] = roots_mp[k].inv_n; } bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (parallel) { auto f0 = sangi::threadPool().submit([&, N]{ ntt_mul_pipeline_mont(da[0], db[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_r[0], inv_r[0], inv_n_vals[0]); }); auto f1 = sangi::threadPool().submit([&, N]{ ntt_mul_pipeline_mont(da[1], db[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_r[1], inv_r[1], inv_n_vals[1]); }); ntt_mul_pipeline_mont(da[2], db[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_r[2], inv_r[2], inv_n_vals[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { ntt_mul_pipeline_mont(da[k], db[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, fwd_r[k], inv_r[k], inv_n_vals[k]); } } #endif // CRT reconstruction: recover N entries into a scratch buffer and extract positions 0..rn-1 // (position 0 of the cyclic convolution corresponds to position bn-1 of the full product) // Even for large N, the full-product length (an+bn) is unnecessary — N entries suffice thread_local std::vector crt_buf; if (crt_buf.size() < N) crt_buf.resize(N); crt_recompose(crt_buf.data(), N, da[0], da[1], da[2], N, ctx.crt); // Copy positions 0..rn-1 to output (= c[bn-1..an-1] of the full product) std::memcpy(rp, crt_buf.data(), rn * sizeof(uint64_t)); } // ================================================================ // Squaring function // ================================================================ // rp[0..2n-1] = ap[0..n-1]² inline void sqr_prime_ntt(uint64_t* rp, const uint64_t* ap, size_t an) { auto& ctx = getPrimeNttContext(); size_t rn = 2 * an; size_t N = next_smooth_size(rn); int radix = 0; if (N % 5 == 0 && is_power_of_2(N / 5)) radix = 5; else if (N % 3 == 0 && is_power_of_2(N / 3)) radix = 3; thread_local std::vector work; size_t total = 3 * N; if (work.size() < total) work.resize(total); uint64_t* da[3]; for (int k = 0; k < 3; ++k) { da[k] = work.data() + k * N; } // Decompose input -> Montgomery form for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p; uint64_t pin = ctx.primes[k].p_inv_neg; uint64_t r2 = ctx.primes[k].r2_mod_p; for (size_t i = 0; i < an; ++i) da[k][i] = to_mont(ap[i] % p, p, pin, r2); for (size_t i = an; i < N; ++i) da[k][i] = 0; } #ifdef __AVX2__ if (radix == 5) { size_t M = N / 5; thread_local NttRootsLayered sub_roots_avx[3]; thread_local MixedRadix5Constants mr5_const[3]; for (int k = 0; k < 3; ++k) { if (sub_roots_avx[k].N != M || sub_roots_avx[k].p != ctx.primes[k].p) sub_roots_avx[k].build(ctx.primes[k], M); if (mr5_const[k].N != N || mr5_const[k].p != ctx.primes[k].p) mr5_const[k].build(ctx.primes[k], N); } { const NttRootsLayered* srp[3] = {&sub_roots_avx[0], &sub_roots_avx[1], &sub_roots_avx[2]}; const MixedRadix5Constants* mrp[3] = {&mr5_const[0], &mr5_const[1], &mr5_const[2]}; auto f0 = sangi::threadPool().submit([&, N, srp, mrp]{ mr5_ntt_sqr_pipeline_avx2(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *srp[0], *mrp[0]); }); auto f1 = sangi::threadPool().submit([&, N, srp, mrp]{ mr5_ntt_sqr_pipeline_avx2(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *srp[1], *mrp[1]); }); mr5_ntt_sqr_pipeline_avx2(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *srp[2], *mrp[2]); f0.get(); f1.get(); } } else if (radix == 3) { size_t M = N / 3; thread_local NttRootsLayered sub_roots_avx[3]; thread_local MixedRadixConstants mr_const[3]; for (int k = 0; k < 3; ++k) { if (sub_roots_avx[k].N != M || sub_roots_avx[k].p != ctx.primes[k].p) sub_roots_avx[k].build(ctx.primes[k], M); if (mr_const[k].N != N || mr_const[k].p != ctx.primes[k].p) mr_const[k].build(ctx.primes[k], N); } { const NttRootsLayered* srp[3] = {&sub_roots_avx[0], &sub_roots_avx[1], &sub_roots_avx[2]}; const MixedRadixConstants* mrp[3] = {&mr_const[0], &mr_const[1], &mr_const[2]}; auto f0 = sangi::threadPool().submit([&, N, srp, mrp]{ mr_ntt_sqr_pipeline_avx2(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *srp[0], *mrp[0]); }); auto f1 = sangi::threadPool().submit([&, N, srp, mrp]{ mr_ntt_sqr_pipeline_avx2(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *srp[1], *mrp[1]); }); mr_ntt_sqr_pipeline_avx2(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *srp[2], *mrp[2]); f0.get(); f1.get(); } } else { thread_local NttRootsLayered roots_avx[3]; for (int k = 0; k < 3; ++k) { if (roots_avx[k].N != N || roots_avx[k].p != ctx.primes[k].p) roots_avx[k].build(ctx.primes[k], N); } bool use_mt = (N >= MULTI_THREAD_NTT_THRESHOLD); bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (use_mt) { unsigned hw = std::thread::hardware_concurrency(); int T = std::max(2, static_cast(hw > 3 ? (hw - 1) / 3 : 1)); const NttRootsLayered* rpp[3] = {&roots_avx[0], &roots_avx[1], &roots_avx[2]}; auto f0 = sangi::threadPool().submit([&, N, T, rpp]{ ntt_sqr_pipeline_avx2_mt(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *rpp[0], T); }); auto f1 = sangi::threadPool().submit([&, N, T, rpp]{ ntt_sqr_pipeline_avx2_mt(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *rpp[1], T); }); ntt_sqr_pipeline_avx2_mt(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *rpp[2], T); f0.get(); f1.get(); } else if (parallel) { const NttRootsLayered* rpp[3] = {&roots_avx[0], &roots_avx[1], &roots_avx[2]}; auto f0 = sangi::threadPool().submit([&, N, rpp]{ ntt_sqr_pipeline_avx2(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, *rpp[0]); }); auto f1 = sangi::threadPool().submit([&, N, rpp]{ ntt_sqr_pipeline_avx2(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, *rpp[1]); }); ntt_sqr_pipeline_avx2(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, *rpp[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) ntt_sqr_pipeline_avx2(da[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, roots_avx[k]); } } #else // Scalar: mixed-radix unsupported -> fall back to power-of-2 if (radix != 0) N = next_power_of_2(rn); thread_local NttRootsMont roots[3]; for (int k = 0; k < 3; ++k) { if (roots[k].N != N || roots[k].p != ctx.primes[k].p) roots[k].build(ctx.primes[k], N); } const uint64_t* fwd_r[3], *inv_r[3]; uint64_t inv_n_vals[3]; for (int k = 0; k < 3; ++k) { fwd_r[k] = roots[k].roots.data(); inv_r[k] = roots[k].inv_roots.data(); inv_n_vals[k] = roots[k].inv_n; } bool use_mt = (N >= MULTI_THREAD_NTT_THRESHOLD); bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); if (use_mt) { unsigned hw = std::thread::hardware_concurrency(); int T = std::max(2, static_cast(hw > 3 ? (hw - 1) / 3 : 1)); auto f0 = sangi::threadPool().submit([&, N, T]{ ntt_sqr_pipeline_mt(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_r[0], inv_r[0], inv_n_vals[0], T); }); auto f1 = sangi::threadPool().submit([&, N, T]{ ntt_sqr_pipeline_mt(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_r[1], inv_r[1], inv_n_vals[1], T); }); ntt_sqr_pipeline_mt(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_r[2], inv_r[2], inv_n_vals[2], T); f0.get(); f1.get(); } else if (parallel) { auto f0 = sangi::threadPool().submit([&, N]{ ntt_sqr_pipeline_mont(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_r[0], inv_r[0], inv_n_vals[0]); }); auto f1 = sangi::threadPool().submit([&, N]{ ntt_sqr_pipeline_mont(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_r[1], inv_r[1], inv_n_vals[1]); }); ntt_sqr_pipeline_mont(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_r[2], inv_r[2], inv_n_vals[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) ntt_sqr_pipeline_mont(da[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, fwd_r[k], inv_r[k], inv_n_vals[k]); } #endif // CRT reconstruction + carry propagation crt_recompose(rp, rn, da[0], da[1], da[2], rn, ctx.crt); } // ================================================================ // NTT cache: store the forward NTT of a constant operand // ================================================================ // When multiplying by the same value repeatedly, compute the forward NTT only once, // cache it, and subsequent multiplications use only pointwise_mul + INTT. // Ordinary multiplication needs four steps (NTT(A) + NTT(B) + pointwise + INTT), but // with the cache, only three steps suffice (NTT(A) + pointwise + INTT). struct NttCache { std::vector storage_; uint64_t* data[3] = {}; // Forward-NTT results for 3 primes (Montgomery form) size_t ntt_len = 0; // Cached NTT length size_t orig_limbs = 0; // Original limb count void invalidate() { ntt_len = 0; } bool valid(size_t N) const { return ntt_len == N && N > 0; } // Compute and cache the forward NTT of bp[0..bn-1] // roots is shared with the caller's thread_local (avoids rebuilding roots) #ifdef __AVX2__ void build(const uint64_t* bp, size_t bn, size_t N, NttRootsLayered (&roots_ext)[3]) { #else void build(const uint64_t* bp, size_t bn, size_t N, NttRootsMont (&roots_ext)[3]) { #endif auto& ctx = getPrimeNttContext(); ntt_len = N; orig_limbs = bn; storage_.resize(3 * N); for (int k = 0; k < 3; ++k) data[k] = storage_.data() + k * N; // Convert input to Montgomery form for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p; uint64_t pin = ctx.primes[k].p_inv_neg; uint64_t r2 = ctx.primes[k].r2_mod_p; for (size_t i = 0; i < bn; ++i) data[k][i] = to_mont(bp[i] % p, p, pin, r2); for (size_t i = bn; i < N; ++i) data[k][i] = 0; } // Forward NTT (3 primes in parallel) — uses caller's roots bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); #ifdef __AVX2__ if (parallel) { const NttRootsLayered* rp[3] = {&roots_ext[0], &roots_ext[1], &roots_ext[2]}; auto f0 = sangi::threadPool().submit([&, N, rp]{ forward_ntt_mont_avx2(data[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, rp[0]->fwd.data(), rp[0]->fwd_offset.data()); }); auto f1 = sangi::threadPool().submit([&, N, rp]{ forward_ntt_mont_avx2(data[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, rp[1]->fwd.data(), rp[1]->fwd_offset.data()); }); forward_ntt_mont_avx2(data[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, roots_ext[2].fwd.data(), roots_ext[2].fwd_offset.data()); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) forward_ntt_mont_avx2(data[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, roots_ext[k].fwd.data(), roots_ext[k].fwd_offset.data()); } #else const uint64_t* fwd_r[3]; for (int k = 0; k < 3; ++k) fwd_r[k] = roots_ext[k].roots.data(); if (parallel) { auto f0 = sangi::threadPool().submit([&, N]{ forward_ntt_mont(data[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_r[0]); }); auto f1 = sangi::threadPool().submit([&, N]{ forward_ntt_mont(data[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_r[1]); }); forward_ntt_mont(data[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_r[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) forward_ntt_mont(data[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, fwd_r[k]); } #endif } }; // Cached multiplication: rp = ap x bp (forward NTT of bp taken from cache) // Automatically build when the cache is invalid or the NTT length mismatches inline void mul_prime_ntt_cached(uint64_t* rp, const uint64_t* ap, size_t an, const uint64_t* bp, size_t bn, NttCache& cache) { auto& ctx = getPrimeNttContext(); size_t rn = an + bn; size_t N = next_power_of_2(rn); // Allocate NTT data arrays for A (for 3 primes) thread_local std::vector cached_work; size_t total = 3 * N; if (cached_work.size() < total) cached_work.resize(total); uint64_t* da[3]; for (int k = 0; k < 3; ++k) da[k] = cached_work.data() + k * N; // Convert A to Montgomery form for (int k = 0; k < 3; ++k) { uint64_t p = ctx.primes[k].p; uint64_t pin = ctx.primes[k].p_inv_neg; uint64_t r2 = ctx.primes[k].r2_mod_p; for (size_t i = 0; i < an; ++i) da[k][i] = to_mont(ap[i] % p, p, pin, r2); for (size_t i = an; i < N; ++i) da[k][i] = 0; } // Pipeline: NTT(A) + pointwise_mul(A, cached_B) + INTT(A) bool parallel = (N >= PRIME_NTT_PARALLEL_THRESHOLD); #ifdef __AVX2__ thread_local NttRootsLayered cached_roots_avx[3]; for (int k = 0; k < 3; ++k) { if (cached_roots_avx[k].N != N || cached_roots_avx[k].p != ctx.primes[k].p) cached_roots_avx[k].build(ctx.primes[k], N); } // Build the cache (first call or when NTT length changes) — shares roots if (!cache.valid(N)) { cache.build(bp, bn, N, cached_roots_avx); } // Extract raw pointers from thread_local (safe for worker threads) const uint64_t* fwd_d[3], *inv_d[3]; const size_t* fwd_o[3], *inv_o[3]; uint64_t inv_n_vals[3]; for (int k = 0; k < 3; ++k) { fwd_d[k] = cached_roots_avx[k].fwd.data(); fwd_o[k] = cached_roots_avx[k].fwd_offset.data(); inv_d[k] = cached_roots_avx[k].inv.data(); inv_o[k] = cached_roots_avx[k].inv_offset.data(); inv_n_vals[k] = cached_roots_avx[k].inv_n; } if (parallel) { auto f0 = sangi::threadPool().submit([&, N]{ forward_ntt_mont_avx2(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_d[0], fwd_o[0]); pointwise_mul_mont_avx2(da[0], cache.data[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg); inverse_ntt_mont_avx2(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, inv_d[0], inv_o[0], inv_n_vals[0]); }); auto f1 = sangi::threadPool().submit([&, N]{ forward_ntt_mont_avx2(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_d[1], fwd_o[1]); pointwise_mul_mont_avx2(da[1], cache.data[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg); inverse_ntt_mont_avx2(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, inv_d[1], inv_o[1], inv_n_vals[1]); }); forward_ntt_mont_avx2(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_d[2], fwd_o[2]); pointwise_mul_mont_avx2(da[2], cache.data[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg); inverse_ntt_mont_avx2(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, inv_d[2], inv_o[2], inv_n_vals[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { forward_ntt_mont_avx2(da[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, fwd_d[k], fwd_o[k]); pointwise_mul_mont_avx2(da[k], cache.data[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg); inverse_ntt_mont_avx2(da[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, inv_d[k], inv_o[k], inv_n_vals[k]); } } #else thread_local NttRootsMont cached_roots[3]; for (int k = 0; k < 3; ++k) { if (cached_roots[k].N != N || cached_roots[k].p != ctx.primes[k].p) cached_roots[k].build(ctx.primes[k], N); } // Build the cache (first call or when NTT length changes) — shares roots if (!cache.valid(N)) { cache.build(bp, bn, N, cached_roots); } const uint64_t* fwd_r[3], *inv_r[3]; uint64_t inv_n_vals[3]; for (int k = 0; k < 3; ++k) { fwd_r[k] = cached_roots[k].roots.data(); inv_r[k] = cached_roots[k].inv_roots.data(); inv_n_vals[k] = cached_roots[k].inv_n; } if (parallel) { auto f0 = sangi::threadPool().submit([&, N]{ forward_ntt_mont(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, fwd_r[0]); pointwise_mul_mont(da[0], cache.data[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg); inverse_ntt_mont(da[0], N, ctx.primes[0].p, ctx.primes[0].p_inv_neg, inv_r[0], inv_n_vals[0]); }); auto f1 = sangi::threadPool().submit([&, N]{ forward_ntt_mont(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, fwd_r[1]); pointwise_mul_mont(da[1], cache.data[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg); inverse_ntt_mont(da[1], N, ctx.primes[1].p, ctx.primes[1].p_inv_neg, inv_r[1], inv_n_vals[1]); }); forward_ntt_mont(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, fwd_r[2]); pointwise_mul_mont(da[2], cache.data[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg); inverse_ntt_mont(da[2], N, ctx.primes[2].p, ctx.primes[2].p_inv_neg, inv_r[2], inv_n_vals[2]); f0.get(); f1.get(); } else { for (int k = 0; k < 3; ++k) { forward_ntt_mont(da[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, fwd_r[k]); pointwise_mul_mont(da[k], cache.data[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg); inverse_ntt_mont(da[k], N, ctx.primes[k].p, ctx.primes[k].p_inv_neg, inv_r[k], inv_n_vals[k]); } } #endif // CRT reconstruction crt_recompose(rp, rn, da[0], da[1], da[2], rn, ctx.crt); } } // namespace prime_ntt } // namespace sangi