use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::polynomial::{ModMulTruncated, ModMulTruncatedAssign};
use crate::unsigned_polynomial::UnsignedPolynomial;
use crate::unsigned_polynomial::arithmetic::mod_add::assert_reduced;
use crate::unsigned_polynomial::arithmetic::mod_mul::{
MOD_MUL_KARATSUBA_THRESHOLD, ModData, column_sum, mod_add_assign_slice, mod_mul_karatsuba,
};
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::from_coefficients_trimmed;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul_truncated::truncated_len;
use alloc::vec;
use core::cmp::min;
pub(crate) fn mod_mul_truncated_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<T>,
) {
let n = xs.len();
let m = ys.len();
for (k, o) in out.iter_mut().enumerate() {
if k > n + m - 2 {
*o = T::ZERO;
continue;
}
let start = k.saturating_sub(m - 1);
let stop = min(k, n - 1);
let acc = column_sum(&xs[start..=stop], &ys[k - stop..=k - start], d);
*o = d.reduce_sum(acc);
}
}
pub(crate) fn mod_mul_truncated_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
d: &ModData<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 {
mod_mul_karatsuba(&mut out[..full_len], xs, ys, d);
out[full_len..].fill(T::ZERO);
return;
}
if ys.len() < MOD_MUL_KARATSUBA_THRESHOLD {
mod_mul_truncated_classical(out, xs, ys, d);
return;
}
let h = len.div_ceil(2);
let x0 = &xs[..min(h, xs.len())];
let y0 = &ys[..min(h, ys.len())];
mod_mul_truncated_karatsuba(out, x0, y0, d);
let mut cross = vec![T::ZERO; len - h];
if xs.len() > h {
mod_mul_truncated_karatsuba(&mut cross, &xs[h..], y0, d);
mod_add_assign_slice(&mut out[h..], &cross, d.m);
}
if ys.len() > h {
mod_mul_truncated_karatsuba(&mut cross, x0, &ys[h..], d);
mod_add_assign_slice(&mut out[h..], &cross, d.m);
}
}
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_mul_truncated_to_out_classical<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
m: T,
) {
assert_lengths(out, xs, ys);
let terms = xs.len().min(ys.len()).min(out.len());
mod_mul_truncated_classical(out, xs, ys, &ModData::new(m, terms));
}}
crate_test_fn! {
#[allow(dead_code)]
mod_mul_truncated_to_out_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
ys: &[T],
m: T,
) {
assert_lengths(out, xs, ys);
let terms = xs.len().min(ys.len()).min(out.len());
mod_mul_truncated_karatsuba(out, xs, ys, &ModData::new(m, terms));
}}
#[doc(hidden)]
pub fn mod_mul_truncated_to_out<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], ys: &[T], m: T) {
assert_lengths(out, xs, ys);
mod_mul_truncated_karatsuba(
out,
xs,
ys,
&ModData::new(m, xs.len().min(ys.len()).min(out.len())),
);
}
pub(crate) fn mod_mul_truncated_helper<T: PrimitiveUnsigned>(
xs: &[T],
ys: &[T],
len: u64,
m: T,
) -> 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_mul_truncated_to_out(&mut out, xs, ys, m);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModMulTruncated<Self, T> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_mul_truncated(self, other: Self, len: u64, m: T) -> Self {
assert_reduced(&self, &other, m);
mod_mul_truncated_helper(&self.coefficients, &other.coefficients, len, m)
}
}
impl<T: PrimitiveUnsigned> ModMulTruncated<&Self, T> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_mul_truncated(self, other: &Self, len: u64, m: T) -> Self {
assert_reduced(&self, other, m);
mod_mul_truncated_helper(&self.coefficients, &other.coefficients, len, m)
}
}
impl<T: PrimitiveUnsigned> ModMulTruncated<UnsignedPolynomial<T>, T> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_mul_truncated(
self,
other: UnsignedPolynomial<T>,
len: u64,
m: T,
) -> UnsignedPolynomial<T> {
assert_reduced(self, &other, m);
mod_mul_truncated_helper(&self.coefficients, &other.coefficients, len, m)
}
}
impl<T: PrimitiveUnsigned> ModMulTruncated<&UnsignedPolynomial<T>, T> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_mul_truncated(
self,
other: &UnsignedPolynomial<T>,
len: u64,
m: T,
) -> UnsignedPolynomial<T> {
assert_reduced(self, other, m);
mod_mul_truncated_helper(&self.coefficients, &other.coefficients, len, m)
}
}
impl<T: PrimitiveUnsigned> ModMulTruncatedAssign<Self, T> for UnsignedPolynomial<T> {
fn mod_mul_truncated_assign(&mut self, other: Self, len: u64, m: T) {
assert_reduced(self, &other, m);
*self = mod_mul_truncated_helper(&self.coefficients, &other.coefficients, len, m);
}
}
impl<T: PrimitiveUnsigned> ModMulTruncatedAssign<&Self, T> for UnsignedPolynomial<T> {
fn mod_mul_truncated_assign(&mut self, other: &Self, len: u64, m: T) {
assert_reduced(self, other, m);
*self = mod_mul_truncated_helper(&self.coefficients, &other.coefficients, len, m);
}
}