use crate::integer_polynomial::arithmetic::coefficient::PolynomialCoefficient;
use crate::integer_polynomial::arithmetic::mul_truncated::karatsuba::{
combine_truncated_karatsuba, padded,
};
use crate::integer_polynomial::arithmetic::square::karatsuba::square_to_out_karatsuba;
use crate::integer_polynomial::arithmetic::square_truncated::classical::*;
use crate::integer_polynomial::arithmetic::vec::vec_add;
use alloc::vec;
use malachite_base::num::arithmetic::traits::{CeilingLogBase2, Parity, PowerOf2};
use malachite_base::num::conversion::traits::ExactFrom;
use malachite_base::split_into_chunks_mut;
fn square_truncated_karatsuba_recursive<C: PolynomialCoefficient>(
out: &mut [C],
xs: &[C],
temp: &mut [C],
len: usize,
) {
let m1 = len >> 1;
let m2 = len - m1;
let odd = len.odd();
if len <= 6 {
square_truncated_to_out_classical(&mut out[..len], &xs[..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();
}
split_into_chunks_mut!(temp, m2, [low, sums], rest);
square_truncated_karatsuba_recursive(low, sums, rest, m2);
let (high, rest) = temp[m2..].split_at_mut(m2);
square_truncated_karatsuba_recursive(high, &xs[m1..], rest, m2);
square_to_out_karatsuba(&mut out[..two_m1 - 1], &xs[..m1]);
out[two_m1 - 1] = C::ZERO;
combine_truncated_karatsuba(out, temp, m1, m2);
}
crate_test_fn! {square_truncated_to_out_karatsuba_n<C: PolynomialCoefficient>(
out: &mut [C],
xs: &[C],
) {
let n = out.len();
assert_ne!(n, 0);
assert!(xs.len() >= n);
if n == 1 {
out[0] = xs[0].square_ref();
return;
}
let len = usize::power_of_2(u64::exact_from(n).ceiling_log_base_2());
let mut temp = vec![C::ZERO; (len << 1) + 2];
square_truncated_karatsuba_recursive(out, xs, &mut temp, n);
}}
crate_test_fn! {square_truncated_to_out_karatsuba<C: PolynomialCoefficient>(
out: &mut [C],
xs: &[C],
) {
let n = out.len();
square_truncated_to_out_karatsuba_n(out, &padded(xs, n));
}}