use crate::num::arithmetic::traits::Parity;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::unsigned_polynomial::arithmetic::mod_mul::{
ModData, column_sum, mod_add_assign_slice, mod_sub_assign_slice,
};
use alloc::vec;
use core::cmp::Ordering;
pub(crate) const MOD_MUL_MIDDLE_KARATSUBA_THRESHOLD: usize = 48;
pub(crate) fn mod_mul_middle_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<T>,
) {
let n = xs.len();
for (j, o) in out.iter_mut().enumerate() {
*o = d.reduce_sum(column_sum(xs, &ys[j..j + n], d));
}
}
const fn mod_mul_middle_scratch_len(mut n: usize) -> usize {
let mut len = 0;
while n >= MOD_MUL_MIDDLE_KARATSUBA_THRESHOLD {
let h = n >> 1;
len += (h << 2) - 1;
n = h;
}
len
}
fn mod_mul_middle_balanced<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<T>,
scratch: &mut [T],
) {
let n = xs.len();
if n < MOD_MUL_MIDDLE_KARATSUBA_THRESHOLD {
mod_mul_middle_classical(out, xs, ys, d);
return;
}
let m = d.m;
if n.odd() {
let n_1 = n - 1;
let (out_low, out_last) = out.split_at_mut(n_1);
mod_mul_middle_balanced(out_low, &xs[..n_1], &ys[1..n_1 << 1], d, scratch);
let top = xs[n_1];
for (o, &y) in out_low.iter_mut().zip(ys) {
let (hi, lo) = T::x_mul_y_to_zz(top, y);
*o = o.mod_add(d.reduce_2(hi, lo), m);
}
out_last[0] = d.reduce_sum(column_sum(xs, &ys[n_1..], d));
return;
}
let h = n >> 1;
let two_h_1 = (h << 1) - 1;
let (x0, x1) = xs.split_at(h);
split_into_chunks_mut!(scratch, h, [x_sum, p], scratch);
let (y_difference, scratch) = scratch.split_at_mut(two_h_1);
x_sum.copy_from_slice(x0);
mod_add_assign_slice(x_sum, x1, m);
let y1 = &ys[h..h + two_h_1];
mod_mul_middle_balanced(p, x_sum, y1, d, scratch);
let (out_low, out_high) = out.split_at_mut(h);
y_difference.copy_from_slice(&ys[..two_h_1]);
mod_sub_assign_slice(y_difference, y1, m);
mod_mul_middle_balanced(out_low, x1, y_difference, d, scratch);
y_difference.copy_from_slice(&ys[h << 1..]);
mod_sub_assign_slice(y_difference, y1, m);
mod_mul_middle_balanced(out_high, x0, y_difference, d, scratch);
mod_add_assign_slice(out_low, p, m);
mod_add_assign_slice(out_high, p, m);
}
pub(crate) fn mod_mul_middle_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<T>,
) {
let n = xs.len();
let k = out.len();
match n.cmp(&k) {
Ordering::Equal => {
let mut scratch = vec![T::ZERO; mod_mul_middle_scratch_len(n)];
mod_mul_middle_balanced(out, xs, ys, d, &mut scratch);
}
Ordering::Greater => {
out.fill(T::ZERO);
let mut piece_out = vec![T::ZERO; k];
for (c, piece) in xs.chunks(k).enumerate() {
let len = piece.len();
let start = n - c * k - len;
mod_mul_middle_karatsuba(&mut piece_out, piece, &ys[start..start + len + k - 1], d);
mod_add_assign_slice(out, &piece_out, d.m);
}
}
Ordering::Less => {
for (b, block) in out.chunks_mut(n).enumerate() {
let start = b * n;
let len = block.len();
mod_mul_middle_karatsuba(block, xs, &ys[start..start + n + len - 1], d);
}
}
}
}
fn assert_lengths<T>(out: &[T], xs: &[T], ys: &[T]) {
assert!(!out.is_empty());
assert!(!xs.is_empty());
assert_eq!(ys.len(), xs.len() + out.len() - 1);
}
crate_test_fn! {
#[allow(dead_code)]
mod_mul_middle_to_out_classical<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T], m: T) {
assert_lengths(out, xs, ys);
mod_mul_middle_classical(out, xs, ys, &ModData::new(m, xs.len()));
}}
crate_test_fn! {
#[allow(dead_code)]
mod_mul_middle_to_out_karatsuba<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T], m: T) {
assert_lengths(out, xs, ys);
mod_mul_middle_karatsuba(out, xs, ys, &ModData::new(m, xs.len()));
}}