use crate::integer_polynomial::arithmetic::coefficient::PolynomialCoefficient;
use crate::integer_polynomial::arithmetic::vec::{vec_add, vec_sub_assign};
use alloc::vec;
use core::borrow::Borrow;
use core::mem::take;
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::num::logic::traits::LowMask;
pub(crate) const fn revbin(n: usize, bits: u64) -> usize {
debug_assert!(bits != 0);
n.reverse_bits() >> (usize::WIDTH - bits)
}
pub(crate) fn revbin_in<'a, C: PolynomialCoefficient>(out: &mut [&'a C], xs: &'a [C], bits: u64) {
for (i, x) in xs.iter().enumerate() {
out[revbin(i, bits)] = x;
}
}
pub(crate) fn revbin_out<C: PolynomialCoefficient>(out: &mut [C], xs: &mut [C], bits: u64) {
for (i, o) in out.iter_mut().enumerate() {
*o = take(&mut xs[revbin(i, bits)]);
}
}
pub(crate) fn add_shifted_rev<C: PolynomialCoefficient>(xs: &mut [C], ys: &[C], bits: u64) {
for (i, y) in ys[..usize::low_mask(bits)].iter().enumerate() {
xs[revbin(revbin(i, bits) + 1, bits)] += y;
}
}
fn mul_karatsuba_recursive<C: PolynomialCoefficient, T: Borrow<C>>(
out: &mut [C],
xs: &[T],
ys: &[T],
temp: &mut [C],
bits: u64,
) {
let length = usize::power_of_2(bits);
let m = length >> 1;
if length == 1 {
out[0] = xs[0].borrow().mul_ref(ys[0].borrow());
out[1] = C::ZERO;
return;
}
let (sums, temp) = temp.split_at_mut(length);
vec_add(&mut sums[..m], &xs[..m], &xs[m..length]);
vec_add(&mut sums[m..], &ys[..m], &ys[m..length]);
let (out_lo, out_hi) = out.split_at_mut(length);
mul_karatsuba_recursive(out_lo, &xs[..m], &ys[..m], temp, bits - 1);
mul_karatsuba_recursive(out_hi, &sums[..m], &sums[m..], temp, bits - 1);
mul_karatsuba_recursive(sums, &xs[m..], &ys[m..], temp, bits - 1);
vec_sub_assign(&mut out_hi[..length], out_lo);
vec_sub_assign(&mut out_hi[..length], sums);
add_shifted_rev(out_lo, sums, bits);
}
crate_test_fn! {mul_to_out_karatsuba<C: PolynomialCoefficient>(out: &mut [C], xs: &[C], ys: &[C]) {
let len1 = xs.len();
assert!(len1 >= ys.len());
assert_ne!(ys.len(), 0);
if len1 == 1 {
out[0] = xs[0].mul_ref(&ys[0]);
return;
}
let loglen = u64::exact_from(len1).ceiling_log_base_2();
let length = usize::power_of_2(loglen);
let zero = C::ZERO;
let mut revs = vec![&zero; length << 1];
let (rev1, rev2) = revs.split_at_mut(length);
let mut scratch = vec![C::ZERO; length << 2];
let (rev_out, temp) = scratch.split_at_mut(length << 1);
revbin_in(rev1, xs, loglen);
revbin_in(rev2, ys, loglen);
mul_karatsuba_recursive(rev_out, &*rev1, &*rev2, temp, loglen);
revbin_out(out, rev_out, loglen + 1);
}}