Skip to main content

openvm_stark_backend/
poly_common.rs

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
22/// Evaluate `mobius_eq_poly(u_tilde)` at an arbitrary point `x`.
23///
24/// ```text
25/// mobius_eq_poly(u_tilde)(x) = ∏_i ((1 - 2*u_tilde_i) * (1 - x_i) + u_tilde_i * x_i)
26/// ```
27pub 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
35/// Evaluate the MLE defined by its hypercube evaluations at an arbitrary point, in place.
36///
37/// `evals` has length `2^n` and contains `f(b)` for each `b ∈ {0,1}^n`.
38/// Returns `f(x)` where `x = x[0..n]`.
39pub 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
52/// Let D be the univariate skip domain, the subgroup of `F^*` of order `2^l_skip`.
53///
54/// Computes the polynomial ```text
55///     eq_D(X, Y) = \sum_{z_1 \in D} \prod_{z_2 \in D, z_2 != z_1} (X - z_1)(Y - z_2) / (z_1 -
56/// z_2)^2 ```
57pub 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
65/// Let D be the univariate skip domain, the subgroup of `F^*` of order `2^l_skip`.
66///
67/// Computes the polynomial eq_D(X, 1); see `eval_eq_uni`.
68pub 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
76/// Returns `eq_D(x, Z)` as a polynomial in `Z` in coefficient form.
77/// Derived from `eq_D(x, Z)` being the Lagrange basis at `x`, which is the character sum over the
78/// roots of unity.
79///
80/// If z in D, then `eq_D(x, z) = 1/N sum_{k=1}^N (x/z)^k = 1/N sum_{k=1}^N x^k
81/// z^{N-k}`.
82pub 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
129/// Length of `xi_1` should be `l_skip`.
130pub 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
169/// `\kappa_\rot(x, y)` should equal `\delta_{x,rot(y)}` on hyperprism.
170///
171/// `omega_pows` must have length `2^{l_skip}`.
172pub 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    // If not at boundary of D, just rotate in D, don't change cube coordinates. Otherwise at
177    // boundary, rotate the cube
178    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
184/// MLE of cyclic rotation kernel on hypercube
185pub 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    // Recursive formula: rot(x, y) = x[0] * (1 - y[0]) * eq(x[1..], y[1..]) + (1 - x[0]) y[0] *
189    // rot(x[1..], y[1..])
190    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// Source: https://github.com/starkware-libs/stwo/blob/dev/crates/stwo/src/prover/lookups/utils.rs#L12
200/// Univariate polynomial in coefficient form.
201#[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
228/// Evaluates univariate polynomial using [Horner's method].
229///
230/// [Horner's method]: https://en.wikipedia.org/wiki/Horner%27s_method
231pub 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/// Interpolates a linear polynomial through points (0, evals[0]), (1, evals[1])
242/// and evaluates it at x.
243#[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/// Interpolates a quadratic polynomial through points (0, evals[0]), (1, evals[1]),
250/// (2, evals[2])  and evaluates it at x.
251#[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/// Interpolates a cubic polynomial through points (0, evals[0]), (1, evals[1]),
261/// (2, evals[2]), (3, evals[3]) and evaluates it at x.
262#[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), // s(0)
319            EF::from_u64(10), // s(1)
320        ];
321
322        // Test interpolation at known points
323        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), // s(0)
331            EF::from_u64(10), // s(1)
332            EF::from_u64(18), // s(2)
333        ];
334
335        // Test interpolation at known points
336        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), // s(0)
348            EF::from_u64(10), // s(1)
349            EF::from_u64(18), // s(2)
350            EF::from_u64(28), // s(3)
351        ];
352
353        // Test interpolation at known points
354        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}