use crate::num::arithmetic::traits::{
ModPowerOf2IsReduced, ModPowerOf2Square, ModPowerOf2SquareAssign,
};
use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::unsigned_polynomial::UnsignedPolynomial;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::{
MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD, add_wrapping_assign, from_coefficients_trimmed,
karatsuba_wrapping_scratch_len, mask_coefficients, sub_wrapping_assign,
};
use alloc::vec;
pub(crate) fn square_classical_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T]) {
out.fill(T::ZERO);
for (i, &x) in xs.iter().enumerate() {
if x != T::ZERO {
for (o, &y) in out[(i << 1) + 1..].iter_mut().zip(&xs[i + 1..]) {
o.wrapping_add_assign(x.wrapping_mul(y));
}
}
}
for o in out.iter_mut() {
*o = o.wrapping_add(*o);
}
for (i, &x) in xs.iter().enumerate() {
out[i << 1].wrapping_add_assign(x.wrapping_mul(x));
}
}
fn square_karatsuba_scratch_wrapping<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
scratch: &mut [T],
) {
let n = xs.len();
if n < MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD {
square_classical_wrapping(out, xs);
return;
}
let h = n >> 1;
let c = n - h;
let two_h = h << 1;
let (x0, x1) = xs.split_at(h);
let (sum, scratch) = scratch.split_at_mut(c);
let (middle, scratch) = scratch.split_at_mut((c << 1) - 1);
let (low, high) = out.split_at_mut(two_h);
square_karatsuba_scratch_wrapping(&mut low[..two_h - 1], x0, scratch);
low[two_h - 1] = T::ZERO;
square_karatsuba_scratch_wrapping(high, x1, scratch);
sum.copy_from_slice(x1);
add_wrapping_assign(sum, x0);
square_karatsuba_scratch_wrapping(middle, sum, scratch);
sub_wrapping_assign(middle, &out[..two_h - 1]);
sub_wrapping_assign(middle, &out[two_h..]);
add_wrapping_assign(&mut out[h..], middle);
}
pub(crate) fn square_karatsuba_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T]) {
if xs.len() < MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD {
square_classical_wrapping(out, xs);
return;
}
let mut scratch = vec![T::ZERO; karatsuba_wrapping_scratch_len(xs.len())];
square_karatsuba_scratch_wrapping(out, xs, &mut scratch);
}
fn assert_lengths<T>(out: &[T], xs: &[T]) {
assert!(!xs.is_empty());
assert_eq!(out.len(), (xs.len() << 1) - 1);
}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_square_to_out_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
pow: u64,
) {
assert_lengths(out, xs);
assert!(pow <= T::WIDTH);
square_classical_wrapping(out, xs);
mask_coefficients(out, pow);
}}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_square_to_out_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
pow: u64,
) {
assert_lengths(out, xs);
assert!(pow <= T::WIDTH);
square_karatsuba_wrapping(out, xs);
mask_coefficients(out, pow);
}}
#[doc(hidden)]
pub fn mod_power_of_2_square_to_out<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], pow: u64) {
assert_lengths(out, xs);
assert!(pow <= T::WIDTH);
square_karatsuba_wrapping(out, xs);
mask_coefficients(out, pow);
}
fn assert_reduced<T: PrimitiveUnsigned>(p: &UnsignedPolynomial<T>, pow: u64) {
assert!(pow <= T::WIDTH);
assert!(
p.mod_power_of_2_is_reduced(pow),
"self must be reduced mod 2^pow, but {p} has a coefficient >= 2^{pow}"
);
}
pub(crate) fn mod_power_of_2_square_helper<T: PrimitiveUnsigned>(
xs: &[T],
pow: u64,
) -> UnsignedPolynomial<T> {
if xs.is_empty() {
return UnsignedPolynomial::ZERO;
}
let mut out = vec![T::ZERO; (xs.len() << 1) - 1];
mod_power_of_2_square_to_out(&mut out, xs, pow);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModPowerOf2Square for UnsignedPolynomial<T> {
type Output = Self;
fn mod_power_of_2_square(self, pow: u64) -> Self {
assert_reduced(&self, pow);
mod_power_of_2_square_helper(&self.coefficients, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2Square for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_power_of_2_square(self, pow: u64) -> UnsignedPolynomial<T> {
assert_reduced(self, pow);
mod_power_of_2_square_helper(&self.coefficients, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2SquareAssign for UnsignedPolynomial<T> {
fn mod_power_of_2_square_assign(&mut self, pow: u64) {
assert_reduced(self, pow);
*self = mod_power_of_2_square_helper(&self.coefficients, pow);
}
}