1- use ark_ff:: Field ;
1+ use ark_ff:: { AdditiveGroup , Field } ;
22
33use crate :: algebra:: embedding:: { Embedding , Identity } ;
44
@@ -8,75 +8,146 @@ pub fn multilinear_extend<F: Field>(evals: &[F], point: &[F]) -> F {
88}
99
1010/// Evaluate the multi-linear extension of `evals` in `point`.
11+ ///
12+ /// Supports implicit zero-padding: when `evals.len() < 1 << point.len()`,
13+ /// the missing tail entries are treated as zeros.
14+ #[ allow( clippy:: too_many_lines) ]
1115pub fn mixed_multilinear_extend < M : Embedding > (
1216 embedding : & M ,
1317 evals : & [ M :: Source ] ,
1418 point : & [ M :: Target ] ,
1519) -> M :: Target {
16- assert_eq ! ( evals. len( ) , 1 << point. len( ) ) ;
17-
18- // Helper to compute (a + (b - a) * c) efficiently with a, b in source field.
19- let mixed = |a, b, c| embedding. mixed_add ( embedding. mixed_mul ( c, b - a) , a) ;
20-
21- match point {
22- [ ] => embedding. map ( evals[ 0 ] ) ,
23- [ x] => mixed ( evals[ 0 ] , evals[ 1 ] , * x) ,
24- [ x0, x1] => {
25- let a0 = mixed ( evals[ 0 ] , evals[ 1 ] , * x1) ;
26- let a1 = mixed ( evals[ 2 ] , evals[ 3 ] , * x1) ;
27- a0 + ( a1 - a0) * * x0
20+ #[ inline]
21+ fn eval_exact < M : Embedding > (
22+ embedding : & M ,
23+ evals : & [ M :: Source ] ,
24+ point : & [ M :: Target ] ,
25+ ) -> M :: Target {
26+ debug_assert_eq ! ( evals. len( ) , 1 << point. len( ) ) ;
27+
28+ // Helper to compute (a + (b - a) * c) efficiently with a, b in source field.
29+ let mixed = |a, b, c| embedding. mixed_add ( embedding. mixed_mul ( c, b - a) , a) ;
30+
31+ match point {
32+ [ ] => embedding. map ( evals[ 0 ] ) ,
33+ [ x] => mixed ( evals[ 0 ] , evals[ 1 ] , * x) ,
34+ [ x0, x1] => {
35+ let a0 = mixed ( evals[ 0 ] , evals[ 1 ] , * x1) ;
36+ let a1 = mixed ( evals[ 2 ] , evals[ 3 ] , * x1) ;
37+ a0 + ( a1 - a0) * * x0
38+ }
39+ [ x0, x1, x2] => {
40+ let a00 = mixed ( evals[ 0 ] , evals[ 1 ] , * x2) ;
41+ let a01 = mixed ( evals[ 2 ] , evals[ 3 ] , * x2) ;
42+ let a10 = mixed ( evals[ 4 ] , evals[ 5 ] , * x2) ;
43+ let a11 = mixed ( evals[ 6 ] , evals[ 7 ] , * x2) ;
44+ let a0 = a00 + ( a01 - a00) * * x1;
45+ let a1 = a10 + ( a11 - a10) * * x1;
46+ a0 + ( a1 - a0) * * x0
47+ }
48+ [ x0, x1, x2, x3] => {
49+ let a000 = mixed ( evals[ 0 ] , evals[ 1 ] , * x3) ;
50+ let a001 = mixed ( evals[ 2 ] , evals[ 3 ] , * x3) ;
51+ let a010 = mixed ( evals[ 4 ] , evals[ 5 ] , * x3) ;
52+ let a011 = mixed ( evals[ 6 ] , evals[ 7 ] , * x3) ;
53+ let a100 = mixed ( evals[ 8 ] , evals[ 9 ] , * x3) ;
54+ let a101 = mixed ( evals[ 10 ] , evals[ 11 ] , * x3) ;
55+ let a110 = mixed ( evals[ 12 ] , evals[ 13 ] , * x3) ;
56+ let a111 = mixed ( evals[ 14 ] , evals[ 15 ] , * x3) ;
57+ let a00 = a000 + ( a001 - a000) * * x2;
58+ let a01 = a010 + ( a011 - a010) * * x2;
59+ let a10 = a100 + ( a101 - a100) * * x2;
60+ let a11 = a110 + ( a111 - a110) * * x2;
61+ let a0 = a00 + ( a01 - a00) * * x1;
62+ let a1 = a10 + ( a11 - a10) * * x1;
63+ a0 + ( a1 - a0) * * x0
64+ }
65+ [ x, tail @ ..] => {
66+ let ( f0, f1) = evals. split_at ( evals. len ( ) / 2 ) ;
67+ #[ cfg( not( feature = "parallel" ) ) ]
68+ let ( f0, f1) = (
69+ eval_exact ( embedding, f0, tail) ,
70+ eval_exact ( embedding, f1, tail) ,
71+ ) ;
72+
73+ #[ cfg( feature = "parallel" ) ]
74+ let ( f0, f1) = {
75+ use crate :: utils:: workload_size;
76+ if evals. len ( ) > workload_size :: < M :: Source > ( ) {
77+ rayon:: join (
78+ || eval_exact ( embedding, f0, tail) ,
79+ || eval_exact ( embedding, f1, tail) ,
80+ )
81+ } else {
82+ (
83+ eval_exact ( embedding, f0, tail) ,
84+ eval_exact ( embedding, f1, tail) ,
85+ )
86+ }
87+ } ;
88+
89+ f0 + ( f1 - f0) * * x
90+ }
2891 }
29- [ x0, x1, x2] => {
30- let a00 = mixed ( evals[ 0 ] , evals[ 1 ] , * x2) ;
31- let a01 = mixed ( evals[ 2 ] , evals[ 3 ] , * x2) ;
32- let a10 = mixed ( evals[ 4 ] , evals[ 5 ] , * x2) ;
33- let a11 = mixed ( evals[ 6 ] , evals[ 7 ] , * x2) ;
34- let a0 = a00 + ( a01 - a00) * * x1;
35- let a1 = a10 + ( a11 - a10) * * x1;
36- a0 + ( a1 - a0) * * x0
92+ }
93+
94+ #[ inline]
95+ fn eval_partial < M : Embedding > (
96+ embedding : & M ,
97+ evals : & [ M :: Source ] ,
98+ point : & [ M :: Target ] ,
99+ ) -> M :: Target {
100+ let size = 1 << point. len ( ) ;
101+ debug_assert ! ( evals. len( ) <= size) ;
102+ if evals. is_empty ( ) {
103+ return M :: Target :: ZERO ;
37104 }
38- [ x0, x1, x2, x3] => {
39- let a000 = mixed ( evals[ 0 ] , evals[ 1 ] , * x3) ;
40- let a001 = mixed ( evals[ 2 ] , evals[ 3 ] , * x3) ;
41- let a010 = mixed ( evals[ 4 ] , evals[ 5 ] , * x3) ;
42- let a011 = mixed ( evals[ 6 ] , evals[ 7 ] , * x3) ;
43- let a100 = mixed ( evals[ 8 ] , evals[ 9 ] , * x3) ;
44- let a101 = mixed ( evals[ 10 ] , evals[ 11 ] , * x3) ;
45- let a110 = mixed ( evals[ 12 ] , evals[ 13 ] , * x3) ;
46- let a111 = mixed ( evals[ 14 ] , evals[ 15 ] , * x3) ;
47- let a00 = a000 + ( a001 - a000) * * x2;
48- let a01 = a010 + ( a011 - a010) * * x2;
49- let a10 = a100 + ( a101 - a100) * * x2;
50- let a11 = a110 + ( a111 - a110) * * x2;
51- let a0 = a00 + ( a01 - a00) * * x1;
52- let a1 = a10 + ( a11 - a10) * * x1;
53- a0 + ( a1 - a0) * * x0
105+ if evals. len ( ) == size {
106+ return eval_exact ( embedding, evals, point) ;
54107 }
55- [ x, tail @ ..] => {
56- let ( f0, f1) = evals. split_at ( evals. len ( ) / 2 ) ;
57- #[ cfg( not( feature = "parallel" ) ) ]
58- let ( f0, f1) = (
59- mixed_multilinear_extend ( embedding, f0, tail) ,
60- mixed_multilinear_extend ( embedding, f1, tail) ,
61- ) ;
62- #[ cfg( feature = "parallel" ) ]
63- let ( f0, f1) = {
64- use crate :: utils:: workload_size;
65- if evals. len ( ) > workload_size :: < M :: Source > ( ) {
66- rayon:: join (
67- || mixed_multilinear_extend ( embedding, f0, tail) ,
68- || mixed_multilinear_extend ( embedding, f1, tail) ,
69- )
70- } else {
71- (
72- mixed_multilinear_extend ( embedding, f0, tail) ,
73- mixed_multilinear_extend ( embedding, f1, tail) ,
74- )
108+
109+ match point {
110+ [ ] => embedding. map ( evals[ 0 ] ) ,
111+ [ x, tail @ ..] => {
112+ let half = size / 2 ;
113+
114+ // Only low half has data; high half is all implicit zeros.
115+ if evals. len ( ) <= half {
116+ let f0 = eval_partial ( embedding, evals, tail) ;
117+ return f0 * ( M :: Target :: ONE - * x) ;
75118 }
76- } ;
77- f0 + ( f1 - f0) * * x
119+
120+ // Low subtree is exact/full, high subtree is partial.
121+ let ( low, high) = evals. split_at ( half) ;
122+
123+ #[ cfg( not( feature = "parallel" ) ) ]
124+ let ( f0, f1) = (
125+ eval_exact ( embedding, low, tail) ,
126+ eval_partial ( embedding, high, tail) ,
127+ ) ;
128+
129+ #[ cfg( feature = "parallel" ) ]
130+ let ( f0, f1) = {
131+ use crate :: utils:: workload_size;
132+ if evals. len ( ) > workload_size :: < M :: Source > ( ) {
133+ rayon:: join (
134+ || eval_exact ( embedding, low, tail) ,
135+ || eval_partial ( embedding, high, tail) ,
136+ )
137+ } else {
138+ (
139+ eval_exact ( embedding, low, tail) ,
140+ eval_partial ( embedding, high, tail) ,
141+ )
142+ }
143+ } ;
144+
145+ f0 + ( f1 - f0) * * x
146+ }
78147 }
79148 }
149+
150+ eval_partial ( embedding, evals, point)
80151}
81152
82153/// Accumulates a scaled evaluation of the equality function.
@@ -88,7 +159,7 @@ pub fn mixed_multilinear_extend<M: Embedding>(
88159/// eq(X) = ∏ (1 - X_i + 2X_i z_i)
89160/// ```
90161///
91- /// where `z_i` are the points.
162+ /// where `z_i` are the points.
92163pub fn eval_eq < F : Field > ( accumulator : & mut [ F ] , point : & [ F ] , scalar : F ) {
93164 assert_eq ! ( accumulator. len( ) , 1 << point. len( ) ) ;
94165 if let [ x0, xs @ ..] = point {
@@ -110,3 +181,45 @@ pub fn eval_eq<F: Field>(accumulator: &mut [F], point: &[F], scalar: F) {
110181 accumulator[ 0 ] += scalar;
111182 }
112183}
184+
185+ #[ cfg( test) ]
186+ mod tests {
187+ use ark_std:: rand:: { rngs:: StdRng , SeedableRng } ;
188+ use proptest:: proptest;
189+
190+ use super :: * ;
191+ use crate :: algebra:: { random_vector, sumcheck:: tests:: zero_pad} ;
192+
193+ pub type F = crate :: algebra:: fields:: Field64 ;
194+
195+ #[ test]
196+ fn test_multilinear_zero_extend ( ) {
197+ proptest ! ( |( seed: u64 , length in 0_usize ..( 1 << 14 ) ) | {
198+ let mut rng = StdRng :: seed_from_u64( seed) ;
199+ let vector: Vec <F > = random_vector( & mut rng, length) ;
200+ let extended_vector = zero_pad( & vector) ;
201+ let num_variables = length. next_power_of_two( ) . trailing_zeros( ) ;
202+ let point = random_vector( & mut rng, num_variables as usize ) ;
203+ assert_eq!(
204+ multilinear_extend( & vector, & point) ,
205+ multilinear_extend( & extended_vector, & point)
206+ ) ;
207+ } ) ;
208+ }
209+
210+ #[ test]
211+ fn test_multilinear_extra_variables ( ) {
212+ proptest ! ( |( seed: u64 , length in 0_usize ..( 1 << 10 ) , excess_variables in 0_usize ..3 ) | {
213+ let mut rng = StdRng :: seed_from_u64( seed) ;
214+ let vector: Vec <F > = random_vector( & mut rng, length) ;
215+ let num_variables = length. next_power_of_two( ) . trailing_zeros( ) as usize + excess_variables;
216+ let point = random_vector( & mut rng, num_variables) ;
217+ let mut extended_vector = vector. clone( ) ;
218+ extended_vector. resize( 1 << num_variables, F :: ZERO ) ;
219+ assert_eq!(
220+ multilinear_extend( & vector, & point) ,
221+ multilinear_extend( & extended_vector, & point)
222+ ) ;
223+ } ) ;
224+ }
225+ }
0 commit comments