use std::fmt;
use std::marker::PhantomData;
use std::ops::{Add, AddAssign, Mul, MulAssign, Neg, Sub, SubAssign};
use num::{BigInt, One, ToPrimitive, Zero};
use crate::fp_field::{montyred, negqinv_modr, r_sq_modq, FpField};
pub(crate) trait PrimeList<const N: usize> {
const PRIMES: [u32; N];
}
macro_rules! dispatch_log2r {
($log2r:expr, $a:expr, $b:expr, $p:expr, $neg_inv:expr) => {{
let a = $a;
let b = $b;
let p = $p;
let neg_inv = $neg_inv;
match $log2r {
1 => montyred::<1>(a, b, p, neg_inv),
2 => montyred::<2>(a, b, p, neg_inv),
3 => montyred::<3>(a, b, p, neg_inv),
4 => montyred::<4>(a, b, p, neg_inv),
5 => montyred::<5>(a, b, p, neg_inv),
6 => montyred::<6>(a, b, p, neg_inv),
7 => montyred::<7>(a, b, p, neg_inv),
8 => montyred::<8>(a, b, p, neg_inv),
9 => montyred::<9>(a, b, p, neg_inv),
10 => montyred::<10>(a, b, p, neg_inv),
11 => montyred::<11>(a, b, p, neg_inv),
12 => montyred::<12>(a, b, p, neg_inv),
13 => montyred::<13>(a, b, p, neg_inv),
14 => montyred::<14>(a, b, p, neg_inv),
15 => montyred::<15>(a, b, p, neg_inv),
16 => montyred::<16>(a, b, p, neg_inv),
17 => montyred::<17>(a, b, p, neg_inv),
18 => montyred::<18>(a, b, p, neg_inv),
19 => montyred::<19>(a, b, p, neg_inv),
20 => montyred::<20>(a, b, p, neg_inv),
21 => montyred::<21>(a, b, p, neg_inv),
22 => montyred::<22>(a, b, p, neg_inv),
23 => montyred::<23>(a, b, p, neg_inv),
24 => montyred::<24>(a, b, p, neg_inv),
25 => montyred::<25>(a, b, p, neg_inv),
26 => montyred::<26>(a, b, p, neg_inv),
27 => montyred::<27>(a, b, p, neg_inv),
28 => montyred::<28>(a, b, p, neg_inv),
29 => montyred::<29>(a, b, p, neg_inv),
30 => montyred::<30>(a, b, p, neg_inv),
31 => montyred::<31>(a, b, p, neg_inv),
32 => montyred::<32>(a, b, p, neg_inv),
_ => unreachable!(),
}
}};
}
pub(crate) struct Rns<const N: usize, P: PrimeList<N>> {
residues: [u32; N],
_phantom: PhantomData<P>,
}
impl<const N: usize, P: PrimeList<N>> Clone for Rns<N, P> {
fn clone(&self) -> Self {
*self
}
}
impl<const N: usize, P: PrimeList<N>> Copy for Rns<N, P> {}
impl<const N: usize, P: PrimeList<N>> PartialEq for Rns<N, P> {
fn eq(&self, other: &Self) -> bool {
self.residues == other.residues
}
}
impl<const N: usize, P: PrimeList<N>> Eq for Rns<N, P> {}
impl<const N: usize, P: PrimeList<N>> fmt::Debug for Rns<N, P> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Rns")
.field("residues", &self.residues)
.finish()
}
}
impl<const N: usize, P: PrimeList<N>> Rns<N, P> {
const LOG2R: [u32; N] = {
let primes = P::PRIMES;
let mut a = [0u32; N];
let mut i = 0;
while i < N {
a[i] = primes[i].ilog2() + 1;
i += 1;
}
a
};
const NEG_INV: [u32; N] = {
let primes = P::PRIMES;
let mut a = [0u32; N];
let mut i = 0;
while i < N {
a[i] = negqinv_modr(primes[i]);
i += 1;
}
a
};
const R_SQ: [u32; N] = {
let primes = P::PRIMES;
let mut a = [0u32; N];
let mut i = 0;
while i < N {
a[i] = r_sq_modq(primes[i]);
i += 1;
}
a
};
fn to_mont_at(v: u32, i: usize) -> u32 {
dispatch_log2r!(
Self::LOG2R[i],
v,
Self::R_SQ[i],
P::PRIMES[i],
Self::NEG_INV[i]
)
}
fn from_mont_at(a: u32, i: usize) -> u32 {
dispatch_log2r!(Self::LOG2R[i], a, 1u32, P::PRIMES[i], Self::NEG_INV[i])
}
pub(crate) fn from_u32(v: u32) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
*residues_i = Self::to_mont_at(v % P::PRIMES[i], i);
}
Rns {
residues,
_phantom: PhantomData,
}
}
pub(crate) fn residue(&self, i: usize) -> u32 {
Self::from_mont_at(self.residues[i], i)
}
pub(crate) fn to_garner(self) -> [u32; N] {
let mut u = [0u32; N];
for (i, u_i) in u.iter_mut().enumerate().take(N) {
*u_i = self.residue(i);
}
for i in 0..N {
for j in i + 1..N {
let pi = P::PRIMES[i] as u64;
let pj = P::PRIMES[j] as u64;
let inv = mod_pow(pi % pj, pj - 2, pj);
let ui_mod_pj = u[i] as u64 % pj;
let diff = (u[j] as u64 + pj - ui_mod_pj) % pj;
u[j] = (diff * inv % pj) as u32;
}
}
u
}
pub(crate) fn from_i32(x: i32) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
let p = P::PRIMES[i];
let r = x.rem_euclid(p as i32) as u32;
*residues_i = Self::to_mont_at(r, i);
}
Rns {
residues,
_phantom: PhantomData,
}
}
pub(crate) fn from_i128(x: i128) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
let r = x.rem_euclid(P::PRIMES[i] as i128) as u32;
*residues_i = Self::to_mont_at(r, i);
}
Rns {
residues,
_phantom: PhantomData,
}
}
pub(crate) fn to_i64(self) -> i64 {
let a = self.to_garner();
let mut result = 0u128;
let mut base = 1u128;
for (i, &a_i) in a.iter().enumerate().take(N) {
result = result.wrapping_add(base.wrapping_mul(a_i as u128));
base = base.wrapping_mul(P::PRIMES[i] as u128);
}
if result > base / 2 {
result.wrapping_sub(base) as i64
} else {
result as i64
}
}
pub(crate) fn from_bigint(x: &BigInt) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
let p = P::PRIMES[i];
let mut r = x % &BigInt::from(p);
if r.sign() == num::bigint::Sign::Minus {
r += BigInt::from(p);
}
*residues_i = Self::to_mont_at(r.to_u32().unwrap(), i);
}
Rns {
residues,
_phantom: PhantomData,
}
}
pub(crate) fn to_bigint(self) -> BigInt {
let a = self.to_garner();
let mut result = BigInt::zero();
let mut base = BigInt::one();
for (i, &a_i) in a.iter().enumerate().take(N) {
result += &base * BigInt::from(a_i);
base *= BigInt::from(P::PRIMES[i]);
}
if result > (&base >> 1) {
result -= &base;
}
result
}
pub(crate) fn to_i128(self) -> i128 {
let a = self.to_garner();
let mut result = 0u128;
let mut base = 1u128;
for (i, &a_i) in a.iter().enumerate().take(N) {
result = result.wrapping_add(base.wrapping_mul(a_i as u128));
base = base.wrapping_mul(P::PRIMES[i] as u128);
}
if result > base / 2 {
result.wrapping_sub(base) as i128
} else {
result as i128
}
}
}
pub(crate) const fn mod_pow(mut base: u64, mut exp: u64, modulus: u64) -> u64 {
let mut result = 1u64;
while exp > 0 {
if exp & 1 == 1 {
result = (result as u128 * base as u128 % modulus as u128) as u64;
}
base = (base as u128 * base as u128 % modulus as u128) as u64;
exp >>= 1;
}
result
}
impl<const N: usize, P: PrimeList<N>> Add for Rns<N, P> {
type Output = Self;
fn add(self, rhs: Self) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
let p = P::PRIMES[i] as u64;
let s = self.residues[i] as u64 + rhs.residues[i] as u64;
*residues_i = (if s >= p { s - p } else { s }) as u32;
}
Rns {
residues,
_phantom: PhantomData,
}
}
}
impl<const N: usize, P: PrimeList<N>> AddAssign for Rns<N, P> {
fn add_assign(&mut self, rhs: Self) {
*self = *self + rhs;
}
}
impl<const N: usize, P: PrimeList<N>> Neg for Rns<N, P> {
type Output = Self;
fn neg(self) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
let r = self.residues[i];
*residues_i = if r == 0 { 0 } else { P::PRIMES[i] - r };
}
Rns {
residues,
_phantom: PhantomData,
}
}
}
impl<const N: usize, P: PrimeList<N>> Sub for Rns<N, P> {
type Output = Self;
fn sub(self, rhs: Self) -> Self {
self + -rhs
}
}
impl<const N: usize, P: PrimeList<N>> SubAssign for Rns<N, P> {
fn sub_assign(&mut self, rhs: Self) {
*self = *self - rhs;
}
}
impl<const N: usize, P: PrimeList<N>> Mul for Rns<N, P> {
type Output = Self;
fn mul(self, rhs: Self) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
*residues_i = dispatch_log2r!(
Self::LOG2R[i],
self.residues[i],
rhs.residues[i],
P::PRIMES[i],
Self::NEG_INV[i]
);
}
Rns {
residues,
_phantom: PhantomData,
}
}
}
impl<const N: usize, P: PrimeList<N>> MulAssign for Rns<N, P> {
fn mul_assign(&mut self, rhs: Self) {
*self = *self * rhs;
}
}
impl<const N: usize, P: PrimeList<N>> Zero for Rns<N, P> {
fn zero() -> Self {
Rns {
residues: [0u32; N],
_phantom: PhantomData,
}
}
fn is_zero(&self) -> bool {
self.residues.iter().all(|&r| r == 0)
}
}
impl<const N: usize, P: PrimeList<N>> One for Rns<N, P> {
fn one() -> Self {
Self::from_u32(1)
}
}
pub(crate) trait NttPrimeList<const N: usize>: PrimeList<N> {
const ROOTS_OF_UNITY_2048: [u32; N];
fn ntt_tables(n: usize) -> NttTables<N>
where
Self: Sized,
{
NttTables::new::<Self>(n)
}
}
pub(crate) const fn primitive_root_2n(root_2048: u32, n: usize, p: u32) -> u32 {
debug_assert!(n >= 1 && n <= 1024 && n.is_power_of_two());
let squarings = 10u32 - n.ilog2(); let mut root = root_2048 as u64;
let mut s = 0;
while s < squarings {
root = root * root % p as u64;
s += 1;
}
root as u32
}
pub(crate) fn bitrev_powers_mont(
root_mont: u32,
n: usize,
p: u32,
log2r: u32,
neg_inv: u32,
r_sq: u32,
) -> Vec<u32> {
let one_mont = dispatch_log2r!(log2r, 1u32, r_sq, p, neg_inv);
let mut array = vec![0u32; n];
let mut alpha = one_mont;
for a in array.iter_mut() {
*a = alpha;
alpha = dispatch_log2r!(log2r, alpha, root_mont, p, neg_inv);
}
crate::cyclotomic_fourier::bitreverse_array(&mut array);
array
}
const fn bitrev_powers_mont_const<const DEG: usize>(
root_mont: u32,
p: u32,
log2r: u32,
neg_inv: u32,
r_sq: u32,
) -> [u32; DEG] {
let one_mont = dispatch_log2r!(log2r, 1u32, r_sq, p, neg_inv);
let mut array = [0u32; DEG];
let mut alpha = one_mont;
let mut idx = 0;
while idx < DEG {
array[idx] = alpha;
alpha = dispatch_log2r!(log2r, alpha, root_mont, p, neg_inv);
idx += 1;
}
let mut i = 0;
while i < DEG {
let j = crate::cyclotomic_fourier::bitreverse_index(i, DEG);
if i < j {
let tmp = array[i];
array[i] = array[j];
array[j] = tmp;
}
i += 1;
}
array
}
const fn ntt_tables_const<const N: usize, const DEG: usize, P: NttPrimeList<N>>(
) -> ([[u32; DEG]; N], [[u32; DEG]; N], [u32; N]) {
let mut psi_rev = [[0u32; DEG]; N];
let mut psi_inv_rev = [[0u32; DEG]; N];
let mut ninv_mont = [0u32; N];
let mut i = 0;
while i < N {
let p = P::PRIMES[i];
let log2r = p.ilog2() + 1;
let neg_inv = negqinv_modr(p);
let r_sq = r_sq_modq(p);
let root_2n = primitive_root_2n(P::ROOTS_OF_UNITY_2048[i], DEG, p);
let root_2n_mont = dispatch_log2r!(log2r, root_2n, r_sq, p, neg_inv);
psi_rev[i] = bitrev_powers_mont_const::<DEG>(root_2n_mont, p, log2r, neg_inv, r_sq);
let root_inv = mod_pow(root_2n as u64, (p - 2) as u64, p as u64) as u32;
let root_inv_mont = dispatch_log2r!(log2r, root_inv, r_sq, p, neg_inv);
psi_inv_rev[i] = bitrev_powers_mont_const::<DEG>(root_inv_mont, p, log2r, neg_inv, r_sq);
let n_inv = mod_pow(DEG as u64, (p - 2) as u64, p as u64) as u32;
ninv_mont[i] = dispatch_log2r!(log2r, n_inv, r_sq, p, neg_inv);
i += 1;
}
(psi_rev, psi_inv_rev, ninv_mont)
}
pub(crate) struct ConstNttTables<const N: usize, const DEG: usize> {
psi_rev: [[u32; DEG]; N],
psi_inv_rev: [[u32; DEG]; N],
ninv_mont: [u32; N],
}
impl<const N: usize, const DEG: usize> ConstNttTables<N, DEG> {
pub(crate) const fn build<P: NttPrimeList<N>>() -> Self {
let (psi_rev, psi_inv_rev, ninv_mont) = ntt_tables_const::<N, DEG, P>();
Self {
psi_rev,
psi_inv_rev,
ninv_mont,
}
}
}
pub(crate) fn montmul_dyn(a: u32, b: u32, p: u32, log2r: u32, neg_inv: u32) -> u32 {
dispatch_log2r!(log2r, a, b, p, neg_inv)
}
pub(crate) fn ntt_u32(a: &mut [u32], psi_rev: &[u32], p: u32, log2r: u32, neg_inv: u32) {
let n = a.len();
let mut t = n;
let mut m = 1;
while m < n {
t >>= 1;
for i in 0..m {
let j1 = 2 * i * t;
let s = psi_rev[m + i];
for j in j1..j1 + t {
let u = a[j];
let v = dispatch_log2r!(log2r, a[j + t], s, p, neg_inv);
let sum = u as u64 + v as u64;
a[j] = if sum >= p as u64 {
(sum - p as u64) as u32
} else {
sum as u32
};
a[j + t] = if u >= v { u - v } else { u + p - v };
}
}
m <<= 1;
}
}
pub(crate) fn intt_u32(
a: &mut [u32],
psi_inv_rev: &[u32],
ninv_mont: u32,
p: u32,
log2r: u32,
neg_inv: u32,
) {
let n = a.len();
let mut t = 1;
let mut m = n;
while m > 1 {
let h = m / 2;
let mut j1 = 0;
for i in 0..h {
let s = psi_inv_rev[h + i];
for j in j1..j1 + t {
let u = a[j];
let v = a[j + t];
let sum = u as u64 + v as u64;
a[j] = if sum >= p as u64 {
(sum - p as u64) as u32
} else {
sum as u32
};
let sub = if u >= v { u - v } else { u + p - v };
a[j + t] = dispatch_log2r!(log2r, sub, s, p, neg_inv);
}
j1 += 2 * t;
}
t <<= 1;
m >>= 1;
}
for ai in a.iter_mut() {
*ai = dispatch_log2r!(log2r, *ai, ninv_mont, p, neg_inv);
}
}
pub(crate) struct NttTables<const N: usize> {
n: usize,
primes: [u32; N],
log2r: [u32; N],
neg_inv: [u32; N],
psi_rev: [Vec<u32>; N],
psi_inv_rev: [Vec<u32>; N],
ninv_mont: [u32; N],
}
impl<const N: usize> NttTables<N> {
pub(crate) fn new<P: NttPrimeList<N>>(n: usize) -> Self {
debug_assert!((1..=1024).contains(&n) && n.is_power_of_two());
let primes = P::PRIMES;
let log2r: [u32; N] = std::array::from_fn(|i| primes[i].ilog2() + 1);
let neg_inv: [u32; N] = std::array::from_fn(|i| negqinv_modr(primes[i]));
let psi_rev: [Vec<u32>; N] = std::array::from_fn(|i| {
let (p, lr, ni, r_sq) = (primes[i], log2r[i], neg_inv[i], r_sq_modq(primes[i]));
let root_2n = primitive_root_2n(P::ROOTS_OF_UNITY_2048[i], n, p);
let root_2n_mont = dispatch_log2r!(lr, root_2n, r_sq, p, ni);
bitrev_powers_mont(root_2n_mont, n, p, lr, ni, r_sq)
});
let psi_inv_rev: [Vec<u32>; N] = std::array::from_fn(|i| {
let (p, lr, ni, r_sq) = (primes[i], log2r[i], neg_inv[i], r_sq_modq(primes[i]));
let root_2n = primitive_root_2n(P::ROOTS_OF_UNITY_2048[i], n, p);
let root_inv = mod_pow(root_2n as u64, (p - 2) as u64, p as u64) as u32;
let root_inv_mont = dispatch_log2r!(lr, root_inv, r_sq, p, ni);
bitrev_powers_mont(root_inv_mont, n, p, lr, ni, r_sq)
});
let ninv_mont: [u32; N] = std::array::from_fn(|i| {
let (p, lr, ni, r_sq) = (primes[i], log2r[i], neg_inv[i], r_sq_modq(primes[i]));
let n_inv = mod_pow(n as u64, (p - 2) as u64, p as u64) as u32;
dispatch_log2r!(lr, n_inv, r_sq, p, ni)
});
Self {
n,
primes,
log2r,
neg_inv,
psi_rev,
psi_inv_rev,
ninv_mont,
}
}
pub(crate) fn from_const<P: NttPrimeList<N>, const DEG: usize>(
c: &ConstNttTables<N, DEG>,
) -> Self {
let primes = P::PRIMES;
let log2r: [u32; N] = std::array::from_fn(|i| primes[i].ilog2() + 1);
let neg_inv: [u32; N] = std::array::from_fn(|i| negqinv_modr(primes[i]));
let psi_rev: [Vec<u32>; N] = std::array::from_fn(|i| c.psi_rev[i].to_vec());
let psi_inv_rev: [Vec<u32>; N] = std::array::from_fn(|i| c.psi_inv_rev[i].to_vec());
Self {
n: DEG,
primes,
log2r,
neg_inv,
psi_rev,
psi_inv_rev,
ninv_mont: c.ninv_mont,
}
}
}
pub(crate) fn ntt_inplace_cached<const N: usize, P: NttPrimeList<N>>(
coeffs: &mut [Rns<N, P>],
tables: &NttTables<N>,
) {
debug_assert_eq!(coeffs.len(), tables.n);
let mut scratch: Vec<u32> = vec![0u32; coeffs.len()];
for prime_idx in 0..N {
for (s, r) in scratch.iter_mut().zip(coeffs.iter()) {
*s = r.residues[prime_idx];
}
ntt_u32(
&mut scratch,
&tables.psi_rev[prime_idx],
tables.primes[prime_idx],
tables.log2r[prime_idx],
tables.neg_inv[prime_idx],
);
for (coeff, &val) in coeffs.iter_mut().zip(scratch.iter()) {
coeff.residues[prime_idx] = val;
}
}
}
pub(crate) fn intt_inplace_cached<const N: usize, P: NttPrimeList<N>>(
coeffs: &mut [Rns<N, P>],
tables: &NttTables<N>,
) {
debug_assert_eq!(coeffs.len(), tables.n);
let mut scratch: Vec<u32> = vec![0u32; coeffs.len()];
for prime_idx in 0..N {
for (s, r) in scratch.iter_mut().zip(coeffs.iter()) {
*s = r.residues[prime_idx];
}
intt_u32(
&mut scratch,
&tables.psi_inv_rev[prime_idx],
tables.ninv_mont[prime_idx],
tables.primes[prime_idx],
tables.log2r[prime_idx],
tables.neg_inv[prime_idx],
);
for (coeff, &val) in coeffs.iter_mut().zip(scratch.iter()) {
coeff.residues[prime_idx] = val;
}
}
}
macro_rules! ntt_prime_list {
($(#[$meta:meta])* $name:ident, $k:literal, $deg:literal, [$($p:literal),+ $(,)?]) => {
$(#[$meta])*
pub(crate) struct $name;
impl PrimeList<$k> for $name {
const PRIMES: [u32; $k] = [$($p),+];
}
impl NttPrimeList<$k> for $name {
const ROOTS_OF_UNITY_2048: [u32; $k] =
[$(FpField::<$p>::primitive_nth_root_of_unity(2048).value()),+];
fn ntt_tables(n: usize) -> NttTables<$k> {
static TABLES: ConstNttTables<$k, $deg> = ConstNttTables::build::<$name>();
if n == $deg {
NttTables::from_const::<$name, $deg>(&TABLES)
} else {
NttTables::new::<$name>(n)
}
}
}
};
}
ntt_prime_list! {
NttPrimes24Bit2, 2, 512, [8_404_993, 8_427_521]
}
ntt_prime_list! {
NttPrimes24Bit4, 4, 256, [8_404_993, 8_427_521, 8_441_857, 8_452_097]
}
ntt_prime_list! {
NttPrimes24Bit5, 5, 128, [8_404_993, 8_427_521, 8_441_857, 8_452_097, 8_466_433]
}
ntt_prime_list! {
NttPrimes24Bit8, 8, 64,
[8_404_993, 8_427_521, 8_441_857, 8_452_097, 8_466_433, 8_513_537, 8_519_681, 8_527_873]
}
#[cfg(test)]
mod tests {
use super::*;
use crate::fp_field::FpField;
struct SmallPrimes;
impl PrimeList<3> for SmallPrimes {
const PRIMES: [u32; 3] = [786_433, 998_244_353, 1_073_754_113];
}
impl NttPrimeList<3> for SmallPrimes {
const ROOTS_OF_UNITY_2048: [u32; 3] = [
FpField::<786_433>::primitive_nth_root_of_unity(2048).value(),
FpField::<998_244_353>::primitive_nth_root_of_unity(2048).value(),
FpField::<1_073_754_113>::primitive_nth_root_of_unity(2048).value(),
];
}
type Rns3 = Rns<3, SmallPrimes>;
struct LargePrimes;
impl PrimeList<2> for LargePrimes {
const PRIMES: [u32; 2] = [1_073_754_113, 4_294_967_291];
}
type Rns2L = Rns<2, LargePrimes>;
impl<const N: usize, P: PrimeList<N>> Rns<N, P> {
fn from_u64(v: u64) -> Self {
let mut residues = [0u32; N];
for (i, residues_i) in residues.iter_mut().enumerate().take(N) {
*residues_i = Self::to_mont_at((v % P::PRIMES[i] as u64) as u32, i);
}
Rns {
residues,
_phantom: PhantomData,
}
}
fn to_u64(self) -> u64 {
let a = self.to_garner();
let mut result = 0u64;
let mut base = 1u64;
for (i, &a_i) in a.iter().enumerate().take(N) {
result = result.wrapping_add(base.wrapping_mul(a_i as u64));
base = base.wrapping_mul(P::PRIMES[i] as u64);
}
result
}
}
#[test]
fn const_tables_match_runtime_build() {
const TABLES: ([[u32; 128]; 5], [[u32; 128]; 5], [u32; 5]) =
ntt_tables_const::<5, 128, NttPrimes24Bit5>();
let runtime = NttTables::<5>::new::<NttPrimes24Bit5>(128);
for i in 0..5 {
assert_eq!(
TABLES.0[i].as_slice(),
runtime.psi_rev[i].as_slice(),
"psi_rev[{i}]"
);
assert_eq!(
TABLES.1[i].as_slice(),
runtime.psi_inv_rev[i].as_slice(),
"psi_inv_rev[{i}]"
);
}
assert_eq!(TABLES.2, runtime.ninv_mont, "ninv_mont");
}
#[test]
fn residues_equal_naive_reduction() {
let primes = SmallPrimes::PRIMES;
for v in [0u32, 1, 12345, 786_432, 998_244_352, 1_073_754_112] {
let r = Rns3::from_u32(v);
for (i, &primes_i) in primes.iter().enumerate() {
assert_eq!(r.residue(i), v % primes_i, "v={v}, i={i}");
}
}
}
#[test]
fn from_u64_residues_correct() {
let primes = SmallPrimes::PRIMES;
let v: u64 = 0x_DEAD_BEEF_1234_5678;
let r = Rns3::from_u64(v);
for (i, &primes_i) in primes.iter().enumerate() {
assert_eq!(r.residue(i), (v % primes_i as u64) as u32, "i={i}");
}
}
#[test]
fn to_u64_roundtrip() {
for v in [
0u64,
1,
999_999_999,
0xFFFF_FFFF,
0x1_0000_0000,
0xDEAD_BEEF_1234,
] {
assert_eq!(Rns3::from_u64(v).to_u64(), v, "v={v}");
}
}
#[test]
fn add_agrees_with_naive() {
let primes = SmallPrimes::PRIMES;
let pairs = [
(0u32, 1u32),
(999, 1_073_754_112),
(786_432, 786_432),
(12345, 67890),
];
for (a, b) in pairs {
let ra = Rns3::from_u32(a);
let rb = Rns3::from_u32(b);
let rc = ra + rb;
for (i, primes_i) in primes.iter().enumerate() {
assert_eq!(
rc.residue(i),
(a as u64 + b as u64) as u32 % primes_i,
"a={a}, b={b}, i={i}"
);
}
}
}
#[test]
fn sub_agrees_with_naive() {
let primes = SmallPrimes::PRIMES;
let pairs = [(10u32, 3u32), (0, 1), (1_073_754_112, 1_073_754_112)];
for (a, b) in pairs {
let ra = Rns3::from_u32(a);
let rb = Rns3::from_u32(b);
let rc = ra - rb;
for (i, primes_i) in primes.iter().enumerate() {
let expected = ((a as i64 - b as i64).rem_euclid(*primes_i as i64)) as u32;
assert_eq!(rc.residue(i), expected, "a={a}, b={b}, i={i}");
}
}
}
#[test]
fn mul_agrees_with_naive() {
let primes = SmallPrimes::PRIMES;
let pairs = [(0u32, 0u32), (1, 1), (2, 3), (786_432, 2), (999, 12345)];
for (a, b) in pairs {
let ra = Rns3::from_u32(a);
let rb = Rns3::from_u32(b);
let rc = ra * rb;
for (i, &primes_i) in primes.iter().enumerate() {
let expected = (a as u64 * b as u64 % primes_i as u64) as u32;
assert_eq!(rc.residue(i), expected, "a={a}, b={b}, i={i}");
}
}
}
#[test]
fn zero_and_one() {
let z = Rns3::zero();
assert!(z.is_zero());
for i in 0..3 {
assert_eq!(z.residue(i), 0, "i={i}");
}
let o = Rns3::one();
for i in 0..3 {
assert_eq!(o.residue(i), 1, "i={i}");
}
}
#[test]
fn garner_reconstruction_matches_to_u64() {
for v in [0u64, 1, 42, 0xDEAD_BEEF, 0x00FF_FFFF_FFFF] {
let r = Rns3::from_u64(v);
let a = r.to_garner();
let primes = SmallPrimes::PRIMES;
let mut horner = a[2] as u64;
horner = horner * primes[1] as u64 + a[1] as u64;
horner = horner * primes[0] as u64 + a[0] as u64;
assert_eq!(horner, v, "v={v}");
}
}
#[test]
fn from_i32_positive() {
let primes = SmallPrimes::PRIMES;
for v in [0i32, 1, 12345, 786_432] {
let r = Rns3::from_i32(v);
for (i, &primes_i) in primes.iter().enumerate() {
assert_eq!(r.residue(i), v as u32 % primes_i, "v={v}, i={i}");
}
}
}
#[test]
fn from_i32_negative() {
let primes = SmallPrimes::PRIMES;
for v in [-1i32, -12345, -786_432] {
let r = Rns3::from_i32(v);
for (i, &primes_i) in primes.iter().enumerate() {
let expected = v.rem_euclid(primes_i as i32) as u32;
assert_eq!(r.residue(i), expected, "v={v}, i={i}");
}
}
}
#[test]
fn to_i64_roundtrip() {
for v in [-1000i64, -1, 0, 1, 999, 786_432, -786_433] {
let r = Rns3::from_i32(v as i32);
assert_eq!(r.to_i64(), v, "v={v}");
}
}
#[test]
fn ntt_intt_roundtrip() {
use super::{intt_inplace_cached, ntt_inplace_cached, NttTables};
let n = 16;
let tables = NttTables::<3>::new::<SmallPrimes>(n);
let orig: Vec<Rns3> = (0..n as i32).map(Rns3::from_i32).collect();
let mut buf = orig.clone();
ntt_inplace_cached::<3, SmallPrimes>(&mut buf, &tables);
assert_ne!(buf[0].residue(0), orig[0].residue(0));
intt_inplace_cached::<3, SmallPrimes>(&mut buf, &tables);
for (i, (a, b)) in orig.iter().zip(buf.iter()).enumerate() {
for pi in 0..3 {
assert_eq!(a.residue(pi), b.residue(pi), "coeff={i}, prime={pi}");
}
}
}
#[test]
fn ntt_multiplication_mod_cyclotomic() {
use super::{intt_inplace_cached, ntt_inplace_cached, NttTables};
let n = 4;
let tables = NttTables::<3>::new::<SmallPrimes>(n);
let a_coeffs = [1i32, 1, 0, 0];
let b_coeffs = [1i32, 1, 0, 0];
let mut a: Vec<Rns3> = a_coeffs.iter().copied().map(Rns3::from_i32).collect();
let mut b: Vec<Rns3> = b_coeffs.iter().copied().map(Rns3::from_i32).collect();
ntt_inplace_cached::<3, SmallPrimes>(&mut a, &tables);
ntt_inplace_cached::<3, SmallPrimes>(&mut b, &tables);
let mut c: Vec<Rns3> = a.iter().zip(b.iter()).map(|(&x, &y)| x * y).collect();
intt_inplace_cached::<3, SmallPrimes>(&mut c, &tables);
let expected = [1i32, 2, 1, 0];
for (i, (e, r)) in expected.iter().zip(c.iter()).enumerate() {
assert_eq!(*e, r.to_i64() as i32, "coeff={i}");
}
}
#[test]
fn large_prime_u128_branch() {
let v: u64 = 3_000_000_000;
let r = Rns2L::from_u64(v);
assert_eq!(r.residue(0), v as u32 % 1_073_754_113);
assert_eq!(r.residue(1), (v % 4_294_967_291) as u32);
assert_eq!(r.to_u64(), v);
let a = Rns2L::from_u64(1_000_000_007);
let b = Rns2L::from_u64(2_000_000_003);
let c = a * b;
for i in 0..2 {
let p = LargePrimes::PRIMES[i] as u64;
let expected = (1_000_000_007u64 * 2_000_000_003u64 % p) as u32;
assert_eq!(c.residue(i), expected, "i={i}");
}
}
}