openvm_stark_backend/utils/
batch_inverse.rs1use p3_field::{Field, FieldArray, PackedValue, PrimeCharacteristicRing};
4use tracing::instrument;
5
6#[instrument(level = "debug", skip_all)]
17pub fn batch_multiplicative_inverse_serial<F: Field>(x: &[F]) -> Vec<F> {
18 let n = x.len();
19 let mut result = F::zero_vec(n);
20
21 batch_multiplicative_inverse_helper(x, &mut result);
22
23 result
24}
25
26fn batch_multiplicative_inverse_helper<F: Field>(x: &[F], result: &mut [F]) {
28 const WIDTH: usize = 4;
31
32 let n = x.len();
33 assert_eq!(result.len(), n);
34 if !n.is_multiple_of(WIDTH) {
35 return batch_multiplicative_inverse_general(x, result, |x| x.inverse());
39 }
40
41 let x_packed = FieldArray::<F, 4>::pack_slice(x);
42 let result_packed = FieldArray::<F, 4>::pack_slice_mut(result);
43
44 let inv = |x_packed: FieldArray<F, 4>| {
45 let mut result = FieldArray::<F, 4>::default();
46 batch_multiplicative_inverse_general(&x_packed.0, &mut result.0, |x| x.inverse());
47 result
48 };
49 batch_multiplicative_inverse_general(x_packed, result_packed, inv);
50}
51
52pub(crate) fn batch_multiplicative_inverse_general<F, Inv>(x: &[F], result: &mut [F], inv: Inv)
55where
56 F: PrimeCharacteristicRing + Copy,
57 Inv: Fn(F) -> F,
58{
59 let n = x.len();
60 assert_eq!(result.len(), n);
61 if n == 0 {
62 return;
63 }
64
65 result[0] = F::ONE;
66 for i in 1..n {
67 result[i] = result[i - 1] * x[i - 1];
68 }
69
70 let product = result[n - 1] * x[n - 1];
71 let mut inv = inv(product);
72
73 for i in (0..n).rev() {
74 result[i] *= inv;
75 inv *= x[i];
76 }
77}