use std::collections::HashMap;
use std::sync::{Arc, LazyLock, Mutex};
use crate::fp_field::{negqinv_modr, r_sq_modq};
use crate::multiword_int::MultiwordInt;
use crate::multiword_poly::MultiwordPoly;
use crate::rns::{bitrev_powers_mont, intt_u32, mod_pow, montmul_dyn, ntt_u32, primitive_root_2n};
static PRIME_POOL: LazyLock<Mutex<Vec<u32>>> = LazyLock::new(|| Mutex::new(Vec::new()));
fn is_prime_u32(n: u32) -> bool {
if n < 2 {
return false;
}
if n.is_multiple_of(2) {
return n == 2;
}
let mut i = 3u64;
while i * i <= n as u64 {
if (n as u64).is_multiple_of(i) {
return false;
}
i += 2;
}
true
}
fn ntt_primes(k: usize) -> Vec<u32> {
let mut pool = PRIME_POOL.lock().unwrap();
if pool.len() < k {
let mut cand: i64 = ((1i64 << 24) - 2048) + 1;
if let Some(&last) = pool.last() {
cand = last as i64 - 2048;
}
while pool.len() < k && cand > 1 {
if is_prime_u32(cand as u32) {
pool.push(cand as u32);
}
cand -= 2048;
}
assert!(
pool.len() >= k,
"exhausted 24-bit NTT primes: needed {k}, found {}",
pool.len()
);
}
pool[..k].to_vec()
}
fn primitive_root_2048(p: u32) -> u32 {
let order = (p - 1) as u64;
let exp = order / 2048;
let mut g = 2u64;
loop {
let root = mod_pow(g, exp, p as u64);
if root != 1 && mod_pow(root, 1024, p as u64) != 1 {
return root as u32;
}
g += 1;
}
}
pub(crate) struct Transformed {
ntt: Vec<Vec<u32>>,
}
type NttCache = HashMap<(usize, usize), Arc<RuntimeNtt>>;
static NTT_CACHE: LazyLock<Mutex<NttCache>> = LazyLock::new(|| Mutex::new(HashMap::new()));
pub(crate) struct RuntimeNtt {
n: usize,
primes: Vec<u32>,
log2r: Vec<u32>,
neg_inv: Vec<u32>,
r_sq: Vec<u32>, r64: Vec<u32>, psi_rev: Vec<Vec<u32>>, psi_inv_rev: Vec<Vec<u32>>, ninv_mont: Vec<u32>, garner_inv: Vec<u32>,
}
impl RuntimeNtt {
pub(crate) fn new(n: usize, k: usize) -> Self {
debug_assert!(n.is_power_of_two());
let primes = ntt_primes(k);
let log2r: Vec<u32> = primes.iter().map(|&p| p.ilog2() + 1).collect();
let neg_inv: Vec<u32> = primes.iter().map(|&p| negqinv_modr(p)).collect();
let r_sq: Vec<u32> = primes.iter().map(|&p| r_sq_modq(p)).collect();
let r64: Vec<u32> = primes
.iter()
.map(|&p| ((1u128 << 64) % p as u128) as u32)
.collect();
let mut psi_rev = Vec::with_capacity(k);
let mut psi_inv_rev = Vec::with_capacity(k);
let mut ninv_mont = Vec::with_capacity(k);
for i in 0..k {
let p = primes[i];
let root_2n = primitive_root_2n(primitive_root_2048(p), n, p);
let root_2n_mont = montmul_dyn(root_2n, r_sq[i], p, log2r[i], neg_inv[i]);
psi_rev.push(bitrev_powers_mont(
root_2n_mont,
n,
p,
log2r[i],
neg_inv[i],
r_sq[i],
));
let root_inv = mod_pow(root_2n as u64, (p - 2) as u64, p as u64) as u32;
let root_inv_mont = montmul_dyn(root_inv, r_sq[i], p, log2r[i], neg_inv[i]);
psi_inv_rev.push(bitrev_powers_mont(
root_inv_mont,
n,
p,
log2r[i],
neg_inv[i],
r_sq[i],
));
let n_inv = mod_pow(n as u64, (p - 2) as u64, p as u64) as u32;
ninv_mont.push(montmul_dyn(n_inv, r_sq[i], p, log2r[i], neg_inv[i]));
}
let mut garner_inv = vec![0u32; k * k];
for i in 0..k {
for j in (i + 1)..k {
let pj = primes[j] as u64;
garner_inv[i * k + j] = mod_pow(primes[i] as u64 % pj, pj - 2, pj) as u32;
}
}
Self {
n,
primes,
log2r,
neg_inv,
r_sq,
r64,
psi_rev,
psi_inv_rev,
ninv_mont,
garner_inv,
}
}
pub(crate) fn cached(n: usize, k: usize) -> Arc<RuntimeNtt> {
let mut cache = NTT_CACHE.lock().unwrap();
cache
.entry((n, k))
.or_insert_with(|| Arc::new(RuntimeNtt::new(n, k)))
.clone()
}
pub(crate) fn primes(&self) -> &[u32] {
&self.primes
}
fn reduce_signed(&self, limbs: &[u64], pi: usize) -> u32 {
let p = self.primes[pi] as u128;
let r64 = self.r64[pi] as u128;
let mut acc = 0u128;
for &limb in limbs.iter().rev() {
acc = (acc * r64 + (limb as u128 % p)) % p;
}
if crate::multiword_int::is_negative(limbs) {
let two_pow = mod_pow(2, 64 * limbs.len() as u64, self.primes[pi] as u64) as u128;
acc = (acc + p - two_pow) % p;
}
acc as u32
}
pub(crate) fn forward(&self, a: &MultiwordPoly) -> Transformed {
debug_assert_eq!(a.n(), self.n);
let n = self.n;
let mut ntt = Vec::with_capacity(self.primes.len());
for pi in 0..self.primes.len() {
let (p, log2r, neg_inv) = (self.primes[pi], self.log2r[pi], self.neg_inv[pi]);
let mut ar = vec![0u32; n];
for (i, ar_i) in ar.iter_mut().enumerate().take(n) {
*ar_i = montmul_dyn(
self.reduce_signed(a.coeff(i), pi),
self.r_sq[pi],
p,
log2r,
neg_inv,
);
}
ntt_u32(&mut ar, &self.psi_rev[pi], p, log2r, neg_inv);
ntt.push(ar);
}
Transformed { ntt }
}
pub(crate) fn mul_transformed(
&self,
kp: &MultiwordPoly,
h: &Transformed,
out_w: usize,
) -> MultiwordPoly {
debug_assert!(kp.n() == self.n && h.ntt.len() == self.primes.len());
let n = self.n;
let k = self.primes.len();
let mut residues = vec![0u32; n * k];
let mut kr = vec![0u32; n];
for pi in 0..k {
let (p, log2r, neg_inv) = (self.primes[pi], self.log2r[pi], self.neg_inv[pi]);
for (i, kr_i) in kr.iter_mut().enumerate().take(n) {
*kr_i = montmul_dyn(
self.reduce_signed(kp.coeff(i), pi),
self.r_sq[pi],
p,
log2r,
neg_inv,
);
}
ntt_u32(&mut kr, &self.psi_rev[pi], p, log2r, neg_inv);
for (i, kr_i) in kr.iter_mut().enumerate().take(n) {
*kr_i = montmul_dyn(*kr_i, h.ntt[pi][i], p, log2r, neg_inv);
}
intt_u32(
&mut kr,
&self.psi_inv_rev[pi],
self.ninv_mont[pi],
p,
log2r,
neg_inv,
);
for i in 0..n {
residues[i * k + pi] = montmul_dyn(kr[i], 1, p, log2r, neg_inv);
}
}
let mut out = MultiwordPoly::zeros(n, out_w);
let mut digits = vec![0u32; k];
for i in 0..n {
digits.copy_from_slice(&residues[i * k..(i + 1) * k]);
self.garner_in_place(&mut digits);
let value = MultiwordInt::from_garner_digits(&digits, &self.primes, out_w);
out.coeff_mut(i).copy_from_slice(value.limbs());
}
out
}
pub(crate) fn negacyclic_mul(
&self,
a: &MultiwordPoly,
b: &MultiwordPoly,
out_w: usize,
) -> MultiwordPoly {
self.mul_transformed(a, &self.forward(b), out_w)
}
fn garner_in_place(&self, u: &mut [u32]) {
let k = self.primes.len();
for i in 0..k {
for j in (i + 1)..k {
let pj = self.primes[j] as u64;
let inv = self.garner_inv[i * k + j] as u64;
let diff = (u[j] as u64 + pj - (u[i] as u64 % pj)) % pj;
u[j] = (diff * inv % pj) as u32;
}
}
}
}
#[cfg(test)]
mod tests {
use super::RuntimeNtt;
use crate::multiword_poly::MultiwordPoly;
use crate::polynomial::Polynomial;
use num::BigInt;
use rand::{rngs::StdRng, RngExt, SeedableRng};
fn negacyclic_bigint(a: &Polynomial<BigInt>, b: &Polynomial<BigInt>) -> Polynomial<BigInt> {
let n = a.coefficients.len();
let mut out = vec![BigInt::from(0); n];
for i in 0..n {
for j in 0..n {
let prod = &a.coefficients[i] * &b.coefficients[j];
let k = i + j;
if k < n {
out[k] += ∏
} else {
out[k - n] -= ∏
}
}
}
Polynomial::new(out)
}
fn rand_poly(rng: &mut StdRng, n: usize, bits: u32) -> Polynomial<BigInt> {
let words = bits.div_ceil(32);
let excess = words * 32 - bits; Polynomial::new(
(0..n)
.map(|_| {
let mut v = BigInt::from(0);
for _ in 0..words {
v = (v << 32) + BigInt::from(rng.random::<u32>());
}
v >>= excess;
if rng.random::<bool>() {
-v
} else {
v
}
})
.collect(),
)
}
fn check(n: usize, bits: u32, k: usize, out_w: usize, seed: u64) {
let mut rng = StdRng::seed_from_u64(seed);
let ctx = RuntimeNtt::new(n, k);
for _ in 0..20 {
let a = rand_poly(&mut rng, n, bits);
let b = rand_poly(&mut rng, n, bits);
let want = negacyclic_bigint(&a, &b);
let ma = MultiwordPoly::from_bigint_poly(&a, out_w);
let mb = MultiwordPoly::from_bigint_poly(&b, out_w);
let got = ctx.negacyclic_mul(&ma, &mb, out_w);
assert_eq!(
got.to_bigint_poly().coefficients,
want.coefficients,
"n={n} bits={bits} k={k}"
);
}
}
#[test]
fn small_operands_few_primes() {
check(128, 50, 5, 4, 1);
}
#[test]
fn medium_operands() {
check(32, 200, 19, 8, 2);
}
#[test]
fn deep_small_n_many_primes() {
check(4, 600, 54, 20, 3);
}
#[test]
fn asymmetric_widths() {
let mut rng = StdRng::seed_from_u64(4);
let n = 8;
let ctx = RuntimeNtt::new(n, 40);
for _ in 0..20 {
let a = rand_poly(&mut rng, n, 400);
let b = rand_poly(&mut rng, n, 60);
let want = negacyclic_bigint(&a, &b);
let ma = MultiwordPoly::from_bigint_poly(&a, 16);
let mb = MultiwordPoly::from_bigint_poly(&b, 16);
let got = ctx.negacyclic_mul(&ma, &mb, 16);
assert_eq!(got.to_bigint_poly().coefficients, want.coefficients);
}
}
}