use crate::num::arithmetic::traits::ModPowerOf2IsReduced;
use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::polynomial::{ModPowerOf2SquareTruncated, ModPowerOf2SquareTruncatedAssign};
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,
mask_coefficients,
};
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul_truncated::*;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_square::square_karatsuba_wrapping;
use alloc::vec;
use core::cmp::min;
pub(crate) fn square_truncated_classical_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T]) {
let len = out.len();
out.fill(T::ZERO);
for (i, &x) in xs.iter().enumerate() {
if (i << 1) + 1 >= len {
break;
}
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 (o, &x) in out.iter_mut().step_by(2).zip(xs) {
o.wrapping_add_assign(x.wrapping_mul(x));
}
}
pub(crate) fn square_truncated_karatsuba_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T]) {
let len = out.len();
let xs = &xs[..min(xs.len(), len)];
let n = xs.len();
let full_len = (n << 1) - 1;
if full_len <= len {
square_karatsuba_wrapping(&mut out[..full_len], xs);
out[full_len..].fill(T::ZERO);
return;
}
if n < MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD {
square_truncated_classical_wrapping(out, xs);
return;
}
let h = len.div_ceil(2);
let x0 = &xs[..min(h, n)];
square_truncated_karatsuba_wrapping(out, x0);
if n > h {
let mut cross = vec![T::ZERO; len - h];
mul_truncated_karatsuba_wrapping(&mut cross, x0, &xs[h..]);
add_wrapping_assign(&mut out[h..], &cross);
add_wrapping_assign(&mut out[h..], &cross);
}
}
fn assert_lengths<T>(out: &[T], xs: &[T]) {
assert!(!out.is_empty());
assert!(!xs.is_empty());
}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_square_truncated_to_out_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
pow: u64,
) {
assert_lengths(out, xs);
assert!(pow <= T::WIDTH);
square_truncated_classical_wrapping(out, xs);
mask_coefficients(out, pow);
}}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_square_truncated_to_out_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
pow: u64,
) {
assert_lengths(out, xs);
assert!(pow <= T::WIDTH);
square_truncated_karatsuba_wrapping(out, xs);
mask_coefficients(out, pow);
}}
#[doc(hidden)]
pub fn mod_power_of_2_square_truncated_to_out<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
pow: u64,
) {
assert_lengths(out, xs);
assert!(pow <= T::WIDTH);
square_truncated_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_truncated_helper<T: PrimitiveUnsigned>(
xs: &[T],
len: u64,
pow: u64,
) -> UnsignedPolynomial<T> {
if len == 0 || xs.is_empty() {
return UnsignedPolynomial::ZERO;
}
let mut out = vec![T::ZERO; truncated_len(xs.len(), xs.len(), len)];
mod_power_of_2_square_truncated_to_out(&mut out, xs, pow);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModPowerOf2SquareTruncated for UnsignedPolynomial<T> {
type Output = Self;
fn mod_power_of_2_square_truncated(self, len: u64, pow: u64) -> Self {
assert_reduced(&self, pow);
mod_power_of_2_square_truncated_helper(&self.coefficients, len, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2SquareTruncated for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_power_of_2_square_truncated(self, len: u64, pow: u64) -> UnsignedPolynomial<T> {
assert_reduced(self, pow);
mod_power_of_2_square_truncated_helper(&self.coefficients, len, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2SquareTruncatedAssign for UnsignedPolynomial<T> {
fn mod_power_of_2_square_truncated_assign(&mut self, len: u64, pow: u64) {
assert_reduced(self, pow);
*self = mod_power_of_2_square_truncated_helper(&self.coefficients, len, pow);
}
}