Skip to main content

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}