use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::polynomial::{ModPowerOf2MulTruncated, ModPowerOf2MulTruncatedAssign};
use crate::unsigned_polynomial::UnsignedPolynomial;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_add::assert_reduced;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::{
MOD_POWER_OF_2_MUL_KARATSUBA_THRESHOLD, add_wrapping_assign, from_coefficients_trimmed,
mask_coefficients, mul_karatsuba_wrapping,
};
use alloc::vec;
use core::cmp::min;
pub(crate) fn mul_truncated_classical_wrapping<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
) {
let len = out.len();
out.fill(T::ZERO);
for (i, &x) in xs.iter().take(len).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) fn mul_truncated_karatsuba_wrapping<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
) {
let len = out.len();
let xs = &xs[..min(xs.len(), len)];
let ys = &ys[..min(ys.len(), len)];
let (xs, ys) = if xs.len() >= ys.len() {
(xs, ys)
} else {
(ys, xs)
};
let full_len = xs.len() + ys.len() - 1;
if full_len <= len {
mul_karatsuba_wrapping(&mut out[..full_len], xs, ys);
out[full_len..].fill(T::ZERO);
return;
}
if ys.len() < MOD_POWER_OF_2_MUL_KARATSUBA_THRESHOLD {
mul_truncated_classical_wrapping(out, xs, ys);
return;
}
let h = len.div_ceil(2);
let x0 = &xs[..min(h, xs.len())];
let y0 = &ys[..min(h, ys.len())];
mul_truncated_karatsuba_wrapping(out, x0, y0);
let mut cross = vec![T::ZERO; len - h];
if xs.len() > h {
mul_truncated_karatsuba_wrapping(&mut cross, &xs[h..], y0);
add_wrapping_assign(&mut out[h..], &cross);
}
if ys.len() > h {
mul_truncated_karatsuba_wrapping(&mut cross, x0, &ys[h..]);
add_wrapping_assign(&mut out[h..], &cross);
}
}
fn assert_lengths<T>(out: &[T], xs: &[T], ys: &[T]) {
assert!(!out.is_empty());
assert!(!xs.is_empty());
assert!(!ys.is_empty());
}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_mul_truncated_to_out_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
pow: u64,
) {
assert_lengths(out, xs, ys);
assert!(pow <= T::WIDTH);
mul_truncated_classical_wrapping(out, xs, ys);
mask_coefficients(out, pow);
}}
crate_test_fn! {
#[allow(dead_code)]
mod_power_of_2_mul_truncated_to_out_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
pow: u64,
) {
assert_lengths(out, xs, ys);
assert!(pow <= T::WIDTH);
mul_truncated_karatsuba_wrapping(out, xs, ys);
mask_coefficients(out, pow);
}}
#[doc(hidden)]
pub fn mod_power_of_2_mul_truncated_to_out<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
pow: u64,
) {
assert_lengths(out, xs, ys);
assert!(pow <= T::WIDTH);
mul_truncated_karatsuba_wrapping(out, xs, ys);
mask_coefficients(out, pow);
}
pub(crate) fn truncated_len(len1: usize, len2: usize, len: u64) -> usize {
min(usize::try_from(len).unwrap_or(usize::MAX), len1 + len2 - 1)
}
pub(crate) fn mod_power_of_2_mul_truncated_helper<T: PrimitiveUnsigned>(
xs: &[T],
ys: &[T],
len: u64,
pow: u64,
) -> UnsignedPolynomial<T> {
if len == 0 || xs.is_empty() || ys.is_empty() {
return UnsignedPolynomial::ZERO;
}
let mut out = vec![T::ZERO; truncated_len(xs.len(), ys.len(), len)];
mod_power_of_2_mul_truncated_to_out(&mut out, xs, ys, pow);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulTruncated<Self> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_power_of_2_mul_truncated(self, other: Self, len: u64, pow: u64) -> Self {
assert_reduced(&self, &other, pow);
mod_power_of_2_mul_truncated_helper(&self.coefficients, &other.coefficients, len, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulTruncated<&Self> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_power_of_2_mul_truncated(self, other: &Self, len: u64, pow: u64) -> Self {
assert_reduced(&self, other, pow);
mod_power_of_2_mul_truncated_helper(&self.coefficients, &other.coefficients, len, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulTruncated<UnsignedPolynomial<T>>
for &UnsignedPolynomial<T>
{
type Output = UnsignedPolynomial<T>;
fn mod_power_of_2_mul_truncated(
self,
other: UnsignedPolynomial<T>,
len: u64,
pow: u64,
) -> UnsignedPolynomial<T> {
assert_reduced(self, &other, pow);
mod_power_of_2_mul_truncated_helper(&self.coefficients, &other.coefficients, len, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulTruncated<&UnsignedPolynomial<T>>
for &UnsignedPolynomial<T>
{
type Output = UnsignedPolynomial<T>;
fn mod_power_of_2_mul_truncated(
self,
other: &UnsignedPolynomial<T>,
len: u64,
pow: u64,
) -> UnsignedPolynomial<T> {
assert_reduced(self, other, pow);
mod_power_of_2_mul_truncated_helper(&self.coefficients, &other.coefficients, len, pow)
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulTruncatedAssign<Self> for UnsignedPolynomial<T> {
fn mod_power_of_2_mul_truncated_assign(&mut self, other: Self, len: u64, pow: u64) {
assert_reduced(self, &other, pow);
*self =
mod_power_of_2_mul_truncated_helper(&self.coefficients, &other.coefficients, len, pow);
}
}
impl<T: PrimitiveUnsigned> ModPowerOf2MulTruncatedAssign<&Self> for UnsignedPolynomial<T> {
fn mod_power_of_2_mul_truncated_assign(&mut self, other: &Self, len: u64, pow: u64) {
assert_reduced(self, other, pow);
*self =
mod_power_of_2_mul_truncated_helper(&self.coefficients, &other.coefficients, len, pow);
}
}