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}