use crate::integer_polynomial::arithmetic::coefficient::PolynomialCoefficient;
use crate::integer_polynomial::arithmetic::mul_middle::classical::mul_middle_to_out_classical;
use crate::integer_polynomial::arithmetic::mul_middle::truncate_mul_middle_inputs;
use crate::integer_polynomial::arithmetic::vec::max_bits::vec_max_bits;
use crate::natural::arithmetic::add::limbs_slice_add_limb_in_place;
use crate::natural::arithmetic::mul::schonhage_strassen::convolution::fft_convolution;
use crate::natural::arithmetic::mul::schonhage_strassen::limbs_neg_to_out;
use crate::natural::arithmetic::mul::schonhage_strassen::mulmod_2expp1::*;
use crate::natural::arithmetic::mul::schonhage_strassen::mulmod_2expp1_basecase::*;
use crate::natural::arithmetic::neg::limbs_neg_in_place;
use crate::natural::{LIMB_HIGH_BIT, Natural};
use crate::platform::Limb;
use alloc::vec;
use alloc::vec::Vec;
use core::cmp::min;
use core::mem::swap;
use core::ptr;
use malachite_base::num::arithmetic::traits::{CeilingLogBase2, PowerOf2};
use malachite_base::num::basic::integers::PrimitiveInt;
use malachite_base::num::conversion::traits::ExactFrom;
use malachite_base::slices::slice_test_zero;
crate_test_fn! {integers_to_fermat_residues<C: PolynomialCoefficient>(
coeffs_f: &mut [Vec<Limb>],
xs: &[C],
limbs: usize,
) {
let size_f = limbs + 1;
for (f, x) in coeffs_f.iter_mut().zip(xs.iter()) {
let coeff = x.unsigned_abs_ref().as_limbs_asc();
let size_j = coeff.len();
let f = &mut f[..size_f];
if x.is_negative() {
limbs_neg_to_out(f, coeff);
f[size_j..].fill(Limb::MAX);
} else {
f[..size_j].copy_from_slice(coeff);
f[size_j..].fill(0);
}
}
}}
crate_test_fn! {integers_from_fermat_residues<C: PolynomialCoefficient>(
out: &mut [C],
coeffs_f: &[Vec<Limb>],
limbs: usize,
sign: bool,
) {
for (x, f) in out.iter_mut().zip(coeffs_f.iter()) {
let top = f[limbs - 1];
*x = if sign
&& (f[limbs] != 0
|| top > LIMB_HIGH_BIT
|| (top == LIMB_HIGH_BIT && !slice_test_zero(&f[..limbs - 1])))
{
let mut data = f[..limbs].to_vec();
limbs_neg_in_place(&mut data);
limbs_slice_add_limb_in_place(&mut data, 1);
C::from_sign_and_abs(false, Natural::from_owned_limbs_asc(data))
} else {
C::from_sign_and_abs(true, Natural::from_limbs_asc(&f[..limbs]))
};
}
}}
crate_test_fn! {mul_middle_to_out_schonhage_strassen<C: PolynomialCoefficient>(
out: &mut [C],
xs: &[C],
ys: &[C],
nlo: usize,
nhi: usize,
) {
assert_ne!(xs.len(), 0);
assert_ne!(ys.len(), 0);
assert!(nlo < nhi);
assert!(nhi < xs.len() + ys.len());
let (mut xs, mut ys, nlo, nhi) = truncate_mul_middle_inputs(xs, ys, nlo, nhi);
if nhi <= 2 {
mul_middle_to_out_classical(out, xs, ys, nlo, nhi);
return;
}
if xs.len() < ys.len() {
swap(&mut xs, &mut ys);
}
let trunc = nhi;
let xs = &xs[..min(xs.len(), trunc)];
let ys = &ys[..min(ys.len(), trunc)];
let square = ptr::eq(xs, ys);
let len1 = xs.len();
let len2 = ys.len();
let len_out = len1 + len2 - 1;
let loglen = u64::exact_from(len_out).ceiling_log_base_2();
let loglen2 = u64::exact_from(len2).ceiling_log_base_2();
let n = usize::power_of_2(loglen - 2);
let (bits1, negative1) = vec_max_bits(xs);
let (bits2, negative2) = if square {
(bits1, negative1)
} else {
vec_max_bits(ys)
};
let size1 = bits1.div_ceil(Limb::WIDTH);
let size2 = bits2.div_ceil(Limb::WIDTH);
let res_bits = ((size1 + size2) << Limb::LOG_WIDTH) + loglen2 + 1;
let res_bits = (((res_bits - 1) >> (loglen - 2)) + 1) << (loglen - 2);
let mut limbs = usize::exact_from((res_bits - 1) >> Limb::LOG_WIDTH) + 1;
if limbs > FFT_MULMOD_2EXPP1_CUTOFF {
limbs = usize::power_of_2(u64::exact_from(limbs).ceiling_log_base_2());
}
let size = limbs + 1;
let mut residues: Vec<Vec<Limb>> = vec![vec![0; size]; (n << 2) + 2];
let (ii, t) = residues.split_at_mut(n << 2);
let (t1, t2) = t.split_at_mut(1);
let (t1, t2) = (&mut t1[0], &mut t2[0]);
integers_to_fermat_residues(ii, xs, limbs);
let mut jj = if square {
None
} else {
let mut jj: Vec<Vec<Limb>> = vec![vec![0; size]; n << 2];
integers_to_fermat_residues(&mut jj, ys, limbs);
Some(jj)
};
let sign = negative1 || negative2;
let res_bits = bits1 + bits2 + loglen2 + u64::from(sign);
if res_bits == 0 {
out[..trunc - nlo].fill(C::ZERO);
return;
}
let res_bits = (((res_bits - 1) >> (loglen - 2)) + 1) << (loglen - 2);
let limbs = usize::exact_from((res_bits - 1) >> Limb::LOG_WIDTH) + 1;
let limbs = fft_adjust_limbs(limbs); let mut scratch = vec![0; limbs + 1 + limbs_mul_mod_2expp1_basecase_scratch_len(limbs)];
let (s1, tt) = scratch.split_at_mut(limbs + 1);
fft_convolution(ii, jj.as_deref_mut(), loglen - 2, limbs, len_out, t1, t2, s1, tt);
integers_from_fermat_residues(&mut out[..trunc - nlo], &ii[nlo..], limbs, sign); }}