Skip to main content

zinc_piop/
shift_predicate.rs

1//! Shift predicate evaluation.
2//!
3//! Evaluates `S_c(x, y)` — the multilinear extension of the shift-by-c
4//! indicator — at arbitrary field points.
5
6use crypto_primitives::SemiringConfig;
7use zinc_poly::utils::next_mle_eval;
8
9/// Evaluate the shift predicate `S_c(x, y)` at arbitrary field points.
10///
11/// Uses the high/low decomposition:
12///   `S_c(x, y) = L_0(x_lo, y_lo) · eq(x_hi, y_hi)
13///              + L_1(x_lo, y_lo) · next_mle(x_hi, y_hi)`
14///
15/// where `k = ceil(log2(2c))` determines the split point.
16///
17/// Cost: O(m + c · log c) field operations.
18#[allow(clippy::arithmetic_side_effects)]
19pub fn eval_shift_predicate<C: SemiringConfig>(
20    cfg: &C,
21    x: &[C::Element],
22    y: &[C::Element],
23    c: usize,
24) -> C::Element {
25    let m = x.len();
26    assert_eq!(y.len(), m);
27
28    // S_0(x, y) = eq(x, y): identity shift.
29    if c == 0 {
30        return eval_eq_poly(cfg, x, y);
31    }
32
33    // S_1(x, y) = next_mle(x, y): the successor predicate is exactly shift-by-1.
34    if c == 1 {
35        return next_mle_eval(cfg, x, y);
36    }
37
38    assert!(c < (1usize << m), "shift c must satisfy c < 2^m");
39    // k = ceil(log2(2*c))
40    let k = (2 * c).next_power_of_two().trailing_zeros() as usize;
41    if k >= m {
42        return eval_shift_small(cfg, x, y, c, m);
43    }
44
45    // LE convention: x[0..k] are the low bits, x[k..] are the high bits.
46    let (x_lo, x_hi) = x.split_at(k);
47    let (y_lo, y_hi) = y.split_at(k);
48
49    let l0 = eval_l0(cfg, x_lo, y_lo, c, k);
50    let l1 = eval_l1(cfg, x_lo, y_lo, c, k);
51    let eq = eval_eq_poly(cfg, x_hi, y_hi);
52    let next = next_mle_eval(cfg, x_hi, y_hi);
53
54    cfg.add(&cfg.mul(&l0, &eq), &cfg.mul(&l1, &next))
55}
56
57/// `eq(u, v) = prod_i (u_i * v_i + (1 - u_i)(1 - v_i))`
58///
59/// Evaluates the Multilinear polynomial for eq polynomial
60pub(crate) fn eval_eq_poly<C: SemiringConfig>(
61    cfg: &C,
62    u: &[C::Element],
63    v: &[C::Element],
64) -> C::Element {
65    let one = cfg.one();
66    u.iter()
67        .zip(v.iter())
68        .map(|(u_i, v_i)| {
69            cfg.add(
70                &cfg.mul(u_i, v_i),
71                &cfg.mul(&cfg.sub(&one, u_i), &cfg.sub(&one, v_i)),
72            )
73        })
74        .fold(one.clone(), |acc, term| cfg.mul(&acc, &term))
75}
76
77/// `delta_{bin_k(a)}(u) = eq(u, bin_k(a))`.
78///
79/// Evaluates the Lagrange basis polynomial for the binary encoding of `a`
80/// with `k` bits at the point `u`.
81///
82/// LE convention: `u[i]` corresponds to bit `i` (LSB = index 0).
83pub(crate) fn eval_delta<C: SemiringConfig>(
84    cfg: &C,
85    u: &[C::Element],
86    a: usize,
87    k: usize,
88) -> C::Element {
89    let one = cfg.one();
90    let mut result = one.clone();
91    for (i, u) in u.iter().take(k).enumerate() {
92        let bit = (a >> i) & 1;
93        if bit == 1 {
94            cfg.mul_assign(&mut result, u);
95        } else {
96            let term = cfg.sub(&one, u);
97            cfg.mul_assign(&mut result, &term);
98        }
99    }
100    result
101}
102
103/// `L_0^{(c)}(x_lo, y_lo)` — no-carry component.
104///
105/// `sum_{a=0}^{2^k - 1 - c} delta(x_lo, a) * delta(y_lo, a + c)`
106///
107/// On Booleans: 1 iff `Val(y_lo) = Val(x_lo) + c` with no carry into the high
108/// block.
109#[allow(clippy::arithmetic_side_effects)]
110pub(crate) fn eval_l0<C: SemiringConfig>(
111    cfg: &C,
112    x_lo: &[C::Element],
113    y_lo: &[C::Element],
114    c: usize,
115    k: usize,
116) -> C::Element {
117    let upper = (1 << k) - c;
118    (0..upper).fold(cfg.zero(), |acc, a| {
119        cfg.add(
120            &acc,
121            &cfg.mul(
122                &eval_delta(cfg, x_lo, a, k),
123                &eval_delta(cfg, y_lo, a + c, k),
124            ),
125        )
126    })
127}
128
129/// `L_1^{(c)}(x_lo, y_lo)` — carry component.
130///
131/// `sum_{a=2^k-c}^{2^k-1} delta(x_lo, a) * delta(y_lo, a + c - 2^k)`
132///
133/// On Booleans: 1 iff the addition carries into the high block.
134#[allow(clippy::arithmetic_side_effects)]
135pub(crate) fn eval_l1<C: SemiringConfig>(
136    cfg: &C,
137    x_lo: &[C::Element],
138    y_lo: &[C::Element],
139    c: usize,
140    k: usize,
141) -> C::Element {
142    let two_k = 1 << k;
143    ((two_k - c)..two_k).fold(cfg.zero(), |acc, a| {
144        cfg.add(
145            &acc,
146            &cfg.mul(
147                &eval_delta(cfg, x_lo, a, k),
148                &eval_delta(cfg, y_lo, a + c - two_k, k),
149            ),
150        )
151    })
152}
153
154/// Special case when `k >= m`: no high block, direct evaluation.
155///
156/// `sum_{a=0}^{n-1-c} delta(x, a, m) * delta(y, a+c, m)`
157#[allow(clippy::arithmetic_side_effects)]
158fn eval_shift_small<C: SemiringConfig>(
159    cfg: &C,
160    x: &[C::Element],
161    y: &[C::Element],
162    c: usize,
163    m: usize,
164) -> C::Element {
165    let upper = (1 << m) - c;
166    (0..upper).fold(cfg.zero(), |acc, a| {
167        cfg.add(
168            &acc,
169            &cfg.mul(&eval_delta(cfg, x, a, m), &eval_delta(cfg, y, a + c, m)),
170        )
171    })
172}
173
174#[cfg(test)]
175mod tests {
176    use super::*;
177    use crate::test_utils::test_config;
178    use crypto_primitives::{
179        ProjectElementWithConfig,
180        crypto_bigint_monty::{MontyField, MontyFieldElement},
181    };
182    use rand::prelude::*;
183    use zinc_poly::utils::{build_eq_x_r, build_next_c_r_mle};
184
185    type F = MontyField<4>;
186    type E = MontyFieldElement<4>;
187
188    /// LE convention: to_bin(val, i) = bit i of val (LSB = index 0).
189    fn to_bin(cfg: &F, val: usize, bit: usize) -> E {
190        if (val >> bit) & 1 == 1 {
191            cfg.one()
192        } else {
193            cfg.zero()
194        }
195    }
196
197    fn rand_field(cfg: &F, rng: &mut impl Rng) -> E {
198        cfg.project(&rng.random::<u32>())
199    }
200
201    /// Check S_c on Boolean inputs: S_c(bin(a), bin(a+c)) = 1,
202    /// and S_c(bin(a), bin(b)) = 0 for b != a+c.
203    #[test]
204    fn test_shift_predicate_boolean() {
205        let cfg = test_config();
206        let m = 4;
207        let n = 1usize << m;
208
209        for c in [1, 2, 5] {
210            for a in 0..n {
211                for b in 0..n {
212                    let x: Vec<E> = (0..m).map(|i| to_bin(&cfg, a, i)).collect();
213                    let y: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
214                    let val = eval_shift_predicate(&cfg, &x, &y, c);
215
216                    if b == a + c && a + c < n {
217                        assert_eq!(val, cfg.one(), "S_c({a},{b}) should be 1 for c={c}");
218                    } else {
219                        assert_eq!(val, cfg.zero(), "S_c({a},{b}) should be 0 for c={c}");
220                    }
221                }
222            }
223        }
224    }
225
226    /// Verify the next_mle on all Boolean inputs.
227    #[test]
228    fn test_next_boolean() {
229        let cfg = test_config();
230        let m = 4;
231        let n = 1usize << m;
232        for a in 0..n {
233            for b in 0..n {
234                let u: Vec<E> = (0..m).map(|i| to_bin(&cfg, a, i)).collect();
235                let v: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
236                let val = next_mle_eval(&cfg, &u, &v);
237
238                if b == a + 1 && a + 1 < n {
239                    assert_eq!(val, cfg.one(), "Next({a},{b}) should be 1");
240                } else {
241                    assert_eq!(val, cfg.zero(), "Next({a},{b}) should be 0");
242                }
243            }
244        }
245    }
246
247    /// Check verifier (`eval_shift_predicate`) against prover
248    /// (`build_next_c_r_mle`) at Boolean points:
249    ///   eval_shift_predicate(r, bin(b), c) == build_next_c_r_mle(r, c)[b]
250    #[test]
251    fn test_shift_predicate_vs_prover_mle() {
252        let cfg = test_config();
253        let mut rng = rand::rng();
254        let m = 4;
255        let n = 1usize << m;
256        let c = 3;
257
258        let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
259        let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
260
261        for b in 0..n {
262            let b_bin: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
263            let val = eval_shift_predicate(&cfg, &r, &b_bin, c);
264            assert_eq!(
265                val, next_c.evaluations[b],
266                "S_{c}(r, bin({b})) mismatch with prover MLE"
267            );
268        }
269    }
270
271    /// Check at random field points via MLE summation:
272    ///   eval_shift_predicate(r, y, c) == sum_b build_next_c_r_mle(r, c)[b] *
273    /// eq(b, y)
274    #[test]
275    fn test_shift_predicate_random_points() {
276        let cfg = test_config();
277        let mut rng = rand::rng();
278        let m = 4;
279        let c = 3;
280
281        for _ in 0..8 {
282            let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
283            let y: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
284
285            let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
286            let eq_y = build_eq_x_r(&cfg, &y).unwrap();
287            let rhs = next_c
288                .evaluations
289                .iter()
290                .zip(eq_y.evaluations.iter())
291                .fold(cfg.zero(), |acc, (ni, ei)| cfg.add(&acc, &cfg.mul(ni, ei)));
292            let lhs = eval_shift_predicate(&cfg, &r, &y, c);
293
294            assert_eq!(lhs, rhs, "random-point MLE mismatch");
295        }
296    }
297
298    /// Test c=0 (identity) and c=1 (successor) fast paths at random points,
299    /// and verify predicate vs prover MLE consistency across multiple shift
300    /// amounts.
301    #[test]
302    fn test_fast_paths_and_multi_c() {
303        let cfg = test_config();
304        let mut rng = rand::rng();
305        let m = 4;
306        let n = 1usize << m;
307
308        for c in [0, 1, 2, 5, 7] {
309            let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
310            let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
311
312            // Predicate vs prover MLE at Boolean y
313            for b in 0..n {
314                let b_bin: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
315                let val = eval_shift_predicate(&cfg, &r, &b_bin, c);
316                assert_eq!(
317                    val, next_c.evaluations[b],
318                    "S_{c}(r, bin({b})) mismatch with prover MLE"
319                );
320            }
321
322            // Predicate vs prover MLE at random y (MLE consistency)
323            for _ in 0..4 {
324                let y: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
325                let eq_y = build_eq_x_r(&cfg, &y).unwrap();
326                let rhs = next_c
327                    .evaluations
328                    .iter()
329                    .zip(eq_y.evaluations.iter())
330                    .fold(cfg.zero(), |acc, (ni, ei)| cfg.add(&acc, &cfg.mul(ni, ei)));
331                let lhs = eval_shift_predicate(&cfg, &r, &y, c);
332                assert_eq!(lhs, rhs, "random-point MLE mismatch for c={c}");
333            }
334        }
335    }
336
337    /// Boundary test: large c values where most rows shift beyond the domain.
338    #[test]
339    fn test_shift_predicate_boundary() {
340        let cfg = test_config();
341        let m = 3;
342        let n = 1usize << m; // 8
343
344        for c in [n / 2, n - 1] {
345            // Boolean correctness: S_c(bin(a), bin(b)) = 1 iff b == a+c < n
346            for a in 0..n {
347                for b in 0..n {
348                    let x: Vec<E> = (0..m).map(|i| to_bin(&cfg, a, i)).collect();
349                    let y: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
350                    let val = eval_shift_predicate(&cfg, &x, &y, c);
351
352                    if b == a + c && a + c < n {
353                        assert_eq!(val, cfg.one(), "S_{c}(bin({a}), bin({b})) should be 1");
354                    } else {
355                        assert_eq!(val, cfg.zero(), "S_{c}(bin({a}), bin({b})) should be 0");
356                    }
357                }
358            }
359
360            // Prover MLE: first c entries zero, rest match eq(r, b-c)
361            let mut rng = rand::rng();
362            let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
363            let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
364            let zero = cfg.zero();
365
366            // First c entries must be zero
367            for b in 0..c {
368                assert_eq!(
369                    next_c.evaluations[b], zero,
370                    "next_c[{b}] should be zero for c={c}"
371                );
372            }
373            // Remaining entries should be nonzero (with overwhelming probability)
374            let nonzero_count = next_c.evaluations[c..]
375                .iter()
376                .filter(|e| **e != zero)
377                .count();
378            assert_eq!(
379                nonzero_count,
380                n - c,
381                "expected {} nonzero entries for c={c}",
382                n - c
383            );
384        }
385    }
386
387    /// Check that build_next_c_r_mle correctly reproduces MLE[shift_c(v)](r)
388    /// via inner product: sum_b next_c(b) * v[b] == sum_b eq(r, b-c) * v[b].
389    #[test]
390    fn test_prover_mle_inner_product() {
391        let cfg = test_config();
392        let mut rng = rand::rng();
393        let m = 4;
394        let n = 1usize << m;
395
396        for c in [1, 2, 3, 7] {
397            let v: Vec<E> = (0..n).map(|_| rand_field(&cfg, &mut rng)).collect();
398            let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
399
400            // Ground truth: sum_{b>=c} eq(r, b-c) * v[b]
401            let eq_r = build_eq_x_r(&cfg, &r).unwrap();
402            let expected = (c..n).fold(cfg.zero(), |acc, b| {
403                cfg.add(&acc, &cfg.mul(&eq_r.evaluations[b - c], &v[b]))
404            });
405
406            // Via prover MLE: sum_b next_c[b] * v[b]
407            let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
408            let got = next_c
409                .evaluations
410                .iter()
411                .zip(v.iter())
412                .fold(cfg.zero(), |acc, (ni, vi)| cfg.add(&acc, &cfg.mul(vi, ni)));
413
414            assert_eq!(got, expected, "prover MLE inner product mismatch for c={c}");
415        }
416    }
417}