binius_math/batch_invert.rs
1// Copyright 2025-2026 The Binius Developers
2// Copyright 2025 Irreducible Inc.
3
4//! Batch multiplicative inversion via Montgomery's trick.
5
6use std::iter;
7
8use binius_field::{Field, PackedField};
9
10/// Reusable batch inversion context that owns its scratch buffers.
11///
12/// Reusing one instance across many same-size calls avoids reallocating on every call.
13pub struct BatchInversion<P: PackedField> {
14 /// Number of packed elements this instance is sized for.
15 n: usize,
16 /// Scratch space used by the pairwise-tree recursion.
17 scratchpad: Vec<P>,
18 /// Flat scalar index of every zero found by the most recent zero-tolerant call.
19 zero_indices: Vec<usize>,
20}
21
22impl<P: PackedField> BatchInversion<P> {
23 /// Creates a new batch inversion context sized for `n` packed elements.
24 ///
25 /// # Arguments
26 /// - `n`: the number of packed elements every future call must be invoked with.
27 ///
28 /// # Panics
29 /// Panics if `n` is 0.
30 pub fn new(n: usize) -> Self {
31 // No elements to invert, and nothing to allocate, when n is 0.
32 assert!(n > 0, "n must be greater than 0");
33
34 Self {
35 n,
36 scratchpad: vec![P::zero(); min_scratchpad_size(n)],
37 zero_indices: Vec::new(),
38 }
39 }
40
41 /// Inverts every element of the slice in place.
42 ///
43 /// # Arguments
44 /// - `elements`: the slice to invert in place.
45 ///
46 /// # Safety
47 /// Every scalar element must be non-zero.
48 /// Behavior is undefined if any scalar is zero.
49 ///
50 /// # Panics
51 /// Panics if the slice length does not equal the `n` given at construction.
52 pub fn invert_nonzero(&mut self, elements: &mut [P]) {
53 assert_eq!(
54 elements.len(),
55 self.n,
56 "elements.len() must equal n (expected {}, got {})",
57 self.n,
58 elements.len()
59 );
60
61 self.batch_invert_nonzero(elements);
62 }
63
64 /// Inverts every element of the slice in place, leaving zero elements as zero.
65 ///
66 /// # Arguments
67 /// - `elements`: the slice to invert in place.
68 ///
69 /// # Panics
70 /// Panics if the slice length does not equal the `n` given at construction.
71 pub fn invert_or_zero(&mut self, elements: &mut [P]) {
72 assert_eq!(
73 elements.len(),
74 self.n,
75 "elements.len() must equal n (expected {}, got {})",
76 self.n,
77 elements.len()
78 );
79
80 // Zero has no inverse, so swap every zero scalar for a one, recording where.
81 self.zero_indices.clear();
82 for (packed_idx, packed) in elements.iter_mut().enumerate() {
83 for lane in 0..P::WIDTH {
84 if packed.get(lane) == P::Scalar::ZERO {
85 packed.set(lane, P::Scalar::ONE);
86 self.zero_indices.push(packed_idx * P::WIDTH + lane);
87 }
88 }
89 }
90
91 // Every scalar is non-zero now, so batch-invert directly.
92 self.invert_nonzero(elements);
93
94 // Restore the zeros — inverting one just gives one back.
95 for &scalar_idx in &self.zero_indices {
96 elements[scalar_idx / P::WIDTH].set(scalar_idx % P::WIDTH, P::Scalar::ZERO);
97 }
98 }
99
100 /// Runs the pairwise-tree inversion using this context's own scratch buffer.
101 fn batch_invert_nonzero(&mut self, elements: &mut [P]) {
102 batch_invert_nonzero_with_scratchpad(elements, &mut self.scratchpad);
103 }
104}
105
106/// Size of the scratchpad needed by the pairwise-tree recursion.
107///
108/// The recursion halves the element count at every level until it reaches 1.
109/// It needs one scratch slot per level: `ceil(n/2) + ceil(n/4) + ... + 1`.
110///
111/// # Arguments
112/// - `n`: the number of elements the recursion starts from.
113///
114/// # Returns
115/// The total number of scratch slots needed across every level below the top.
116///
117/// # Panics
118/// Panics if `n` is 0.
119fn min_scratchpad_size(mut n: usize) -> usize {
120 assert!(n > 0);
121
122 let mut size = 0;
123 // Sum each level's element count until only one element is left.
124 while n > 1 {
125 n = n.div_ceil(2);
126 size += n;
127 }
128 size
129}
130
131/// Inverts every element of the slice in place, organized as a balanced binary tree.
132///
133/// # Arguments
134/// - `elements`: the slice to invert in place.
135/// - `scratchpad`: scratch space for the recursion, with one slot per element at every level below
136/// the top.
137///
138/// # Safety
139/// Every element must be non-zero.
140/// Behavior is undefined if any scalar is zero.
141///
142/// # Algorithm
143/// Each level pairs element `i` with element `half + i`, multiplying to halve the count.
144/// Recursing down reaches a single combined product, which gets inverted directly.
145/// Unwinding multiplies that inverse back against the saved products.
146/// This recovers every individual inverse.
147///
148/// Elements paired at the same level never depend on each other.
149/// So the CPU can pipeline their multiplications instead of stalling on one chain.
150///
151/// Walking through 4 elements `a, b, c, d`:
152/// ```text
153/// elements: [ a, b, c, d ]
154/// pairwise products: [ a*c, b*d ] (pairs i with half + i)
155/// recurse to 1 element: invert (a*c)*(b*d) once
156/// unwind one level: [ (a*c)^-1, (b*d)^-1 ]
157/// unwind one level: [ a^-1, b^-1, c^-1, d^-1 ]
158/// ```
159fn batch_invert_nonzero_with_scratchpad<P: PackedField>(elements: &mut [P], scratchpad: &mut [P]) {
160 debug_assert!(!elements.is_empty());
161
162 if elements.len() == 1 {
163 // Safety: inputs are non-zero, so their product is non-zero in every lane.
164 // A packed type inverts every lane on its own — no manual unpacking needed.
165 elements[0] = unsafe { elements[0].invert() };
166 return;
167 }
168
169 // The next level's products go in the front of the scratch buffer.
170 // The rest stays free for deeper levels of the recursion.
171 let next_layer_len = elements.len().div_ceil(2);
172 let (next_layer, remaining) = scratchpad.split_at_mut(next_layer_len);
173
174 // Down: combine pairs into the next, half-as-long level.
175 product_layer(elements, next_layer);
176 // Recurse until a single combined product is left, then invert it directly.
177 batch_invert_nonzero_with_scratchpad(next_layer, remaining);
178 // Up: turn the single inverse for this level back into one inverse per element.
179 unproduct_layer(next_layer, elements);
180}
181
182/// Computes element-wise products of the top and bottom halves of a slice.
183///
184/// Pairs `input[i]` with `input[half + i]`.
185/// The middle element is copied through unpaired when the input length is odd.
186///
187/// # Arguments
188/// - `input`: the elements to pair up and multiply.
189/// - `output`: destination for the products, with length `input.len().div_ceil(2)`.
190///
191/// # Panics
192/// Panics in debug builds if `output.len() != input.len().div_ceil(2)`.
193#[inline]
194fn product_layer<P: PackedField>(input: &[P], output: &mut [P]) {
195 debug_assert_eq!(output.len(), input.len().div_ceil(2));
196
197 // The bottom half has exactly output.len() elements.
198 // The top half is whatever remains, one shorter when the length is odd.
199 let (lo, hi) = input.split_at(output.len());
200 let mut out_lo_iter = iter::zip(output, lo);
201
202 // Odd length: the last bottom-half element has no partner — copy it through.
203 if hi.len() < out_lo_iter.len() {
204 let Some((out_i, lo_i)) = out_lo_iter.next_back() else {
205 // Always called with 2 or more elements, so this iterator is never empty.
206 unreachable!("out_lo_iter.len() must be greater than zero");
207 };
208 *out_i = *lo_i;
209 }
210 // Every remaining pair has both halves: multiply them together.
211 for ((out_i, &lo_i), &hi_i) in iter::zip(out_lo_iter, hi) {
212 *out_i = lo_i * hi_i;
213 }
214}
215
216/// Unwinds a pairwise product pass to recover individual inverses.
217///
218/// Given inverted pair-products and the original paired values, recovers:
219/// - `output[i] = input[i] * output[half + i]` (inverse of the bottom-half element)
220/// - `output[half + i] = input[i] * output[i]` (inverse of the top-half element)
221///
222/// # Arguments
223/// - `input`: the inverted product for each pair, from the level above.
224/// - `output`: the original paired elements on entry, overwritten with their inverses.
225///
226/// # Panics
227/// Panics in debug builds if `input.len() != output.len().div_ceil(2)`.
228#[inline]
229fn unproduct_layer<P: PackedField>(input: &[P], output: &mut [P]) {
230 debug_assert_eq!(input.len(), output.len().div_ceil(2));
231
232 // Mirrors the split from the product pass.
233 // The bottom half pairs one-to-one with `input`.
234 // The top half is whatever remains.
235 let (lo, hi) = output.split_at_mut(input.len());
236 let mut lo_in_iter = iter::zip(lo, input);
237
238 // Odd length: the last element was unpaired, so its own product is its inverse.
239 if hi.len() < lo_in_iter.len() {
240 let Some((lo_i, in_i)) = lo_in_iter.next_back() else {
241 // Always called with 1 or more pairs, so this iterator is never empty.
242 unreachable!("out_lo_iter.len() must be greater than zero");
243 };
244 *lo_i = *in_i;
245 }
246 // Each pair recovers both halves, using their shared inverse and saved values.
247 for ((lo_i, &in_i), hi_i) in iter::zip(lo_in_iter, hi) {
248 let lo_tmp = *lo_i;
249 let hi_tmp = *hi_i;
250 *lo_i = in_i * hi_tmp;
251 *hi_i = in_i * lo_tmp;
252 }
253}
254
255#[cfg(test)]
256mod tests {
257 use binius_field::{Ghash128b, Random, arithmetic_traits::InvertOrZero};
258 use proptest::prelude::*;
259 use rand::{Rng, SeedableRng, rngs::StdRng, seq::IteratorRandom};
260
261 use super::*;
262
263 /// Shared helper to test batch inversion with a given inverter.
264 fn invert_with_inverter(
265 inverter: &mut BatchInversion<Ghash128b>,
266 n: usize,
267 n_zeros: usize,
268 rng: &mut impl Rng,
269 ) {
270 assert!(n_zeros <= n, "n_zeros must be <= n");
271
272 // Pick n_zeros distinct positions out of n to force to zero.
273 // Every other position stays random and non-zero.
274 let zero_indices: Vec<usize> = (0..n).sample(rng, n_zeros);
275
276 // Build the input slice from those positions.
277 let mut state = Vec::with_capacity(n);
278 for i in 0..n {
279 if zero_indices.contains(&i) {
280 state.push(Ghash128b::ZERO);
281 } else {
282 state.push(Ghash128b::random(&mut *rng));
283 }
284 }
285
286 // Reference result: invert every element independently, one call per element.
287 let expected: Vec<Ghash128b> = state
288 .iter()
289 .map(|x| InvertOrZero::invert_or_zero(*x))
290 .collect();
291
292 // Result under test: invert the whole batch through the zero-tolerant entry point.
293 inverter.invert_or_zero(&mut state);
294
295 // The batched result must match the per-element reference exactly, zeros included.
296 assert_eq!(state, expected);
297 }
298
299 fn test_batch_inversion_for_size(n: usize, n_zeros: usize, rng: &mut impl Rng) {
300 // Fresh context sized for exactly n elements.
301 let mut inverter = BatchInversion::<Ghash128b>::new(n);
302 invert_with_inverter(&mut inverter, n, n_zeros, rng);
303 }
304
305 fn test_batch_inversion_nonzero_for_size(n: usize, rng: &mut impl Rng) {
306 // Every element is random and non-zero.
307 // So the non-zero-only entry point is safe to use directly.
308 let mut state = Vec::with_capacity(n);
309 for _ in 0..n {
310 state.push(Ghash128b::random(&mut *rng));
311 }
312
313 // Reference result: invert every element independently, one call per element.
314 let expected: Vec<Ghash128b> = state
315 .iter()
316 .map(|x| InvertOrZero::invert_or_zero(*x))
317 .collect();
318
319 let mut inverter = BatchInversion::<Ghash128b>::new(n);
320 inverter.invert_nonzero(&mut state);
321
322 // The batched result must match the per-element reference exactly.
323 assert_eq!(state, expected);
324 }
325
326 proptest! {
327 #[test]
328 fn test_batch_inversion(n in 1usize..=16, n_zeros in 0usize..=16) {
329 // n_zeros counts positions to zero out of n, so it can never exceed n.
330 // Discard the proptest cases where the generator picked past that.
331 prop_assume!(n_zeros <= n);
332 let mut rng = StdRng::seed_from_u64(0);
333 test_batch_inversion_for_size(n, n_zeros, &mut rng);
334 }
335
336 #[test]
337 fn test_batch_inversion_nonzero(n in 1usize..=16) {
338 let mut rng = StdRng::seed_from_u64(0);
339 test_batch_inversion_nonzero_for_size(n, &mut rng);
340 }
341 }
342
343 #[test]
344 fn test_batch_inversion_reuse() {
345 let mut rng = StdRng::seed_from_u64(0);
346 // One context, reused across every zero count from 0 to 8 below.
347 // This checks that the zero mask from one call never leaks into the next.
348 let mut inverter = BatchInversion::<Ghash128b>::new(8);
349
350 for n_zeros in 0..=8 {
351 invert_with_inverter(&mut inverter, 8, n_zeros, &mut rng);
352 }
353 }
354
355 #[test]
356 fn test_batch_inversion_packed() {
357 use crate::test_utils::Packed128b;
358
359 let mut rng = StdRng::seed_from_u64(0);
360 const N: usize = 4;
361
362 // Packed128b packs 4 scalar lanes into one packed element.
363 // So 4 packed elements cover 16 scalars in total.
364 // Place zeros at 2 of those 16 positions: word 1's lane 0, and word 2's lane 2.
365 //
366 // word: 0 1 2 3
367 // lane: [_, _, _, _] [0, _, _, _] [_, _, 0, _] [_, _, _, _]
368 let mut state: Vec<Packed128b> = (0..N)
369 .map(|i| {
370 Packed128b::from_fn(|lane| {
371 if (i == 1 && lane == 0) || (i == 2 && lane == 2) {
372 Ghash128b::ZERO
373 } else {
374 Ghash128b::random(&mut rng)
375 }
376 })
377 })
378 .collect();
379
380 // Reference result: invert every scalar lane independently.
381 let expected: Vec<Packed128b> = state
382 .iter()
383 .map(|packed| Packed128b::from_scalars(packed.iter().map(InvertOrZero::invert_or_zero)))
384 .collect();
385
386 // Result under test: invert the whole batch of packed elements at once.
387 let mut inverter = BatchInversion::<Packed128b>::new(N);
388 inverter.invert_or_zero(&mut state);
389
390 assert_eq!(state, expected);
391 }
392}