use crate::integer_polynomial::arithmetic::coefficient::PolynomialCoefficient;
use crate::integer_polynomial::arithmetic::mul::karatsuba::mul_to_out_karatsuba;
use crate::integer_polynomial::arithmetic::mul_truncated::classical::mul_truncated_to_out_classical;
use crate::integer_polynomial::arithmetic::vec::{vec_add, vec_add_assign, vec_sub_assign};
use alloc::borrow::Cow;
use alloc::vec;
use alloc::vec::Vec;
use core::mem::take;
use malachite_base::num::arithmetic::traits::{CeilingLogBase2, Parity, PowerOf2};
use malachite_base::num::conversion::traits::ExactFrom;
use malachite_base::split_into_chunks_mut;
fn mul_truncated_karatsuba_recursive<C: PolynomialCoefficient>(
out: &mut [C],
xs: &[C],
ys: &[C],
temp: &mut [C],
len: usize,
) {
let m1 = len >> 1;
let m2 = len - m1;
let odd = len.odd();
if len <= 6 {
mul_truncated_to_out_classical(&mut out[..len], &xs[..len], &ys[..len]);
return;
}
let two_m1 = m1 << 1;
vec_add(&mut temp[m2..m2 + m1], &xs[..m1], &xs[m1..two_m1]);
if odd {
temp[m2 + m1] = xs[two_m1].clone();
}
vec_add(
&mut temp[m2 << 1..(m2 << 1) + m1],
&ys[..m1],
&ys[m1..two_m1],
);
if odd {
temp[(m2 << 1) + m1] = ys[two_m1].clone();
}
mul_to_out_karatsuba(&mut out[..two_m1 - 1], &xs[..m1], &ys[..m1]);
out[two_m1 - 1] = C::ZERO;
split_into_chunks_mut!(temp, m2, [low, sums_1, sums_2], rest);
mul_truncated_karatsuba_recursive(low, sums_1, sums_2, rest, m2);
let (high, rest) = temp[m2..].split_at_mut(m2);
mul_truncated_karatsuba_recursive(high, &xs[m1..], &ys[m1..], rest, m2);
combine_truncated_karatsuba(out, temp, m1, m2);
}
pub(crate) fn combine_truncated_karatsuba<C: PolynomialCoefficient>(
out: &mut [C],
temp: &mut [C],
m1: usize,
m2: usize,
) {
vec_sub_assign(&mut temp[..m2], &out[..m2]);
let (low, high) = temp.split_at_mut(m2);
vec_sub_assign(low, &high[..m2]);
if m2 != m1 {
out[m1 << 1] = take(&mut high[0]);
}
vec_add_assign(&mut out[m1..m1 + m2], low);
}
crate_test_fn! {mul_truncated_to_out_karatsuba_n<C: PolynomialCoefficient>(
out: &mut [C],
xs: &[C],
ys: &[C],
) {
let n = out.len();
assert_ne!(n, 0);
assert!(xs.len() >= n);
assert!(ys.len() >= n);
if n == 1 {
out[0] = xs[0].mul_ref(&ys[0]);
return;
}
let len = usize::power_of_2(u64::exact_from(n).ceiling_log_base_2());
let mut temp = vec![C::ZERO; 3 * len];
mul_truncated_karatsuba_recursive(out, xs, ys, &mut temp, n);
}}
pub(crate) fn padded<C: PolynomialCoefficient>(xs: &[C], n: usize) -> Cow<'_, [C]> {
if xs.len() >= n {
Cow::Borrowed(&xs[..n])
} else {
let mut padded = Vec::with_capacity(n);
padded.extend_from_slice(xs);
padded.resize(n, C::ZERO);
Cow::Owned(padded)
}
}
crate_test_fn! {mul_truncated_to_out_karatsuba<C: PolynomialCoefficient>(
out: &mut [C],
xs: &[C],
ys: &[C],
) {
let n = out.len();
mul_truncated_to_out_karatsuba_n(out, &padded(xs, n), &padded(ys, n));
}}