use crate::integer_polynomial::arithmetic::square::square_ref;
use crate::natural::Natural;
use crate::natural::arithmetic::add::limbs_slice_add_same_length_in_place_left;
use crate::natural::arithmetic::mod_power_of_2_square::limbs_square_low;
use crate::natural::arithmetic::mul::mul_low::limbs_mul_low_same_length;
use crate::natural::arithmetic::shl::limbs_slice_shl_in_place;
use crate::natural_polynomial::NaturalPolynomial;
use crate::natural_polynomial::arithmetic::mod_power_of_2_mul::{
low_preferred, naturals_to_slots, reduce_coefficients, slot_len, slots_add_assign,
slots_sub_assign, slots_to_naturals, square_slots_karatsuba_threshold,
};
use crate::platform::Limb;
use alloc::vec;
use alloc::vec::Vec;
use malachite_base::num::arithmetic::traits::{
ModPowerOf2Assign, ModPowerOf2IsReduced, ModPowerOf2Square, ModPowerOf2SquareAssign,
};
use malachite_base::num::basic::integers::PrimitiveInt;
use malachite_base::num::basic::traits::Zero;
use malachite_base::num::conversion::traits::ExactFrom;
use malachite_base::slices::slice_test_zero;
use malachite_base::unsigned_polynomial::arithmetic::mod_power_of_2_square::*;
fn slots_double_assign(xs: &mut [Limb], slot_len: usize) {
for x in xs.chunks_exact_mut(slot_len) {
limbs_slice_shl_in_place(x, 1);
}
}
pub(crate) fn square_slots_classical(out: &mut [Limb], xs: &[Limb], slot_len: usize) {
out.fill(0);
let mut product = vec![0; slot_len];
for (i, x) in xs.chunks_exact(slot_len).enumerate() {
if slice_test_zero(x) {
continue;
}
for (o, y) in out[((i << 1) + 1) * slot_len..]
.chunks_exact_mut(slot_len)
.zip(xs[(i + 1) * slot_len..].chunks_exact(slot_len))
{
limbs_mul_low_same_length(&mut product, x, y);
limbs_slice_add_same_length_in_place_left(o, &product);
}
}
slots_double_assign(out, slot_len);
for (o, x) in out
.chunks_exact_mut(slot_len)
.step_by(2)
.zip(xs.chunks_exact(slot_len))
{
limbs_square_low(&mut product, x);
limbs_slice_add_same_length_in_place_left(o, &product);
}
}
pub(crate) fn square_slots_karatsuba(out: &mut [Limb], xs: &[Limb], slot_len: usize) {
let n = xs.len() / slot_len;
if n < square_slots_karatsuba_threshold(slot_len) {
square_slots_classical(out, xs, slot_len);
return;
}
let h = n >> 1;
let split = h * slot_len;
let low_len = ((h << 1) - 1) * slot_len;
let (x0, x1) = xs.split_at(split);
let (low, high) = out.split_at_mut(split << 1);
square_slots_karatsuba(&mut low[..low_len], x0, slot_len);
low[low_len..].fill(0);
square_slots_karatsuba(high, x1, slot_len);
let mut sum = x1.to_vec();
slots_add_assign(&mut sum, x0, slot_len);
let mut middle = vec![0; (((n - h) << 1) - 1) * slot_len];
square_slots_karatsuba(&mut middle, &sum, slot_len);
slots_sub_assign(&mut middle, &out[..low_len], slot_len);
slots_sub_assign(&mut middle, &out[split << 1..], slot_len);
slots_add_assign(&mut out[split..], &middle, slot_len);
}
crate_test_fn! {mod_power_of_2_square_low_classical(xs: &[Natural], pow: u64) -> Vec<Natural> {
let len = (xs.len() << 1) - 1;
if pow == 0 {
return vec![Natural::ZERO; len];
}
let slot_len = slot_len(pow);
let mut out = vec![0; len * slot_len];
square_slots_classical(&mut out, &naturals_to_slots(xs, slot_len), slot_len);
slots_to_naturals(&out, slot_len, pow)
}}
crate_test_fn! {mod_power_of_2_square_low_karatsuba(xs: &[Natural], pow: u64) -> Vec<Natural> {
let len = (xs.len() << 1) - 1;
if pow == 0 {
return vec![Natural::ZERO; len];
}
if pow <= Limb::WIDTH {
let xs: Vec<Limb> = xs.iter().map(Limb::exact_from).collect();
let mut out = vec![0; len];
mod_power_of_2_square_to_out(&mut out, &xs, pow);
return out.into_iter().map(Natural::from).collect();
}
let slot_len = slot_len(pow);
let mut out = vec![0; len * slot_len];
square_slots_karatsuba(&mut out, &naturals_to_slots(xs, slot_len), slot_len);
slots_to_naturals(&out, slot_len, pow)
}}
pub(crate) fn assert_reduced(p: &NaturalPolynomial, pow: u64) {
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) const SQUARE_LOW_WINDOWS: [(u64, usize); 11] = [
(16, 150),
(32, 1000),
(64, 2000),
(96, 48),
(200, 64),
(300, 100),
(2000, 32),
(5000, 16),
(10000, 8),
(20000, 4),
(u64::MAX, 0),
];
pub(crate) fn mod_power_of_2_square_ref(xs: &[Natural], pow: u64) -> NaturalPolynomial {
if xs.len() > 1 && low_preferred(&SQUARE_LOW_WINDOWS, xs.len(), pow) {
reduce_coefficients(mod_power_of_2_square_low_karatsuba(xs, pow), pow)
} else {
reduce_coefficients(square_ref(xs), pow)
}
}
crate_test_fn! {mod_power_of_2_square_full(xs: &[Natural], pow: u64) -> Vec<Natural> {
let mut out = square_ref(xs);
for x in &mut out {
x.mod_power_of_2_assign(pow);
}
out
}}
impl ModPowerOf2Square for NaturalPolynomial {
type Output = Self;
#[inline]
fn mod_power_of_2_square(mut self, pow: u64) -> Self {
self.mod_power_of_2_square_assign(pow);
self
}
}
impl ModPowerOf2Square for &NaturalPolynomial {
type Output = NaturalPolynomial;
fn mod_power_of_2_square(self, pow: u64) -> NaturalPolynomial {
assert_reduced(self, pow);
mod_power_of_2_square_ref(&self.coefficients, pow)
}
}
impl ModPowerOf2SquareAssign for NaturalPolynomial {
fn mod_power_of_2_square_assign(&mut self, pow: u64) {
assert_reduced(self, pow);
if let [c] = self.coefficients.as_mut_slice() {
c.mod_power_of_2_square_assign(pow);
self.trim();
} else {
*self = mod_power_of_2_square_ref(&self.coefficients, pow);
}
}
}