binius_ip_prover/sumcheck/
factored_multilinear.rs1use binius_compute::BufferData;
6use binius_field::{Field, PackedField};
7use binius_math::{FieldBuffer, multilinear::fold::fold_highest_var_inplace};
8
9#[derive(Debug, Clone)]
53pub struct FactoredMultilinear<P: PackedField, Data: BufferData<P> = Vec<P>> {
54 factors: Vec<FieldBuffer<P, Data>>,
58
59 bound: P::Scalar,
64}
65
66impl<P: PackedField, Data: BufferData<P>> FactoredMultilinear<P, Data> {
67 pub fn new(factors: impl IntoIterator<Item = FieldBuffer<P, Data>>) -> Self {
72 let mut bound = P::Scalar::ONE;
73 let factors = factors
74 .into_iter()
75 .filter(|factor| {
76 if factor.log_len() == 0 {
77 bound *= factor.get(0);
78 false
79 } else {
80 true
81 }
82 })
83 .collect();
84 Self { factors, bound }
85 }
86
87 pub fn n_vars(&self) -> usize {
89 self.factors.iter().map(FieldBuffer::log_len).sum()
90 }
91
92 pub fn get(&self, index: usize) -> P::Scalar {
101 let n_vars = self.n_vars();
102 assert!(
103 n_vars >= usize::BITS as usize || index < 1 << n_vars,
104 "precondition: index {index} must fit {n_vars} variables"
105 );
106
107 let mut value = self.bound;
108 let mut rest = index;
109 for factor in &self.factors {
110 let mask = (1 << factor.log_len()) - 1;
112 value *= factor.get(rest & mask);
113 rest >>= factor.log_len();
114 }
115 value
116 }
117
118 pub fn fold_highest_var(&mut self, challenge: P::Scalar) {
127 let factor = self
128 .factors
129 .last_mut()
130 .expect("precondition: at least one variable must be free");
131 fold_highest_var_inplace(factor, challenge);
132
133 if factor.log_len() == 0 {
135 self.bound *= factor.get(0);
136 self.factors.pop();
137 }
138 }
139}
140
141#[cfg(test)]
142mod tests {
143 use binius_compute::GlobalAllocator;
144 use binius_field::{Ghash128b, PackedGhash1x128b, Random, arch::OptimalPackedB128};
145 use proptest::prelude::*;
146 use rand::{SeedableRng, prelude::StdRng};
147
148 use super::*;
149
150 fn materialize<P: PackedField>(factors: &[FieldBuffer<P>]) -> FieldBuffer<P> {
152 let n_vars: usize = factors.iter().map(FieldBuffer::log_len).sum();
153 let values = (0..1usize << n_vars)
154 .map(|index| {
155 let mut rest = index;
156 let mut value = P::Scalar::ONE;
157 for factor in factors {
158 let mask = (1 << factor.log_len()) - 1;
159 value *= factor.get(rest & mask);
160 rest >>= factor.log_len();
161 }
162 value
163 })
164 .collect::<Vec<_>>();
165 FieldBuffer::from_values(&values)
166 }
167
168 fn random_factors<P: PackedField>(rng: &mut StdRng, shape: &[usize]) -> Vec<FieldBuffer<P>>
169 where
170 P::Scalar: Random,
171 {
172 shape
173 .iter()
174 .map(|&log_len| {
175 let values = (0..1usize << log_len)
176 .map(|_| P::Scalar::random(&mut *rng))
177 .collect::<Vec<_>>();
178 FieldBuffer::from_values(&values)
179 })
180 .collect()
181 }
182
183 #[test]
184 fn a_factor_with_no_variables_folds_into_the_product() {
185 type P = PackedGhash1x128b;
190 let mut rng = StdRng::seed_from_u64(1);
191
192 let factors = random_factors::<P>(&mut rng, &[0, 2, 0]);
193 let expected = materialize(&factors);
194
195 let factored = FactoredMultilinear::new(factors);
196
197 assert_eq!(factored.n_vars(), 2);
199 for index in 0..4 {
200 assert_eq!(factored.get(index), expected.get(index));
201 }
202 }
203
204 #[test]
205 fn folding_consumes_the_factors_from_the_highest_variable_down() {
206 type P = PackedGhash1x128b;
215 let mut rng = StdRng::seed_from_u64(2);
216
217 let factors = random_factors::<P>(&mut rng, &[1, 2, 1]);
218 let mut factored = FactoredMultilinear::new(factors.clone());
219 let mut reference = materialize(&factors);
220 assert_eq!(factored.n_vars(), 4);
221
222 let challenge = Ghash128b::random(&mut rng);
223 factored.fold_highest_var(challenge);
224 fold_highest_var_inplace(&mut reference, challenge);
225
226 assert_eq!(factored.n_vars(), 3);
228 for index in 0..8 {
229 assert_eq!(factored.get(index), reference.get(index), "index {index}");
230 }
231 }
232
233 #[test]
234 fn factors_may_live_in_an_arena_rather_than_the_heap() {
235 type P = PackedGhash1x128b;
249 let mut rng = StdRng::seed_from_u64(23);
250 let alloc = GlobalAllocator;
251
252 let shape = [2usize, 1];
253 let scalars = shape
254 .iter()
255 .map(|&log_len| {
256 (0..1usize << log_len)
257 .map(|_| Ghash128b::random(&mut rng))
258 .collect::<Vec<_>>()
259 })
260 .collect::<Vec<_>>();
261
262 let mut heap = FactoredMultilinear::<P>::new(
263 scalars
264 .iter()
265 .map(|values| FieldBuffer::<P>::from_values(values)),
266 );
267 let mut arena = FactoredMultilinear::new(
268 scalars
269 .iter()
270 .map(|values| FieldBuffer::<P>::from_values_in(&alloc, values)),
271 );
272
273 assert_eq!(arena.n_vars(), heap.n_vars());
274
275 while heap.n_vars() > 0 {
277 for index in 0..1usize << heap.n_vars() {
278 assert_eq!(arena.get(index), heap.get(index), "index {index}");
279 }
280 let challenge = Ghash128b::random(&mut rng);
281 heap.fold_highest_var(challenge);
282 arena.fold_highest_var(challenge);
283 }
284 assert_eq!(arena.get(0), heap.get(0));
285 }
286
287 proptest! {
288 #[test]
293 fn folding_tracks_the_materialized_product(
294 shape in prop::collection::vec(0usize..=3, 1..=4),
295 seed: u64,
296 ) {
297 type P = OptimalPackedB128;
298 let mut rng = StdRng::seed_from_u64(seed);
299
300 let factors = random_factors::<P>(&mut rng, &shape);
301 let mut factored = FactoredMultilinear::new(factors.clone());
302 let mut reference = materialize(&factors);
303
304 prop_assert_eq!(factored.n_vars(), reference.log_len());
306 for index in 0..reference.len() {
307 prop_assert_eq!(factored.get(index), reference.get(index));
308 }
309
310 while factored.n_vars() > 0 {
312 let challenge = Ghash128b::random(&mut rng);
313 factored.fold_highest_var(challenge);
314 fold_highest_var_inplace(&mut reference, challenge);
315
316 prop_assert_eq!(factored.n_vars(), reference.log_len());
317 for index in 0..reference.len() {
318 prop_assert_eq!(factored.get(index), reference.get(index));
319 }
320 }
321
322 prop_assert_eq!(factored.get(0), reference.get(0));
324 }
325 }
326}