use crate::num::arithmetic::traits::{ModPowerOf2Mul, ModPowerOf2MulAssign};
use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::unsigned_polynomial::UnsignedPolynomial;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_add::assert_reduced;
use alloc::vec;
use alloc::vec::Vec;
pub(crate) const MOD_POWER_OF_2_MUL_KARATSUBA_THRESHOLD: usize = 32;
pub(crate) const MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD: usize = 64;
pub(crate) fn mask_coefficients<T: PrimitiveUnsigned>(xs: &mut [T], pow: u64) {
if pow < T::WIDTH {
let mask = T::low_mask(pow);
for x in xs {
*x &= mask;
}
}
}
pub(crate) fn add_wrapping_assign<T: PrimitiveUnsigned>(xs: &mut [T], ys: &[T]) {
for (x, &y) in xs.iter_mut().zip(ys) {
x.wrapping_add_assign(y);
}
}
pub(crate) fn sub_wrapping_assign<T: PrimitiveUnsigned>(xs: &mut [T], ys: &[T]) {
for (x, &y) in xs.iter_mut().zip(ys) {
x.wrapping_sub_assign(y);
}
}
pub(crate) fn mul_classical_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T]) {
out.fill(T::ZERO);
for (i, &x) in xs.iter().enumerate() {
if x != T::ZERO {
for (o, &y) in out[i..].iter_mut().zip(ys) {
o.wrapping_add_assign(x.wrapping_mul(y));
}
}
}
}
pub(crate) const fn karatsuba_wrapping_scratch_len(mut n: usize) -> usize {
let mut len = 0;
while n >= MOD_POWER_OF_2_MUL_KARATSUBA_THRESHOLD {
let c = n - (n >> 1);
len += (c << 2) - 1;
n = c;
}
len
}
fn mul_karatsuba_balanced_wrapping<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
scratch: &mut [T],
) {
let n = xs.len();
if n < MOD_POWER_OF_2_MUL_KARATSUBA_THRESHOLD {
mul_classical_wrapping(out, xs, ys);
return;
}
let h = n >> 1;
let c = n - h;
let two_h = h << 1;
let (x0, x1) = xs.split_at(h);
let (y0, y1) = ys.split_at(h);
split_into_chunks_mut!(scratch, c, [x_sum, y_sum], scratch);
let (middle, scratch) = scratch.split_at_mut((c << 1) - 1);
let (low, high) = out.split_at_mut(two_h);
mul_karatsuba_balanced_wrapping(&mut low[..two_h - 1], x0, y0, scratch);
low[two_h - 1] = T::ZERO;
mul_karatsuba_balanced_wrapping(high, x1, y1, scratch);
x_sum.copy_from_slice(x1);
add_wrapping_assign(x_sum, x0);
y_sum.copy_from_slice(y1);
add_wrapping_assign(y_sum, y0);
mul_karatsuba_balanced_wrapping(middle, x_sum, y_sum, scratch);
sub_wrapping_assign(middle, &out[..two_h - 1]);
sub_wrapping_assign(middle, &out[two_h..]);
add_wrapping_assign(&mut out[h..], middle);
}
pub(crate) fn mul_karatsuba_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T]) {
let (xs, ys) = if xs.len() >= ys.len() {
(xs, ys)
} else {
(ys, xs)
};
let n = xs.len();
let m = ys.len();
if m < MOD_POWER_OF_2_MUL_KARATSUBA_THRESHOLD {
mul_classical_wrapping(out, xs, ys);
return;
}
let mut scratch = vec![T::ZERO; karatsuba_wrapping_scratch_len(m)];
if n == m {
mul_karatsuba_balanced_wrapping(out, xs, ys, &mut scratch);
return;
}
out.fill(T::ZERO);
let mut product = vec![T::ZERO; (m << 1) - 1];
for (k, piece) in xs.chunks(m).enumerate() {
let product = &mut product[..piece.len() + m - 1];
if piece.len() == m {
mul_karatsuba_balanced_wrapping(product, piece, ys, &mut scratch);
} else {
mul_karatsuba_wrapping(product, ys, piece);
}
add_wrapping_assign(&mut out[k * m..], product);
}
}
fn assert_lengths<T>(out: &[T], xs: &[T], ys: &[T]) {
assert!(!xs.is_empty());
assert!(!ys.is_empty());
assert_eq!(out.len(), xs.len() + ys.len() - 1);
}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_mul_to_out_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
pow: u64,
) {
assert_lengths(out, xs, ys);
assert!(pow <= T::WIDTH);
mul_classical_wrapping(out, xs, ys);
mask_coefficients(out, pow);
}}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_mul_to_out_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
pow: u64,
) {
assert_lengths(out, xs, ys);
assert!(pow <= T::WIDTH);
mul_karatsuba_wrapping(out, xs, ys);
mask_coefficients(out, pow);
}}
#[doc(hidden)]
pub fn mod_power_of_2_mul_to_out<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
pow: u64,
) {
assert_lengths(out, xs, ys);
assert!(pow <= T::WIDTH);
mul_karatsuba_wrapping(out, xs, ys);
mask_coefficients(out, pow);
}
pub(crate) fn from_coefficients_trimmed<T: PrimitiveUnsigned>(xs: Vec<T>) -> UnsignedPolynomial<T> {
let mut p = UnsignedPolynomial { coefficients: xs };
p.trim();
p
}
pub(crate) fn mod_power_of_2_mul_helper<T: PrimitiveUnsigned>(
xs: &[T],
ys: &[T],
pow: u64,
) -> UnsignedPolynomial<T> {
if xs.is_empty() || ys.is_empty() {
return UnsignedPolynomial::ZERO;
}
let mut out = vec![T::ZERO; xs.len() + ys.len() - 1];
mod_power_of_2_mul_to_out(&mut out, xs, ys, pow);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModPowerOf2Mul<Self> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_power_of_2_mul(self, other: Self, pow: u64) -> Self {
assert_reduced(&self, &other, pow);
mod_power_of_2_mul_helper(&self.coefficients, &other.coefficients, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2Mul<&Self> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_power_of_2_mul(self, other: &Self, pow: u64) -> Self {
assert_reduced(&self, other, pow);
mod_power_of_2_mul_helper(&self.coefficients, &other.coefficients, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2Mul<UnsignedPolynomial<T>> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_power_of_2_mul(self, other: UnsignedPolynomial<T>, pow: u64) -> UnsignedPolynomial<T> {
assert_reduced(self, &other, pow);
mod_power_of_2_mul_helper(&self.coefficients, &other.coefficients, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2Mul<&UnsignedPolynomial<T>> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_power_of_2_mul(self, other: &UnsignedPolynomial<T>, pow: u64) -> UnsignedPolynomial<T> {
assert_reduced(self, other, pow);
mod_power_of_2_mul_helper(&self.coefficients, &other.coefficients, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulAssign<Self> for UnsignedPolynomial<T> {
fn mod_power_of_2_mul_assign(&mut self, other: Self, pow: u64) {
assert_reduced(self, &other, pow);
*self = mod_power_of_2_mul_helper(&self.coefficients, &other.coefficients, pow);
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulAssign<&Self> for UnsignedPolynomial<T> {
fn mod_power_of_2_mul_assign(&mut self, other: &Self, pow: u64) {
assert_reduced(self, other, pow);
*self = mod_power_of_2_mul_helper(&self.coefficients, &other.coefficients, pow);
}
}