use crate::num::arithmetic::traits::{ModIsReduced, ModSquare, ModSquareAssign, Parity};
use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::unsigned_polynomial::UnsignedPolynomial;
use crate::unsigned_polynomial::arithmetic::mod_mul::{
MOD_SQUARE_KARATSUBA_THRESHOLD, ModData, accumulate, column_sum, mod_add_assign_slice,
mod_karatsuba_scratch_len, mod_sub_assign_slice,
};
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::from_coefficients_trimmed;
use alloc::vec;
#[inline]
pub(crate) fn double<T: PrimitiveUnsigned>(acc: &mut (T, T, T)) {
let (a2, a1, a0) = *acc;
let top = T::WIDTH - 1;
*acc = ((a2 << 1) | (a1 >> top), (a1 << 1) | (a0 >> top), a0 << 1);
}
pub(crate) fn mod_square_classical_prefix<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
d: &ModData<T>,
) {
let n = xs.len();
for (k, o) in out.iter_mut().enumerate() {
let start = k.saturating_sub(n - 1);
let stop = (k + 1) >> 1;
let mut acc = if start < stop {
column_sum(&xs[start..stop], &xs[k + 1 - stop..=k - start], d)
} else {
(T::ZERO, T::ZERO, T::ZERO)
};
double(&mut acc);
if k.even() && (k >> 1) < n {
let x = xs[k >> 1];
accumulate(&mut acc, x, x);
}
*o = d.reduce_sum(acc);
}
}
fn mod_square_karatsuba_scratch<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
d: &ModData<T>,
scratch: &mut [T],
) {
let n = xs.len();
if n < MOD_SQUARE_KARATSUBA_THRESHOLD {
mod_square_classical_prefix(out, xs, 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 (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);
mod_square_karatsuba_scratch(&mut low[..two_h - 1], x0, d, scratch);
low[two_h - 1] = T::ZERO;
mod_square_karatsuba_scratch(high, x1, d, scratch);
sum.copy_from_slice(x1);
mod_add_assign_slice(sum, x0, m);
mod_square_karatsuba_scratch(middle, 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_square_karatsuba<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], d: &ModData<T>) {
if xs.len() < MOD_SQUARE_KARATSUBA_THRESHOLD {
mod_square_classical_prefix(out, xs, d);
return;
}
let mut scratch =
vec![T::ZERO; mod_karatsuba_scratch_len(xs.len(), MOD_SQUARE_KARATSUBA_THRESHOLD)];
mod_square_karatsuba_scratch(out, xs, d, &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_square_to_out_classical<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
assert_lengths(out, xs);
mod_square_classical_prefix(out, xs, &ModData::new(m, xs.len()));
}}
crate_test_fn! {
#[allow(dead_code)]
mod_square_to_out_karatsuba<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
assert_lengths(out, xs);
mod_square_karatsuba(out, xs, &ModData::new(m, xs.len()));
}}
#[doc(hidden)]
pub fn mod_square_to_out<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
assert_lengths(out, xs);
mod_square_karatsuba(out, xs, &ModData::new(m, xs.len()));
}
fn assert_reduced<T: PrimitiveUnsigned>(p: &UnsignedPolynomial<T>, m: T) {
assert!(
p.mod_is_reduced(&m),
"self must be reduced mod m, but {p} has a coefficient >= {m}"
);
}
pub(crate) fn mod_square_helper<T: PrimitiveUnsigned>(xs: &[T], m: T) -> UnsignedPolynomial<T> {
if xs.is_empty() {
return UnsignedPolynomial::ZERO;
}
let mut out = vec![T::ZERO; (xs.len() << 1) - 1];
mod_square_to_out(&mut out, xs, m);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModSquare<T> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_square(self, m: T) -> Self {
assert_reduced(&self, m);
mod_square_helper(&self.coefficients, m)
}
}
impl<T: PrimitiveUnsigned> ModSquare<T> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_square(self, m: T) -> UnsignedPolynomial<T> {
assert_reduced(self, m);
mod_square_helper(&self.coefficients, m)
}
}
impl<T: PrimitiveUnsigned> ModSquareAssign<T> for UnsignedPolynomial<T> {
fn mod_square_assign(&mut self, m: T) {
assert_reduced(self, m);
*self = mod_square_helper(&self.coefficients, m);
}
}