1use core::ops::{Add, Sub};
2use std::{iter::zip, ops::Mul};
3
4use itertools::Itertools;
5use p3_field::{ExtensionField, Field, PrimeCharacteristicRing, TwoAdicField};
6
7pub fn eval_eq_mle<F1, F2, F3>(x: &[F1], y: &[F2]) -> F3
8where
9 F1: Field,
10 F2: Field,
11 F3: Field,
12 F1: Mul<F2, Output = F3>,
13 F3: Sub<F1, Output = F3>,
14 F3: Sub<F2, Output = F3>,
15{
16 debug_assert_eq!(x.len(), y.len());
17 zip(x, y).fold(F3::ONE, |acc, (&x_i, &y_i)| {
18 acc * (F3::ONE - y_i - x_i + (x_i * y_i).double())
19 })
20}
21
22pub fn eval_mobius_eq_mle<F: Field>(u: &[F], x: &[F]) -> F {
28 debug_assert_eq!(u.len(), x.len());
29 zip(u, x).fold(F::ONE, |acc, (&u_i, &x_i)| {
30 let w0 = F::ONE - u_i.double();
31 acc * (w0 * (F::ONE - x_i) + u_i * x_i)
32 })
33}
34
35pub fn eval_mle_evals_at_point<F: Field>(evals: &mut [F], x: &[F]) -> F {
40 debug_assert_eq!(evals.len(), 1 << x.len());
41 let mut len = evals.len();
42 for &xj in x.iter().rev() {
43 len >>= 1;
44 let (lo, hi) = evals.split_at_mut(len);
45 for i in 0..len {
46 lo[i] = lo[i] * (F::ONE - xj) + hi[i] * xj;
47 }
48 }
49 evals[0]
50}
51
52pub fn eval_eq_uni<F: Field>(l_skip: usize, x: F, y: F) -> F {
58 let mut res = F::ONE;
59 for (x_pow, y_pow) in zip(x.exp_powers_of_2(), y.exp_powers_of_2()).take(l_skip) {
60 res = (x_pow + y_pow) * res + (x_pow - F::ONE) * (y_pow - F::ONE);
61 }
62 res * F::ONE.halve().exp_u64(l_skip as u64)
63}
64
65pub fn eval_eq_uni_at_one<F: Field>(l_skip: usize, x: F) -> F {
69 let mut res = F::ONE;
70 for x_pow in x.exp_powers_of_2().take(l_skip) {
71 res *= x_pow + F::ONE;
72 }
73 res * F::ONE.halve().exp_u64(l_skip as u64)
74}
75
76pub fn eq_uni_poly<F, EF>(l_skip: usize, x: EF) -> UnivariatePoly<EF>
83where
84 F: Field,
85 EF: ExtensionField<F>,
86{
87 let n_inv = F::ONE.halve().exp_u64(l_skip as u64);
88 let mut coeffs = x
89 .powers()
90 .skip(1)
91 .take(1 << l_skip)
92 .map(|x_pow| x_pow * n_inv)
93 .collect_vec();
94 coeffs.reverse();
95 coeffs[0] = n_inv.into();
96 UnivariatePoly::new(coeffs)
97}
98
99pub fn eval_in_uni<F: Field>(l_skip: usize, n: isize, z: F) -> F {
100 debug_assert!(n >= -(l_skip as isize));
101 if n.is_negative() {
102 eval_eq_uni_at_one(
103 n.unsigned_abs(),
104 z.exp_power_of_2(l_skip.wrapping_add_signed(n)),
105 )
106 } else {
107 F::ONE
108 }
109}
110
111pub fn eval_eq_prism<F: Field>(l_skip: usize, x: &[F], y: &[F]) -> F {
112 eval_eq_uni(l_skip, x[0], y[0]) * eval_eq_mle(&x[1..], &y[1..])
113}
114
115pub fn evals_eq_hypercube_serial<F: Field>(x: &[F]) -> Vec<F> {
116 let n = x.len();
117 let mut out = F::zero_vec(1 << n);
118 out[0] = F::ONE;
119 for (i, &x_i) in x.iter().enumerate() {
120 let (los, his) = out[..2 << i].split_at_mut(1 << i);
121 for (lo, hi) in los.iter_mut().zip(his.iter_mut()) {
122 *hi = *lo * x_i;
123 *lo *= F::ONE - x_i;
124 }
125 }
126 out
127}
128
129pub fn eval_eq_sharp_uni<F, EF>(omega_skip_pows: &[F], xi_1: &[EF], z: EF) -> EF
131where
132 F: Field,
133 EF: ExtensionField<F>,
134{
135 let l_skip = xi_1.len();
136 debug_assert_eq!(omega_skip_pows.len(), 1 << l_skip);
137
138 let mut res = EF::ZERO;
139 let eq_xi_evals = evals_eq_hypercube_serial(xi_1);
140 for (&omega_pow, eq_xi_eval) in omega_skip_pows.iter().zip(eq_xi_evals) {
141 res += eval_eq_uni(l_skip, z, omega_pow.into()) * eq_xi_eval;
142 }
143 #[cfg(debug_assertions)]
144 {
145 let coeffs = (0..(1 << l_skip))
146 .map(|k| {
147 let mut c = EF::ONE;
148 #[allow(clippy::needless_range_loop)]
149 for i in 0..l_skip {
150 let idx = (k << i) % (1 << l_skip);
151 c *= EF::ONE - xi_1[i]
152 + xi_1[i] * omega_skip_pows[((1 << l_skip) - idx) % (1 << l_skip)];
153 }
154 c
155 })
156 .collect_vec();
157 let mut rpow = EF::ONE;
158 let mut other = EF::ZERO;
159 for c in coeffs {
160 other += rpow * c;
161 rpow *= z;
162 }
163 other *= EF::TWO.inverse().exp_u64(l_skip as u64);
164 debug_assert_eq!(other, res);
165 }
166 res
167}
168
169pub fn eval_rot_kernel_prism<F: TwoAdicField>(l_skip: usize, x: &[F], y: &[F]) -> F {
173 let omega = F::two_adic_generator(l_skip);
174
175 let (eq_cube, rot_cube) = eval_eq_rot_cube(&x[1..], &y[1..]);
176 eval_eq_uni(l_skip, x[0], y[0] * omega) * eq_cube
179 + eval_eq_uni_at_one(l_skip, x[0])
180 * eval_eq_uni_at_one(l_skip, y[0] * omega)
181 * (rot_cube - eq_cube)
182}
183
184pub fn eval_eq_rot_cube<F: Field>(x: &[F], y: &[F]) -> (F, F) {
186 let n = x.len();
187 debug_assert_eq!(n, y.len());
188 let mut rot = F::ONE;
191 let mut eq = F::ONE;
192 for i in (0..n).rev() {
193 rot = x[i] * (F::ONE - y[i]) * eq + (F::ONE - x[i]) * y[i] * rot;
194 eq *= x[i] * y[i] + (F::ONE - x[i]) * (F::ONE - y[i]);
195 }
196 (eq, rot)
197}
198
199#[derive(Clone, Debug)]
202pub struct UnivariatePoly<F>(pub(crate) Vec<F>);
203
204impl<F> UnivariatePoly<F> {
205 pub fn new(coeffs: Vec<F>) -> Self {
206 Self(coeffs)
207 }
208
209 pub fn coeffs(&self) -> &[F] {
210 &self.0
211 }
212
213 pub fn coeffs_mut(&mut self) -> &mut Vec<F> {
214 &mut self.0
215 }
216
217 pub fn into_coeffs(self) -> Vec<F> {
218 self.0
219 }
220}
221
222impl<F: Field> UnivariatePoly<F> {
223 pub fn eval_at_point<EF: ExtensionField<F>>(&self, x: EF) -> EF {
224 horner_eval(&self.0, x)
225 }
226}
227
228pub fn horner_eval<F1, F2, F3>(coeffs: &[F1], x: F2) -> F3
232where
233 F1: Field,
234 F2: Field,
235 F3: Field + Add<F1, Output = F3>,
236 F3: Mul<F2, Output = F3>,
237{
238 coeffs.iter().rfold(F3::ZERO, |acc, coeff| acc * x + *coeff)
239}
240
241#[inline(always)]
244pub fn interpolate_linear_at_01<F: Field>(evals: &[F; 2], x: F) -> F {
245 let p = evals[1] - evals[0];
246 p * x + evals[0]
247}
248
249#[inline(always)]
252pub fn interpolate_quadratic_at_012<F: Field>(evals: &[F; 3], x: F) -> F {
253 let s1 = evals[1] - evals[0];
254 let s2 = evals[2] - evals[1];
255 let p = (s2 - s1).halve();
256 let q = s1 - p;
257 (p * x + q) * x + evals[0]
258}
259
260#[inline(always)]
263pub fn interpolate_cubic_at_0123<F: Field>(evals: &[F; 4], x: F) -> F {
264 let inv6 = F::from_u64(6).inverse();
265
266 let s1 = evals[1] - evals[0];
267 let s2 = evals[2] - evals[0];
268 let s3 = evals[3] - evals[0];
269
270 let d3 = s3 - (s2 - s1) * F::from_u64(3);
271
272 let p = d3 * inv6;
273 let q = (s2 - d3).halve() - s1;
274 let r = s1 - p - q;
275
276 ((p * x + q) * x + r) * x + evals[0]
277}
278
279pub struct ExpPowers2<T> {
280 current: Option<T>,
281}
282
283impl<T: Squarable + PrimeCharacteristicRing> Iterator for ExpPowers2<T> {
284 type Item = T;
285 fn next(&mut self) -> Option<Self::Item> {
286 if let Some(curr) = self.current.take() {
287 let next = curr.square();
288 self.current = Some(next);
289 Some(curr)
290 } else {
291 None
292 }
293 }
294}
295
296pub trait Squarable: PrimeCharacteristicRing + Clone {
297 #[inline]
298 fn exp_powers_of_2(&self) -> ExpPowers2<Self> {
299 ExpPowers2 {
300 current: Some(self.clone()),
301 }
302 }
303}
304
305impl<T: PrimeCharacteristicRing + Clone> Squarable for T {}
306
307#[cfg(test)]
308mod tests {
309 use itertools::Itertools;
310 use openvm_stark_sdk::config::baby_bear_poseidon2::*;
311 use p3_field::PrimeCharacteristicRing;
312
313 use super::*;
314
315 #[test]
316 fn test_interpolate_linear() {
317 let evals = [
318 EF::from_u64(20), EF::from_u64(10), ];
321
322 assert_eq!(interpolate_linear_at_01(&evals, EF::ZERO), evals[0]);
324 assert_eq!(interpolate_linear_at_01(&evals, EF::ONE), evals[1]);
325 }
326
327 #[test]
328 fn test_interpolate_quadratic() {
329 let evals = [
330 EF::from_u64(20), EF::from_u64(10), EF::from_u64(18), ];
334
335 assert_eq!(interpolate_quadratic_at_012(&evals, EF::ZERO), evals[0]);
337 assert_eq!(interpolate_quadratic_at_012(&evals, EF::ONE), evals[1]);
338 assert_eq!(
339 interpolate_quadratic_at_012(&evals, EF::from_u64(2)),
340 evals[2]
341 );
342 }
343
344 #[test]
345 fn test_interpolate_cubic() {
346 let evals = [
347 EF::from_u64(20), EF::from_u64(10), EF::from_u64(18), EF::from_u64(28), ];
352
353 assert_eq!(interpolate_cubic_at_0123(&evals, EF::ZERO), evals[0]);
355 assert_eq!(interpolate_cubic_at_0123(&evals, EF::ONE), evals[1]);
356 assert_eq!(interpolate_cubic_at_0123(&evals, EF::from_u64(2)), evals[2]);
357 assert_eq!(interpolate_cubic_at_0123(&evals, EF::from_u64(3)), evals[3]);
358 }
359
360 #[test]
361 fn test_exp_powers_of_2() {
362 let x = F::from_u32(3);
363 let s = x.exp_powers_of_2().take(3).collect_vec();
364 assert_eq!(s, vec![x, x * x, x * x * x * x],);
365 }
366
367 #[test]
368 fn test_eval_in_uni() {
369 let l = 3;
370 let n = -2;
371 let u_0 = F::from_u32(12345);
372 let ind = eval_in_uni(l, n, u_0);
373 let expected = (u_0.exp_power_of_2(l) - F::ONE)
374 * (u_0.exp_power_of_2(l.wrapping_add_signed(n)) - F::ONE).inverse()
375 * F::from_usize(1 << n.unsigned_abs()).inverse();
376 assert_eq!(ind, expected);
377 }
378}