use crate::integer_polynomial::arithmetic::square_truncated::square_truncated_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_to_naturals, square_slots_karatsuba_threshold,
};
use crate::natural_polynomial::arithmetic::mod_power_of_2_mul_truncated::*;
use crate::natural_polynomial::arithmetic::mod_power_of_2_square::{
assert_reduced, square_slots_karatsuba,
};
use crate::platform::Limb;
use alloc::vec;
use alloc::vec::Vec;
use core::cmp::min;
use malachite_base::num::arithmetic::traits::{ModPowerOf2Assign, 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::polynomial::{ModPowerOf2SquareTruncated, ModPowerOf2SquareTruncatedAssign};
use malachite_base::slices::slice_test_zero;
use malachite_base::unsigned_polynomial::arithmetic::mod_power_of_2_square_truncated::*;
pub(crate) fn square_truncated_slots_classical(out: &mut [Limb], xs: &[Limb], slot_len: usize) {
let len = out.len() / slot_len;
out.fill(0);
let mut product = vec![0; slot_len];
for (i, x) in xs.chunks_exact(slot_len).enumerate() {
if (i << 1) + 1 >= len {
break;
}
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);
}
}
for o in out.chunks_exact_mut(slot_len) {
limbs_slice_shl_in_place(o, 1);
}
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_truncated_slots_karatsuba(out: &mut [Limb], xs: &[Limb], slot_len: usize) {
let len = out.len() / slot_len;
let xs = &xs[..min(xs.len(), len * slot_len)];
let n = xs.len() / slot_len;
let full_len = (n << 1) - 1;
if full_len <= len {
square_slots_karatsuba(&mut out[..full_len * slot_len], xs, slot_len);
out[full_len * slot_len..].fill(0);
return;
}
if n < square_slots_karatsuba_threshold(slot_len) {
square_truncated_slots_classical(out, xs, slot_len);
return;
}
let h = len.div_ceil(2);
let split = h * slot_len;
let x0 = &xs[..min(split, xs.len())];
square_truncated_slots_karatsuba(out, x0, slot_len);
if n > h {
let mut cross = vec![0; (len - h) * slot_len];
mul_truncated_slots_karatsuba(&mut cross, x0, &xs[split..], slot_len);
slots_add_assign(&mut out[split..], &cross, slot_len);
slots_add_assign(&mut out[split..], &cross, slot_len);
}
}
crate_test_fn! {mod_power_of_2_square_truncated_low_classical(
xs: &[Natural],
len: usize,
pow: u64,
) -> Vec<Natural> {
if pow == 0 {
return vec![Natural::ZERO; len];
}
let slot_len = slot_len(pow);
let mut out = vec![0; len * slot_len];
square_truncated_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_truncated_low_karatsuba(
xs: &[Natural],
len: usize,
pow: u64,
) -> Vec<Natural> {
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_truncated_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_truncated_slots_karatsuba(&mut out, &naturals_to_slots(xs, slot_len), slot_len);
slots_to_naturals(&out, slot_len, pow)
}}
pub(crate) const SQUARE_TRUNCATED_LOW_WINDOWS: [(u64, usize); 13] = [
(16, 200),
(32, 1000),
(64, 2000),
(96, 64),
(128, 100),
(200, 150),
(300, 200),
(500, 100),
(2000, 64),
(5000, 16),
(10000, 8),
(20000, 4),
(u64::MAX, 0),
];
pub(crate) fn mod_power_of_2_square_truncated_ref(
xs: &[Natural],
len: u64,
pow: u64,
) -> NaturalPolynomial {
let n = xs.len();
let len_usize = usize::try_from(len).unwrap_or(usize::MAX);
if len != 0 && n > 1 && low_preferred(&SQUARE_TRUNCATED_LOW_WINDOWS, min(n, len_usize), pow) {
let len = min(len_usize, (n << 1) - 1);
reduce_coefficients(
mod_power_of_2_square_truncated_low_karatsuba(xs, len, pow),
pow,
)
} else {
reduce_coefficients(square_truncated_ref(xs, len), pow)
}
}
crate_test_fn! {mod_power_of_2_square_truncated_full(
xs: &[Natural],
len: usize,
pow: u64,
) -> Vec<Natural> {
let mut out = square_truncated_ref(xs, u64::exact_from(len));
for x in &mut out {
x.mod_power_of_2_assign(pow);
}
out
}}
impl ModPowerOf2SquareTruncated for NaturalPolynomial {
type Output = Self;
#[inline]
fn mod_power_of_2_square_truncated(mut self, len: u64, pow: u64) -> Self {
self.mod_power_of_2_square_truncated_assign(len, pow);
self
}
}
impl ModPowerOf2SquareTruncated for &NaturalPolynomial {
type Output = NaturalPolynomial;
fn mod_power_of_2_square_truncated(self, len: u64, pow: u64) -> NaturalPolynomial {
assert_reduced(self, pow);
mod_power_of_2_square_truncated_ref(&self.coefficients, len, pow)
}
}
impl ModPowerOf2SquareTruncatedAssign for NaturalPolynomial {
fn mod_power_of_2_square_truncated_assign(&mut self, len: u64, pow: u64) {
assert_reduced(self, pow);
if len != 0
&& let [c] = self.coefficients.as_mut_slice()
{
c.mod_power_of_2_square_assign(pow);
self.trim();
} else {
*self = mod_power_of_2_square_truncated_ref(&self.coefficients, len, pow);
}
}
}