malachite_base/unsigned_polynomial/arithmetic/mod_power_of_2_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::{
9 ModPowerOf2IsReduced, ModPowerOf2Square, ModPowerOf2SquareAssign,
10};
11use crate::num::basic::traits::Zero;
12use crate::num::basic::unsigneds::PrimitiveUnsigned;
13use crate::unsigned_polynomial::UnsignedPolynomial;
14use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::{
15 MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD, add_wrapping_assign, from_coefficients_trimmed,
16 karatsuba_wrapping_scratch_len, mask_coefficients, sub_wrapping_assign,
17};
18use alloc::vec;
19
20// Sets `out` to the square of the polynomial with coefficients `xs`, modulo $2^\text{W}$, by
21// schoolbook multiplication: each product of two different coefficients is computed once and
22// doubled. `out` must have length `2 * xs.len() - 1`.
23pub(crate) fn square_classical_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T]) {
24 out.fill(T::ZERO);
25 for (i, &x) in xs.iter().enumerate() {
26 if x != T::ZERO {
27 for (o, &y) in out[(i << 1) + 1..].iter_mut().zip(&xs[i + 1..]) {
28 o.wrapping_add_assign(x.wrapping_mul(y));
29 }
30 }
31 }
32 for o in out.iter_mut() {
33 *o = o.wrapping_add(*o);
34 }
35 for (i, &x) in xs.iter().enumerate() {
36 out[i << 1].wrapping_add_assign(x.wrapping_mul(x));
37 }
38}
39
40// Sets `out` to the square of the polynomial with coefficients `xs`, of nonzero length $n$, modulo
41// $2^\text{W}$, by Karatsuba multiplication, falling back to schoolbook multiplication below the
42// threshold. `out` must have length $2n - 1$, and `scratch` at least
43// `karatsuba_wrapping_scratch_len(n)`.
44fn square_karatsuba_scratch_wrapping<T: PrimitiveUnsigned>(
45 out: &mut [T],
46 xs: &[T],
47 scratch: &mut [T],
48) {
49 let n = xs.len();
50 if n < MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD {
51 square_classical_wrapping(out, xs);
52 return;
53 }
54 // 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}
55 // x_1^2.
56 let h = n >> 1;
57 let c = n - h;
58 let two_h = h << 1;
59 let (x0, x1) = xs.split_at(h);
60 let (sum, scratch) = scratch.split_at_mut(c);
61 let (middle, scratch) = scratch.split_at_mut((c << 1) - 1);
62 let (low, high) = out.split_at_mut(two_h);
63 square_karatsuba_scratch_wrapping(&mut low[..two_h - 1], x0, scratch);
64 low[two_h - 1] = T::ZERO;
65 square_karatsuba_scratch_wrapping(high, x1, scratch);
66 sum.copy_from_slice(x1);
67 add_wrapping_assign(sum, x0);
68 square_karatsuba_scratch_wrapping(middle, sum, scratch);
69 sub_wrapping_assign(middle, &out[..two_h - 1]);
70 sub_wrapping_assign(middle, &out[two_h..]);
71 add_wrapping_assign(&mut out[h..], middle);
72}
73
74// Sets `out` to the square of the polynomial with coefficients `xs`, which is nonempty, modulo
75// $2^\text{W}$, by Karatsuba multiplication, falling back to schoolbook multiplication for short
76// polynomials. `out` must have length `2 * xs.len() - 1`.
77pub(crate) fn square_karatsuba_wrapping<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T]) {
78 if xs.len() < MOD_POWER_OF_2_SQUARE_KARATSUBA_THRESHOLD {
79 square_classical_wrapping(out, xs);
80 return;
81 }
82 let mut scratch = vec![T::ZERO; karatsuba_wrapping_scratch_len(xs.len())];
83 square_karatsuba_scratch_wrapping(out, xs, &mut scratch);
84}
85
86fn assert_lengths<T>(out: &[T], xs: &[T]) {
87 assert!(!xs.is_empty());
88 assert_eq!(out.len(), (xs.len() << 1) - 1);
89}
90
91// Sets `out` to the square of the polynomial with coefficients `xs`, nonempty and reduced modulo
92// $2^k$, where $k$ is `pow`, by schoolbook multiplication. `out` must have length `2 * xs.len() -
93// 1`, and `pow` must be no greater than `T::WIDTH`.
94crate_test_fn! {
95#[allow(dead_code)]
96mod_power_of_2_square_to_out_classical<T: PrimitiveUnsigned>(
97 out: &mut [T],
98 xs: &[T],
99 pow: u64,
100) {
101 assert_lengths(out, xs);
102 assert!(pow <= T::WIDTH);
103 square_classical_wrapping(out, xs);
104 mask_coefficients(out, pow);
105}}
106
107// Sets `out` to the square of the polynomial with coefficients `xs`, nonempty and reduced modulo
108// $2^k$, where $k$ is `pow`, by Karatsuba multiplication. `out` must have length `2 * xs.len() -
109// 1`, and `pow` must be no greater than `T::WIDTH`.
110crate_test_fn! {
111#[allow(dead_code)]
112mod_power_of_2_square_to_out_karatsuba<T: PrimitiveUnsigned>(
113 out: &mut [T],
114 xs: &[T],
115 pow: u64,
116) {
117 assert_lengths(out, xs);
118 assert!(pow <= T::WIDTH);
119 square_karatsuba_wrapping(out, xs);
120 mask_coefficients(out, pow);
121}}
122
123/// Sets `out` to the square of the polynomial with coefficients `xs`, nonempty and reduced modulo
124/// $2^k$, where $k$ is `pow`. `out` must have length `2 * xs.len() - 1`, and `pow` must be no
125/// greater than `T::WIDTH`.
126///
127/// This is not part of the public API; it is public so that `malachite-nz` can square
128/// `NaturalPolynomial`s with word-sized coefficients modulo $2^k$.
129#[doc(hidden)]
130pub fn mod_power_of_2_square_to_out<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], pow: u64) {
131 assert_lengths(out, xs);
132 assert!(pow <= T::WIDTH);
133 square_karatsuba_wrapping(out, xs);
134 mask_coefficients(out, pow);
135}
136
137fn assert_reduced<T: PrimitiveUnsigned>(p: &UnsignedPolynomial<T>, pow: u64) {
138 assert!(pow <= T::WIDTH);
139 assert!(
140 p.mod_power_of_2_is_reduced(pow),
141 "self must be reduced mod 2^pow, but {p} has a coefficient >= 2^{pow}"
142 );
143}
144
145// The square of the polynomial with coefficients `xs`, reduced modulo $2^k$, where $k$ is `pow`,
146// modulo $2^k$.
147pub(crate) fn mod_power_of_2_square_helper<T: PrimitiveUnsigned>(
148 xs: &[T],
149 pow: u64,
150) -> UnsignedPolynomial<T> {
151 if xs.is_empty() {
152 return UnsignedPolynomial::ZERO;
153 }
154 let mut out = vec![T::ZERO; (xs.len() << 1) - 1];
155 mod_power_of_2_square_to_out(&mut out, xs, pow);
156 from_coefficients_trimmed(out)
157}
158
159impl<T: PrimitiveUnsigned> ModPowerOf2Square for UnsignedPolynomial<T> {
160 type Output = Self;
161
162 /// Squares an [`UnsignedPolynomial`] modulo $2^k$, taking it by value. Its coefficients must
163 /// already be reduced modulo $2^k$.
164 ///
165 /// $$
166 /// f(p, k) = p^2 \bmod 2^k.
167 /// $$
168 ///
169 /// The leading coefficient of the square can vanish modulo $2^k$, and then the degree of the
170 /// square is lower than twice the degree.
171 ///
172 /// # Worst-case complexity
173 /// $T(n) = O(n^{\log_2 3})$
174 ///
175 /// $M(n) = O(n)$
176 ///
177 /// where $T$ is time, $M$ is additional memory, and $n$ is `self.len()`.
178 ///
179 /// # Panics
180 /// Panics if `pow` is greater than `T::WIDTH`, or if any coefficient of `self` is greater than
181 /// or equal to $2^k$.
182 ///
183 /// # Examples
184 /// ```
185 /// use core::str::FromStr;
186 /// use malachite_base::num::arithmetic::traits::ModPowerOf2Square;
187 /// use malachite_base::unsigned_polynomial::UnsignedPolynomial;
188 ///
189 /// // The square is x^4+6*x^3+13*x^2+12*x+4; its coefficients wrap around modulo 8.
190 /// assert_eq!(
191 /// UnsignedPolynomial::<u8>::from_str("x^2+3*x+2")
192 /// .unwrap()
193 /// .mod_power_of_2_square(3)
194 /// .to_string(),
195 /// "x^4+6*x^3+5*x^2+4*x+4"
196 /// );
197 /// // The square is 16*x^2+8*x+1, which is 1 modulo 8.
198 /// assert_eq!(
199 /// UnsignedPolynomial::<u8>::from_str("4*x+1")
200 /// .unwrap()
201 /// .mod_power_of_2_square(3)
202 /// .to_string(),
203 /// "1"
204 /// );
205 /// ```
206 ///
207 /// This is equivalent to `nmod_poly_mul` from `nmod_poly/mul.c`, FLINT 3.6.0, with both factors
208 /// equal and the modulus $2^k$.
209 fn mod_power_of_2_square(self, pow: u64) -> Self {
210 assert_reduced(&self, pow);
211 mod_power_of_2_square_helper(&self.coefficients, pow)
212 }
213}
214
215impl<T: PrimitiveUnsigned> ModPowerOf2Square for &UnsignedPolynomial<T> {
216 type Output = UnsignedPolynomial<T>;
217
218 /// Squares an [`UnsignedPolynomial`] modulo $2^k$, taking it by reference. Its coefficients
219 /// must already be reduced modulo $2^k$.
220 ///
221 /// $$
222 /// f(p, k) = p^2 \bmod 2^k.
223 /// $$
224 ///
225 /// The leading coefficient of the square can vanish modulo $2^k$, and then the degree of the
226 /// square is lower than twice the degree.
227 ///
228 /// # Worst-case complexity
229 /// $T(n) = O(n^{\log_2 3})$
230 ///
231 /// $M(n) = O(n)$
232 ///
233 /// where $T$ is time, $M$ is additional memory, and $n$ is `self.len()`.
234 ///
235 /// # Panics
236 /// Panics if `pow` is greater than `T::WIDTH`, or if any coefficient of `self` is greater than
237 /// or equal to $2^k$.
238 ///
239 /// # Examples
240 /// ```
241 /// use core::str::FromStr;
242 /// use malachite_base::num::arithmetic::traits::ModPowerOf2Square;
243 /// use malachite_base::unsigned_polynomial::UnsignedPolynomial;
244 ///
245 /// // The square is x^4+6*x^3+13*x^2+12*x+4; its coefficients wrap around modulo 8.
246 /// assert_eq!(
247 /// (&UnsignedPolynomial::<u8>::from_str("x^2+3*x+2").unwrap())
248 /// .mod_power_of_2_square(3)
249 /// .to_string(),
250 /// "x^4+6*x^3+5*x^2+4*x+4"
251 /// );
252 /// // The square is 16*x^2+8*x+1, which is 1 modulo 8.
253 /// assert_eq!(
254 /// (&UnsignedPolynomial::<u8>::from_str("4*x+1").unwrap())
255 /// .mod_power_of_2_square(3)
256 /// .to_string(),
257 /// "1"
258 /// );
259 /// ```
260 ///
261 /// This is equivalent to `nmod_poly_mul` from `nmod_poly/mul.c`, FLINT 3.6.0, with both factors
262 /// equal and the modulus $2^k$.
263 fn mod_power_of_2_square(self, pow: u64) -> UnsignedPolynomial<T> {
264 assert_reduced(self, pow);
265 mod_power_of_2_square_helper(&self.coefficients, pow)
266 }
267}
268
269impl<T: PrimitiveUnsigned> ModPowerOf2SquareAssign for UnsignedPolynomial<T> {
270 /// Squares an [`UnsignedPolynomial`] modulo $2^k$ in place. Its coefficients must already be
271 /// reduced modulo $2^k$.
272 ///
273 /// $$
274 /// p \gets p^2 \bmod 2^k.
275 /// $$
276 ///
277 /// # Worst-case complexity
278 /// $T(n) = O(n^{\log_2 3})$
279 ///
280 /// $M(n) = O(n)$
281 ///
282 /// where $T$ is time, $M$ is additional memory, and $n$ is `self.len()`.
283 ///
284 /// # Panics
285 /// Panics if `pow` is greater than `T::WIDTH`, or if any coefficient of `self` is greater than
286 /// or equal to $2^k$.
287 ///
288 /// # Examples
289 /// ```
290 /// use core::str::FromStr;
291 /// use malachite_base::num::arithmetic::traits::ModPowerOf2SquareAssign;
292 /// use malachite_base::unsigned_polynomial::UnsignedPolynomial;
293 ///
294 /// let mut p = UnsignedPolynomial::<u8>::from_str("x^2+3*x+2").unwrap();
295 /// p.mod_power_of_2_square_assign(3);
296 /// assert_eq!(p.to_string(), "x^4+6*x^3+5*x^2+4*x+4");
297 /// ```
298 ///
299 /// This is equivalent to `nmod_poly_mul` from `nmod_poly/mul.c`, FLINT 3.6.0, with both factors
300 /// equal and the modulus $2^k$.
301 fn mod_power_of_2_square_assign(&mut self, pow: u64) {
302 assert_reduced(self, pow);
303 *self = mod_power_of_2_square_helper(&self.coefficients, pow);
304 }
305}