use crate::num::arithmetic::mod_mul::{limbs_invert_limb_u64, mod_preinverted_double};
use crate::num::arithmetic::traits::{ModMul, ModMulAssign};
use crate::num::basic::integers::PrimitiveInt;
use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::num::logic::traits::LeadingZeros;
use crate::unsigned_polynomial::UnsignedPolynomial;
use crate::unsigned_polynomial::arithmetic::mod_add::assert_reduced;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::from_coefficients_trimmed;
use alloc::vec;
pub(crate) const MOD_MUL_KARATSUBA_THRESHOLD: usize = 80;
pub(crate) const MOD_SQUARE_KARATSUBA_THRESHOLD: usize = 256;
pub(crate) struct ModData<T: PrimitiveUnsigned> {
pub(crate) m: T,
inv: u64,
radix: T,
pub(crate) words: u8,
pub(crate) max_terms: usize,
}
impl<T: PrimitiveUnsigned> ModData<T> {
pub(crate) fn new(m: T, terms: usize) -> Self {
assert_ne!(m, T::ZERO);
let words = match T::try_from(terms) {
Ok(terms) => {
let (hi, lo) = T::x_mul_y_to_zz(m - T::ONE, m - T::ONE);
let carry = T::x_mul_y_to_zz(lo, terms).0;
let (top, mid) = T::x_mul_y_to_zz(hi, terms);
let (top, mid) = T::xx_add_yy_to_zz(top, mid, T::ZERO, carry);
if top != T::ZERO {
3
} else if mid != T::ZERO {
2
} else {
1
}
}
Err(_) => 3,
};
let inv = if T::WIDTH == u64::WIDTH {
let m: u64 = m.wrapping_into();
limbs_invert_limb_u64(m << LeadingZeros::leading_zeros(m))
} else {
0
};
Self {
m,
inv,
radix: if T::WIDTH > u64::WIDTH {
m.wrapping_neg() % m
} else {
T::ZERO
},
words,
max_terms: if T::WIDTH > usize::WIDTH {
usize::MAX
} else {
1 << (T::WIDTH - 1)
},
}
}
#[inline]
pub(crate) fn reduce_2(&self, x1: T, x0: T) -> T {
let m = self.m;
if T::WIDTH <= u32::WIDTH {
let x: u64 = x1.wrapping_into();
let x0: u64 = x0.wrapping_into();
let m: u64 = m.wrapping_into();
T::wrapping_from(((x << T::WIDTH) | x0) % m)
} else if T::WIDTH == u64::WIDTH {
T::wrapping_from(mod_preinverted_double::<u64, u128>(
x1.wrapping_into(),
x0.wrapping_into(),
m.wrapping_into(),
self.inv,
))
} else {
let data = T::precompute_mod_mul_data(&m);
(x1 % m)
.mod_mul_precomputed(self.radix, m, &data)
.mod_add(x0 % m, m)
}
}
#[inline]
pub(crate) fn reduce(&self, x2: T, x1: T, x0: T) -> T {
self.reduce_2(self.reduce_2(x2, x1), x0)
}
#[inline]
pub(crate) fn reduce_sum(&self, (x2, x1, x0): (T, T, T)) -> T {
match self.words {
1 => x0 % self.m,
2 => self.reduce_2(x1, x0),
_ => self.reduce(x2, x1, x0),
}
}
}
#[inline]
pub(crate) fn accumulate<T: PrimitiveUnsigned>(acc: &mut (T, T, T), x: T, y: T) {
let (hi, lo) = T::x_mul_y_to_zz(x, y);
let (a2, a1, a0) = *acc;
*acc = T::xxx_add_yyy_to_zzz(a2, a1, a0, T::ZERO, hi, lo);
}
#[inline]
pub(crate) fn column_sum<T: PrimitiveUnsigned>(xs: &[T], ys: &[T], d: &ModData<T>) -> (T, T, T) {
match d.words {
1 => {
let mut sum = T::ZERO;
for (&x, &y) in xs.iter().zip(ys.iter().rev()) {
sum.wrapping_add_assign(x.wrapping_mul(y));
}
(T::ZERO, T::ZERO, sum)
}
2 => {
let (mut hi, mut lo) = (T::ZERO, T::ZERO);
for (&x, &y) in xs.iter().zip(ys.iter().rev()) {
let (p_hi, p_lo) = T::x_mul_y_to_zz(x, y);
(hi, lo) = T::xx_add_yy_to_zz(hi, lo, p_hi, p_lo);
}
(T::ZERO, hi, lo)
}
_ => {
let mut acc = (T::ZERO, T::ZERO, T::ZERO);
for (i, (x_chunk, y_chunk)) in xs
.chunks(d.max_terms)
.zip(ys.rchunks(d.max_terms))
.enumerate()
{
if i != 0 {
acc = (T::ZERO, T::ZERO, d.reduce(acc.0, acc.1, acc.2));
}
for (&x, &y) in x_chunk.iter().zip(y_chunk.iter().rev()) {
accumulate(&mut acc, x, y);
}
}
acc
}
}
}
pub(crate) fn mod_add_assign_slice<T: PrimitiveUnsigned>(xs: &mut [T], ys: &[T], m: T) {
for (x, &y) in xs.iter_mut().zip(ys) {
*x = x.mod_add(y, m);
}
}
pub(crate) fn mod_sub_assign_slice<T: PrimitiveUnsigned>(xs: &mut [T], ys: &[T], m: T) {
for (x, &y) in xs.iter_mut().zip(ys) {
*x = x.mod_sub(y, m);
}
}
pub(crate) fn mod_mul_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<T>,
) {
let n = xs.len();
let m = ys.len();
for (k, o) in out.iter_mut().enumerate() {
let start = k.saturating_sub(m - 1);
let stop = core::cmp::min(k, n - 1);
let acc = column_sum(&xs[start..=stop], &ys[k - stop..=k - start], d);
*o = d.reduce_sum(acc);
}
}
pub(crate) const fn mod_karatsuba_scratch_len(mut n: usize, threshold: usize) -> usize {
let mut len = 0;
while n >= threshold {
let c = n - (n >> 1);
len += (c << 2) - 1;
n = c;
}
len
}
fn mod_mul_karatsuba_balanced<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<T>,
scratch: &mut [T],
) {
let n = xs.len();
if n < MOD_MUL_KARATSUBA_THRESHOLD {
mod_mul_classical(out, xs, ys, d);
return;
}
let m = d.m;
let h = n >> 1;
let c = n - h;
let two_h = h << 1;
let (x0, x1) = xs.split_at(h);
let (y0, y1) = ys.split_at(h);
split_into_chunks_mut!(scratch, c, [x_sum, y_sum], scratch);
let (middle, scratch) = scratch.split_at_mut((c << 1) - 1);
let (low, high) = out.split_at_mut(two_h);
mod_mul_karatsuba_balanced(&mut low[..two_h - 1], x0, y0, d, scratch);
low[two_h - 1] = T::ZERO;
mod_mul_karatsuba_balanced(high, x1, y1, d, scratch);
x_sum.copy_from_slice(x1);
mod_add_assign_slice(x_sum, x0, m);
y_sum.copy_from_slice(y1);
mod_add_assign_slice(y_sum, y0, m);
mod_mul_karatsuba_balanced(middle, x_sum, y_sum, d, scratch);
mod_sub_assign_slice(middle, &out[..two_h - 1], m);
mod_sub_assign_slice(middle, &out[two_h..], m);
mod_add_assign_slice(&mut out[h..], middle, m);
}
pub(crate) fn mod_mul_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<T>,
) {
let (xs, ys) = if xs.len() >= ys.len() {
(xs, ys)
} else {
(ys, xs)
};
let n = xs.len();
let m = ys.len();
if m < MOD_MUL_KARATSUBA_THRESHOLD {
mod_mul_classical(out, xs, ys, d);
return;
}
let mut scratch = vec![T::ZERO; mod_karatsuba_scratch_len(m, MOD_MUL_KARATSUBA_THRESHOLD)];
if n == m {
mod_mul_karatsuba_balanced(out, xs, ys, d, &mut scratch);
return;
}
out.fill(T::ZERO);
let mut product = vec![T::ZERO; (m << 1) - 1];
for (k, piece) in xs.chunks(m).enumerate() {
let product = &mut product[..piece.len() + m - 1];
if piece.len() == m {
mod_mul_karatsuba_balanced(product, piece, ys, d, &mut scratch);
} else {
mod_mul_karatsuba(product, ys, piece, d);
}
mod_add_assign_slice(&mut out[k * m..], product, d.m);
}
}
fn assert_lengths<T>(out: &[T], xs: &[T], ys: &[T]) {
assert!(!xs.is_empty());
assert!(!ys.is_empty());
assert_eq!(out.len(), xs.len() + ys.len() - 1);
}
crate_test_fn! {
#[allow(dead_code)]
mod_mul_to_out_classical<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T], m: T) {
assert_lengths(out, xs, ys);
mod_mul_classical(out, xs, ys, &ModData::new(m, xs.len().min(ys.len())));
}}
crate_test_fn! {
#[allow(dead_code)]
mod_mul_to_out_karatsuba<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T], m: T) {
assert_lengths(out, xs, ys);
mod_mul_karatsuba(out, xs, ys, &ModData::new(m, xs.len().min(ys.len())));
}}
#[doc(hidden)]
pub fn mod_mul_to_out<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T], m: T) {
assert_lengths(out, xs, ys);
mod_mul_karatsuba(out, xs, ys, &ModData::new(m, xs.len().min(ys.len())));
}
pub(crate) fn mod_mul_helper<T: PrimitiveUnsigned>(
xs: &[T],
ys: &[T],
m: T,
) -> UnsignedPolynomial<T> {
if xs.is_empty() || ys.is_empty() {
return UnsignedPolynomial::ZERO;
}
let mut out = vec![T::ZERO; xs.len() + ys.len() - 1];
mod_mul_to_out(&mut out, xs, ys, m);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModMul<Self, T> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_mul(self, other: Self, m: T) -> Self {
assert_reduced(&self, &other, m);
mod_mul_helper(&self.coefficients, &other.coefficients, m)
}
}
impl<T: PrimitiveUnsigned> ModMul<&Self, T> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_mul(self, other: &Self, m: T) -> Self {
assert_reduced(&self, other, m);
mod_mul_helper(&self.coefficients, &other.coefficients, m)
}
}
impl<T: PrimitiveUnsigned> ModMul<UnsignedPolynomial<T>, T> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_mul(self, other: UnsignedPolynomial<T>, m: T) -> UnsignedPolynomial<T> {
assert_reduced(self, &other, m);
mod_mul_helper(&self.coefficients, &other.coefficients, m)
}
}
impl<T: PrimitiveUnsigned> ModMul<&UnsignedPolynomial<T>, T> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_mul(self, other: &UnsignedPolynomial<T>, m: T) -> UnsignedPolynomial<T> {
assert_reduced(self, other, m);
mod_mul_helper(&self.coefficients, &other.coefficients, m)
}
}
impl<T: PrimitiveUnsigned> ModMulAssign<Self, T> for UnsignedPolynomial<T> {
fn mod_mul_assign(&mut self, other: Self, m: T) {
assert_reduced(self, &other, m);
*self = mod_mul_helper(&self.coefficients, &other.coefficients, m);
}
}
impl<T: PrimitiveUnsigned> ModMulAssign<&Self, T> for UnsignedPolynomial<T> {
fn mod_mul_assign(&mut self, other: &Self, m: T) {
assert_reduced(self, other, m);
*self = mod_mul_helper(&self.coefficients, &other.coefficients, m);
}
}