Skip to main content

malachite_base/unsigned_polynomial/arithmetic/
mod_square.rs

1// Copyright © 2026 Mikhail Hogrefe
2//
3// This file is part of Malachite.
4//
5// Malachite is free software: you can redistribute it and/or modify it under the terms of the GNU
6// Lesser General Public License (LGPL) as published by the Free Software Foundation; either version
7// 3 of the License, or (at your option) any later version. See <https://www.gnu.org/licenses/>.
8use crate::num::arithmetic::traits::{ModIsReduced, ModSquare, ModSquareAssign, Parity};
9use crate::num::basic::traits::Zero;
10use crate::num::basic::unsigneds::PrimitiveUnsigned;
11use crate::unsigned_polynomial::UnsignedPolynomial;
12use crate::unsigned_polynomial::arithmetic::mod_mul::{
13    MOD_SQUARE_KARATSUBA_THRESHOLD, ModData, accumulate, column_sum, mod_add_assign_slice,
14    mod_karatsuba_scratch_len, mod_sub_assign_slice,
15};
16use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::from_coefficients_trimmed;
17use alloc::vec;
18
19// Doubles the three-word accumulator `(a2, a1, a0)`, which must be less than $2^{3\text{W} - 1}$;
20// sums from `column_sum` are.
21#[inline]
22pub(crate) fn double<T: PrimitiveUnsigned>(acc: &mut (T, T, T)) {
23    let (a2, a1, a0) = *acc;
24    let top = T::WIDTH - 1;
25    *acc = ((a2 << 1) | (a1 >> top), (a1 << 1) | (a0 >> top), a0 << 1);
26}
27
28// Sets `out[k]`, for each `k` less than `out.len()`, to coefficient `k` of the square of the
29// polynomial with coefficients `xs`, reduced modulo $m$, by schoolbook multiplication: each product
30// of two different coefficients is accumulated once and doubled.
31pub(crate) fn mod_square_classical_prefix<T: PrimitiveUnsigned>(
32    out: &mut [T],
33    xs: &[T],
34    d: &ModData<T>,
35) {
36    let n = xs.len();
37    for (k, o) in out.iter_mut().enumerate() {
38        let start = k.saturating_sub(n - 1);
39        // Pairs (i, k - i) with i < k - i.
40        let stop = (k + 1) >> 1;
41        let mut acc = if start < stop {
42            column_sum(&xs[start..stop], &xs[k + 1 - stop..=k - start], d)
43        } else {
44            (T::ZERO, T::ZERO, T::ZERO)
45        };
46        double(&mut acc);
47        if k.even() && (k >> 1) < n {
48            let x = xs[k >> 1];
49            accumulate(&mut acc, x, x);
50        }
51        *o = d.reduce_sum(acc);
52    }
53}
54
55// Sets `out` to the square of the polynomial with coefficients `xs`, of nonzero length $n$, modulo
56// $m$, by Karatsuba multiplication, falling back to schoolbook multiplication below the threshold.
57// `out` must have length $2n - 1$.
58fn mod_square_karatsuba_scratch<T: PrimitiveUnsigned>(
59    out: &mut [T],
60    xs: &[T],
61    d: &ModData<T>,
62    scratch: &mut [T],
63) {
64    let n = xs.len();
65    if n < MOD_SQUARE_KARATSUBA_THRESHOLD {
66        mod_square_classical_prefix(out, xs, d);
67        return;
68    }
69    // Write x = x_0 + x^h x_1. Then x^2 = x_0^2 + x^h ((x_0 + x_1)^2 - x_0^2 - x_1^2) + x^{2h}
70    // x_1^2.
71    let m = d.m;
72    let h = n >> 1;
73    let c = n - h;
74    let two_h = h << 1;
75    let (x0, x1) = xs.split_at(h);
76    let (sum, scratch) = scratch.split_at_mut(c);
77    let (middle, scratch) = scratch.split_at_mut((c << 1) - 1);
78    let (low, high) = out.split_at_mut(two_h);
79    mod_square_karatsuba_scratch(&mut low[..two_h - 1], x0, d, scratch);
80    low[two_h - 1] = T::ZERO;
81    mod_square_karatsuba_scratch(high, x1, d, scratch);
82    sum.copy_from_slice(x1);
83    mod_add_assign_slice(sum, x0, m);
84    mod_square_karatsuba_scratch(middle, sum, d, scratch);
85    mod_sub_assign_slice(middle, &out[..two_h - 1], m);
86    mod_sub_assign_slice(middle, &out[two_h..], m);
87    mod_add_assign_slice(&mut out[h..], middle, m);
88}
89
90// Sets `out` to the square of the polynomial with coefficients `xs`, which is nonempty, modulo $m$,
91// by Karatsuba multiplication, falling back to schoolbook multiplication for short polynomials.
92pub(crate) fn mod_square_karatsuba<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], d: &ModData<T>) {
93    if xs.len() < MOD_SQUARE_KARATSUBA_THRESHOLD {
94        mod_square_classical_prefix(out, xs, d);
95        return;
96    }
97    let mut scratch =
98        vec![T::ZERO; mod_karatsuba_scratch_len(xs.len(), MOD_SQUARE_KARATSUBA_THRESHOLD)];
99    mod_square_karatsuba_scratch(out, xs, d, &mut scratch);
100}
101
102fn assert_lengths<T>(out: &[T], xs: &[T]) {
103    assert!(!xs.is_empty());
104    assert_eq!(out.len(), (xs.len() << 1) - 1);
105}
106
107// Sets `out` to the square of the polynomial with coefficients `xs`, nonempty and reduced modulo
108// `m`, modulo `m`, by schoolbook multiplication. `out` must have length `2 * xs.len() - 1`.
109crate_test_fn! {
110#[allow(dead_code)]
111mod_square_to_out_classical<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
112    assert_lengths(out, xs);
113    mod_square_classical_prefix(out, xs, &ModData::new(m, xs.len()));
114}}
115
116// Sets `out` to the square of the polynomial with coefficients `xs`, nonempty and reduced modulo
117// `m`, modulo `m`, by Karatsuba multiplication. `out` must have length `2 * xs.len() - 1`.
118crate_test_fn! {
119#[allow(dead_code)]
120mod_square_to_out_karatsuba<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
121    assert_lengths(out, xs);
122    mod_square_karatsuba(out, xs, &ModData::new(m, xs.len()));
123}}
124
125/// Sets `out` to the square of the polynomial with coefficients `xs`, nonempty and reduced modulo
126/// `m`, modulo `m`. `out` must have length `2 * xs.len() - 1`, and `m` must be positive.
127///
128/// This is not part of the public API; it is public so that `malachite-nz` can square
129/// `NaturalPolynomial`s modulo a word.
130#[doc(hidden)]
131pub fn mod_square_to_out<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
132    assert_lengths(out, xs);
133    mod_square_karatsuba(out, xs, &ModData::new(m, xs.len()));
134}
135
136fn assert_reduced<T: PrimitiveUnsigned>(p: &UnsignedPolynomial<T>, m: T) {
137    assert!(
138        p.mod_is_reduced(&m),
139        "self must be reduced mod m, but {p} has a coefficient >= {m}"
140    );
141}
142
143// The square of the polynomial with coefficients `xs`, reduced modulo `m`, modulo `m`.
144pub(crate) fn mod_square_helper<T: PrimitiveUnsigned>(xs: &[T], m: T) -> UnsignedPolynomial<T> {
145    if xs.is_empty() {
146        return UnsignedPolynomial::ZERO;
147    }
148    let mut out = vec![T::ZERO; (xs.len() << 1) - 1];
149    mod_square_to_out(&mut out, xs, m);
150    from_coefficients_trimmed(out)
151}
152
153impl<T: PrimitiveUnsigned> ModSquare<T> for UnsignedPolynomial<T> {
154    type Output = Self;
155
156    /// Squares an [`UnsignedPolynomial`] modulo $m$, taking it by value. Its coefficients must
157    /// already be reduced modulo $m$.
158    ///
159    /// $$
160    /// f(p, m) = p^2 \bmod m.
161    /// $$
162    ///
163    /// When $m$ is not prime, the leading coefficient of the square can vanish modulo $m$, and then
164    /// the degree of the square is lower than twice the degree.
165    ///
166    /// # Worst-case complexity
167    /// $T(n) = O(n^{\log_2 3})$
168    ///
169    /// $M(n) = O(n)$
170    ///
171    /// where $T$ is time, $M$ is additional memory, and $n$ is `self.len()`.
172    ///
173    /// # Panics
174    /// Panics if `m` is 0, or if any coefficient of `self` is greater than or equal to `m`.
175    ///
176    /// # Examples
177    /// ```
178    /// use core::str::FromStr;
179    /// use malachite_base::num::arithmetic::traits::ModSquare;
180    /// use malachite_base::unsigned_polynomial::UnsignedPolynomial;
181    ///
182    /// // The square is x^4+6*x^3+13*x^2+12*x+4; its coefficients modulo 7.
183    /// assert_eq!(
184    ///     UnsignedPolynomial::<u8>::from_str("x^2+3*x+2")
185    ///         .unwrap()
186    ///         .mod_square(7)
187    ///         .to_string(),
188    ///     "x^4+6*x^3+6*x^2+5*x+4"
189    /// );
190    /// // The square is 4*x^2+4*x+1, which is 1 modulo 4.
191    /// assert_eq!(
192    ///     UnsignedPolynomial::<u8>::from_str("2*x+1")
193    ///         .unwrap()
194    ///         .mod_square(4)
195    ///         .to_string(),
196    ///     "1"
197    /// );
198    /// ```
199    ///
200    /// This is equivalent to `nmod_poly_mul` from `nmod_poly/mul.c`, FLINT 3.6.0, with both factors
201    /// equal.
202    fn mod_square(self, m: T) -> Self {
203        assert_reduced(&self, m);
204        mod_square_helper(&self.coefficients, m)
205    }
206}
207
208impl<T: PrimitiveUnsigned> ModSquare<T> for &UnsignedPolynomial<T> {
209    type Output = UnsignedPolynomial<T>;
210
211    /// Squares an [`UnsignedPolynomial`] modulo $m$, taking it by reference. Its coefficients must
212    /// already be reduced modulo $m$.
213    ///
214    /// $$
215    /// f(p, m) = p^2 \bmod m.
216    /// $$
217    ///
218    /// When $m$ is not prime, the leading coefficient of the square can vanish modulo $m$, and then
219    /// the degree of the square is lower than twice the degree.
220    ///
221    /// # Worst-case complexity
222    /// $T(n) = O(n^{\log_2 3})$
223    ///
224    /// $M(n) = O(n)$
225    ///
226    /// where $T$ is time, $M$ is additional memory, and $n$ is `self.len()`.
227    ///
228    /// # Panics
229    /// Panics if `m` is 0, or if any coefficient of `self` is greater than or equal to `m`.
230    ///
231    /// # Examples
232    /// ```
233    /// use core::str::FromStr;
234    /// use malachite_base::num::arithmetic::traits::ModSquare;
235    /// use malachite_base::unsigned_polynomial::UnsignedPolynomial;
236    ///
237    /// // The square is x^4+6*x^3+13*x^2+12*x+4; its coefficients modulo 7.
238    /// assert_eq!(
239    ///     (&UnsignedPolynomial::<u8>::from_str("x^2+3*x+2").unwrap())
240    ///         .mod_square(7)
241    ///         .to_string(),
242    ///     "x^4+6*x^3+6*x^2+5*x+4"
243    /// );
244    /// // The square is 4*x^2+4*x+1, which is 1 modulo 4.
245    /// assert_eq!(
246    ///     (&UnsignedPolynomial::<u8>::from_str("2*x+1").unwrap())
247    ///         .mod_square(4)
248    ///         .to_string(),
249    ///     "1"
250    /// );
251    /// ```
252    ///
253    /// This is equivalent to `nmod_poly_mul` from `nmod_poly/mul.c`, FLINT 3.6.0, with both factors
254    /// equal.
255    fn mod_square(self, m: T) -> UnsignedPolynomial<T> {
256        assert_reduced(self, m);
257        mod_square_helper(&self.coefficients, m)
258    }
259}
260
261impl<T: PrimitiveUnsigned> ModSquareAssign<T> for UnsignedPolynomial<T> {
262    /// Squares an [`UnsignedPolynomial`] modulo $m$ in place. Its coefficients must already be
263    /// reduced modulo $m$.
264    ///
265    /// $$
266    /// p \gets p^2 \bmod m.
267    /// $$
268    ///
269    /// # Worst-case complexity
270    /// $T(n) = O(n^{\log_2 3})$
271    ///
272    /// $M(n) = O(n)$
273    ///
274    /// where $T$ is time, $M$ is additional memory, and $n$ is `self.len()`.
275    ///
276    /// # Panics
277    /// Panics if `m` is 0, or if any coefficient of `self` is greater than or equal to `m`.
278    ///
279    /// # Examples
280    /// ```
281    /// use core::str::FromStr;
282    /// use malachite_base::num::arithmetic::traits::ModSquareAssign;
283    /// use malachite_base::unsigned_polynomial::UnsignedPolynomial;
284    ///
285    /// let mut p = UnsignedPolynomial::<u8>::from_str("x^2+3*x+2").unwrap();
286    /// p.mod_square_assign(7);
287    /// assert_eq!(p.to_string(), "x^4+6*x^3+6*x^2+5*x+4");
288    /// ```
289    ///
290    /// This is equivalent to `nmod_poly_mul` from `nmod_poly/mul.c`, FLINT 3.6.0, with both factors
291    /// equal.
292    fn mod_square_assign(&mut self, m: T) {
293        assert_reduced(self, m);
294        *self = mod_square_helper(&self.coefficients, m);
295    }
296}