1use std::iter;
23
24use binius_compute::Allocator;
25use binius_field::{Field, PackedField};
26use binius_math::{
27 FieldBuffer, FieldSlice, FieldVec,
28 line::extrapolate_line,
29 multilinear::fold::{fold_highest_var, fold_highest_var_inplace},
30};
31use binius_utils::rayon;
32use itertools::izip;
33
34use super::eq_tracker::EqTracker;
35
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
38pub struct ColId(usize);
39
40impl ColId {
41 pub const fn index(self) -> usize {
44 self.0
45 }
46}
47
48#[derive(Debug, Clone, Copy, PartialEq, Eq)]
50pub struct EqId(usize);
51
52impl EqId {
53 pub const fn index(self) -> usize {
55 self.0
56 }
57}
58
59enum Column<'a, A: Allocator, P: PackedField> {
65 Borrowed(FieldSlice<'a, P>),
66 Owned(FieldVec<P, A>),
67 SplitHalf(FieldVec<P, A>),
74}
75
76pub struct MleStore<'a, A: Allocator, P: PackedField> {
80 n_vars: usize,
81 columns: Vec<Column<'a, A, P>>,
82 n_cols: usize,
85 eq_trackers: Vec<EqTracker<P>>,
86 alloc: &'a A,
88}
89
90impl<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>> MleStore<'a, A, P> {
91 pub const fn new(n_vars: usize, alloc: &'a A) -> Self {
93 Self {
94 n_vars,
95 columns: Vec::new(),
96 n_cols: 0,
97 eq_trackers: Vec::new(),
98 alloc,
99 }
100 }
101
102 pub const fn n_vars(&self) -> usize {
106 self.n_vars
107 }
108
109 pub fn push(&mut self, column: FieldSlice<'a, P>) -> ColId {
113 assert_eq!(
115 column.log_len(),
116 self.n_vars,
117 "column must have number of variables equal to the store"
118 );
119 self.columns.push(Column::Borrowed(column));
120 self.next_col_id()
121 }
122
123 pub fn push_owned(&mut self, column: FieldVec<P, A>) -> ColId {
125 assert_eq!(
127 column.log_len(),
128 self.n_vars,
129 "column must have number of variables equal to the store"
130 );
131 self.columns.push(Column::Owned(column));
132 self.next_col_id()
133 }
134
135 const fn next_col_id(&mut self) -> ColId {
137 let id = ColId(self.n_cols);
138 self.n_cols += 1;
139 id
140 }
141
142 pub fn push_split_half(&mut self, buffer: FieldVec<P, A>) -> [ColId; 2] {
151 assert_eq!(
153 buffer.log_len(),
154 self.n_vars + 1,
155 "buffer must have one more variable than the store so each half matches it"
156 );
157 self.columns.push(Column::SplitHalf(buffer));
158 let low = ColId(self.n_cols);
159 let high = ColId(self.n_cols + 1);
160 self.n_cols += 2;
161 [low, high]
162 }
163
164 pub fn register_eq_tracker(&mut self, eval_point: &[F]) -> EqId {
169 assert_eq!(
171 eval_point.len(),
172 self.n_vars,
173 "evaluation point length must equal the store's number of variables"
174 );
175 let existing = self
178 .eq_trackers
179 .iter()
180 .position(|tracker| &tracker.eval_point()[..self.n_vars] == eval_point);
181 let index = existing.unwrap_or_else(|| {
182 self.eq_trackers.push(EqTracker::new(eval_point));
183 self.eq_trackers.len() - 1
184 });
185 EqId(index)
186 }
187
188 pub fn eq_expansions(&self) -> Vec<&FieldBuffer<P>> {
194 self.eq_trackers
195 .iter()
196 .map(|tracker| tracker.expansion())
197 .collect()
198 }
199
200 pub fn round_context(&self) -> RoundContext<'_, P> {
204 RoundContext {
205 n_vars: self.n_vars,
206 eq_trackers: &self.eq_trackers,
207 }
208 }
209
210 pub fn fold(&mut self, challenge: F) {
215 assert!(self.n_vars > 0, "fold requires at least one remaining variable");
217
218 let n_vars = self.n_vars;
221 let alloc = self.alloc;
222 for column in &mut self.columns {
223 match column {
224 Column::Owned(buffer) => fold_highest_var_inplace(buffer, challenge),
225 Column::Borrowed(slice) => {
226 *column = Column::Owned(fold_highest_var(alloc, slice, challenge));
229 }
230 Column::SplitHalf(buffer) => {
231 let mut split = buffer.split_half_mut();
237 let (mut low, mut high) = split.halves();
238 low.truncate(n_vars);
239 high.truncate(n_vars);
240 fold_highest_var_inplace(&mut low, challenge);
241 fold_highest_var_inplace(&mut high, challenge);
242 }
243 }
244 }
245 for tracker in &mut self.eq_trackers {
246 tracker.fold(challenge);
247 }
248 self.n_vars -= 1;
249 }
250
251 pub fn column(&self, id: ColId) -> FieldSlice<'_, P> {
255 let mut index = id.index();
258 for column in &self.columns {
259 match column {
260 Column::Borrowed(slice) if index == 0 => return slice.as_view(),
261 Column::Owned(buffer) if index == 0 => return buffer.as_view(),
262 Column::SplitHalf(buffer) if index < 2 => {
263 let half_start = index << (buffer.log_len() - 1 - self.n_vars);
266 return buffer.chunk(self.n_vars, half_start);
267 }
268 Column::SplitHalf(_) => index -= 2,
269 _ => index -= 1,
270 }
271 }
272 panic!("column id {} is out of range for a store of {} columns", id.index(), self.n_cols);
273 }
274
275 pub fn column_slices(&self) -> Vec<FieldSlice<'_, P>> {
281 let mut slices = Vec::with_capacity(self.n_cols);
282 for column in &self.columns {
283 match column {
284 Column::Borrowed(slice) => slices.push(slice.as_view()),
285 Column::Owned(buffer) => slices.push(buffer.as_view()),
286 Column::SplitHalf(buffer) => {
287 let high_start = 1 << (buffer.log_len() - 1 - self.n_vars);
291 slices.push(buffer.chunk(self.n_vars, 0));
292 slices.push(buffer.chunk(self.n_vars, high_start));
293 }
294 }
295 }
296 slices
297 }
298
299 pub fn final_evals(&self) -> Vec<F> {
303 assert_eq!(self.n_vars, 0, "final_evals requires all variables to be folded");
305
306 self.column_slices()
307 .iter()
308 .map(|slice| slice.get(0))
309 .collect()
310 }
311
312 pub fn map_reduce<T: Send>(
332 &self,
333 chunk_vars: usize,
334 map: impl (for<'c> Fn(EvaluationChunk<'c, P>) -> T) + Sync,
335 reduce: impl (Fn(T, T, usize) -> T) + Sync,
336 ) -> T {
337 assert!(self.n_vars > 0);
338 let chunk_vars = chunk_vars.min(self.n_vars - 1);
339
340 let col_slices = self.column_slices();
341 let cols = col_slices
342 .iter()
343 .map(|col| {
344 let (lo, hi) = col.split_half();
345 ColumnChunk { lo, hi }
346 })
347 .collect();
348 let eqs = self.eq_expansions().iter().map(|eq| eq.as_view()).collect();
349 let chunk = EvaluationChunk {
350 n_vars: self.n_vars - 1,
351 cols,
352 eqs,
353 };
354 map_reduce_helper(chunk, chunk_vars, &map, &reduce)
355 }
356
357 pub fn map_reduce_with_fold<T: Send>(
371 &mut self,
372 chunk_vars: usize,
373 challenge: F,
374 map: impl (for<'c> Fn(EvaluationChunk<'c, P>) -> T) + Sync,
375 reduce: impl (Fn(T, T, usize) -> T) + Sync,
376 ) -> T {
377 assert!(self.n_vars > 1);
378
379 let n_vars = self.n_vars - 1;
381 let chunk_vars = chunk_vars.min(n_vars - 1);
382
383 if chunk_vars < P::LOG_WIDTH {
386 self.fold(challenge);
387 return self.map_reduce(chunk_vars, map, reduce);
388 }
389
390 let challenge_broadcast = P::broadcast(challenge);
391
392 let alloc = self.alloc;
395 let mut dsts = self
396 .columns
397 .iter()
398 .map(|column| match column {
399 Column::Borrowed(_) => Some(FieldBuffer::zeros_in(alloc, n_vars)),
400 _ => None,
401 })
402 .collect::<Vec<_>>();
403
404 let mut cols = Vec::with_capacity(self.n_cols);
408 for (column, dst) in iter::zip(&mut self.columns, &mut dsts) {
409 match column {
410 Column::Borrowed(src) => {
411 let dst = dst
412 .as_mut()
413 .expect("borrowed columns get a destination buffer")
414 .as_mut();
415 let src = (src as &FieldSlice<'_, P>).as_ref();
416 debug_assert_eq!(src.len(), 1 << (n_vars + 1 - P::LOG_WIDTH));
417
418 let (seg_0, seg_1) = src.split_at(1 << (n_vars - P::LOG_WIDTH));
419 cols.push(PreFoldColumnChunk::OutOfPlace { dst, seg_0, seg_1 });
420 }
421 Column::Owned(buffer) => {
422 let seg = buffer.as_mut();
423 debug_assert_eq!(seg.len(), 1 << (n_vars + 1 - P::LOG_WIDTH));
424
425 let (seg_0, seg_1) = seg.split_at_mut(1 << (n_vars - P::LOG_WIDTH));
426 cols.push(PreFoldColumnChunk::InPlace { seg_0, seg_1 });
427 }
428 Column::SplitHalf(buffer) => {
429 let buffer_log_len = buffer.log_len();
430 let data = buffer.as_mut();
431 let (lo_half, hi_half) =
432 data.split_at_mut(1 << (buffer_log_len - 1 - P::LOG_WIDTH));
433
434 let seg_lo = &mut lo_half[..1 << (n_vars + 1 - P::LOG_WIDTH)];
435 let (seg_lo_0, seg_lo_1) = seg_lo.split_at_mut(1 << (n_vars - P::LOG_WIDTH));
436 cols.push(PreFoldColumnChunk::InPlace {
437 seg_0: seg_lo_0,
438 seg_1: seg_lo_1,
439 });
440
441 let seg_hi = &mut hi_half[..1 << (n_vars + 1 - P::LOG_WIDTH)];
442 let (seg_hi_0, seg_hi_1) = seg_hi.split_at_mut(1 << (n_vars - P::LOG_WIDTH));
443 cols.push(PreFoldColumnChunk::InPlace {
444 seg_0: seg_hi_0,
445 seg_1: seg_hi_1,
446 });
447 }
448 }
449 }
450
451 let cols = cols.into_iter().map(|col| col.split_half()).collect();
454
455 let eqs = self
459 .eq_trackers
460 .iter_mut()
461 .map(|tracker| {
462 let data = tracker.expansion_mut().as_mut();
463 debug_assert_eq!(data.len(), 1 << (n_vars - P::LOG_WIDTH));
464
465 let (seg_0, seg_1) = data.split_at_mut(1 << (n_vars - 1 - P::LOG_WIDTH));
466 PreFoldColumnChunk::InPlace { seg_0, seg_1 }
467 })
468 .collect::<Vec<_>>();
469
470 let chunk = PreFoldEvaluationChunk {
471 n_vars: n_vars - 1,
472 challenge_broadcast: &challenge_broadcast,
473 cols,
474 eqs,
475 };
476 let result = map_reduce_with_fold_helper(chunk, chunk_vars, &map, &reduce);
477
478 for (column, dst) in iter::zip(&mut self.columns, &mut dsts) {
481 match column {
482 Column::Borrowed(_) => {
483 *column = Column::Owned(
484 dst.take()
485 .expect("borrowed columns get a destination buffer"),
486 );
487 }
488 Column::Owned(buffer) => buffer.truncate(n_vars),
489 Column::SplitHalf(_) => {}
490 }
491 }
492 for eq_tracker in &mut self.eq_trackers {
493 eq_tracker.truncate_one_var(challenge);
494 }
495 self.n_vars = n_vars;
496
497 result
498 }
499}
500
501enum PreFoldColumnChunk<'a, P: PackedField> {
508 InPlace {
509 seg_0: &'a mut [P],
510 seg_1: &'a [P],
511 },
512 OutOfPlace {
513 dst: &'a mut [P],
514 seg_0: &'a [P],
515 seg_1: &'a [P],
516 },
517}
518
519impl<'a, P: PackedField> PreFoldColumnChunk<'a, P> {
520 const fn split_half(self) -> [Self; 2] {
522 match self {
523 Self::InPlace { seg_0, seg_1 } => {
524 let (seg_0_lo, seg_0_hi) = seg_0.split_at_mut(seg_0.len() / 2);
525 let (seg_1_lo, seg_1_hi) = seg_1.split_at(seg_1.len() / 2);
526 [
527 Self::InPlace {
528 seg_0: seg_0_lo,
529 seg_1: seg_1_lo,
530 },
531 Self::InPlace {
532 seg_0: seg_0_hi,
533 seg_1: seg_1_hi,
534 },
535 ]
536 }
537 Self::OutOfPlace { dst, seg_0, seg_1 } => {
538 let (dst_lo, dst_hi) = dst.split_at_mut(dst.len() / 2);
539 let (seg_0_lo, seg_0_hi) = seg_0.split_at(seg_0.len() / 2);
540 let (seg_1_lo, seg_1_hi) = seg_1.split_at(seg_1.len() / 2);
541 [
542 Self::OutOfPlace {
543 dst: dst_lo,
544 seg_0: seg_0_lo,
545 seg_1: seg_1_lo,
546 },
547 Self::OutOfPlace {
548 dst: dst_hi,
549 seg_0: seg_0_hi,
550 seg_1: seg_1_hi,
551 },
552 ]
553 }
554 }
555 }
556
557 fn fold_with(self, combine: impl Fn(P, P) -> P) -> &'a [P] {
562 match self {
563 Self::InPlace { seg_0, seg_1 } => {
564 for (out, &hi) in iter::zip(&mut *seg_0, seg_1) {
565 *out = combine(*out, hi);
566 }
567 seg_0
568 }
569 Self::OutOfPlace { dst, seg_0, seg_1 } => {
570 for (out, &lo, &hi) in izip!(&mut *dst, seg_0, seg_1) {
571 *out = combine(lo, hi);
572 }
573 dst
574 }
575 }
576 }
577
578 fn fold(self, challenge_broadcast: &P) -> &'a [P] {
580 self.fold_with(|lo, hi| extrapolate_line(lo, hi, *challenge_broadcast))
581 }
582
583 fn fold_eq(self) -> &'a [P] {
589 self.fold_with(|lo, hi| lo + hi)
590 }
591}
592
593struct PreFoldEvaluationChunk<'a, P: PackedField> {
598 n_vars: usize,
599 challenge_broadcast: &'a P,
600 cols: Vec<[PreFoldColumnChunk<'a, P>; 2]>,
601 eqs: Vec<PreFoldColumnChunk<'a, P>>,
602}
603
604impl<'a, P: PackedField> PreFoldEvaluationChunk<'a, P> {
605 fn split_half(self) -> [Self; 2] {
608 let Self {
609 n_vars,
610 challenge_broadcast,
611 cols,
612 eqs,
613 } = self;
614 let n_vars = n_vars - 1;
615 let (cols_0, cols_1) = cols
616 .into_iter()
617 .map(|[lo, hi]| {
618 let [lo_0, lo_1] = lo.split_half();
619 let [hi_0, hi_1] = hi.split_half();
620 ([lo_0, hi_0], [lo_1, hi_1])
621 })
622 .unzip();
623 let (eqs_0, eqs_1) = eqs
624 .into_iter()
625 .map(|eq| {
626 let [eq_0, eq_1] = eq.split_half();
627 (eq_0, eq_1)
628 })
629 .unzip();
630 [
631 Self {
632 n_vars,
633 challenge_broadcast,
634 cols: cols_0,
635 eqs: eqs_0,
636 },
637 Self {
638 n_vars,
639 challenge_broadcast,
640 cols: cols_1,
641 eqs: eqs_1,
642 },
643 ]
644 }
645
646 fn fold(self) -> EvaluationChunk<'a, P> {
648 let Self {
649 n_vars,
650 challenge_broadcast,
651 cols,
652 eqs,
653 } = self;
654 let cols = cols
655 .into_iter()
656 .map(|[lo, hi]| ColumnChunk {
657 lo: FieldSlice::from_slice(n_vars, lo.fold(challenge_broadcast)),
658 hi: FieldSlice::from_slice(n_vars, hi.fold(challenge_broadcast)),
659 })
660 .collect();
661 let eqs = eqs
662 .into_iter()
663 .map(|eq| FieldSlice::from_slice(n_vars, eq.fold_eq()))
664 .collect();
665 EvaluationChunk { n_vars, cols, eqs }
666 }
667}
668
669pub struct RoundContext<'a, P: PackedField> {
674 n_vars: usize,
675 eq_trackers: &'a [EqTracker<P>],
676}
677
678impl<F: Field, P: PackedField<Scalar = F>> RoundContext<'_, P> {
679 pub const fn n_vars(&self) -> usize {
681 self.n_vars
682 }
683
684 pub fn eq_alpha(&self, id: EqId) -> F {
688 self.eq_trackers[id.index()].next_coordinate()
689 }
690
691 pub const fn eq_prefix(&self, id: EqId) -> F {
698 self.eq_trackers[id.index()].eq_prefix_eval()
699 }
700}
701
702pub struct ColumnChunk<'c, P: PackedField> {
707 pub lo: FieldSlice<'c, P>,
708 pub hi: FieldSlice<'c, P>,
709}
710
711pub struct EvaluationChunk<'c, P: PackedField> {
721 n_vars: usize,
722 cols: Vec<ColumnChunk<'c, P>>,
723 eqs: Vec<FieldSlice<'c, P>>,
724}
725
726impl<'c, P: PackedField> EvaluationChunk<'c, P> {
727 pub fn col(&self, id: ColId) -> &ColumnChunk<'c, P> {
729 &self.cols[id.index()]
730 }
731
732 pub fn eq(&self, id: EqId) -> &FieldSlice<'c, P> {
737 &self.eqs[id.index()]
738 }
739
740 fn split_half(&self) -> [EvaluationChunk<'_, P>; 2] {
743 let Self { n_vars, cols, eqs } = self;
744 let (cols_0, cols_1) = cols
745 .iter()
746 .map(|ColumnChunk { lo, hi }| {
747 let (lo_0, lo_1) = lo.split_half();
748 let (hi_0, hi_1) = hi.split_half();
749 (ColumnChunk { lo: lo_0, hi: hi_0 }, ColumnChunk { lo: lo_1, hi: hi_1 })
750 })
751 .unzip();
752 let (eqs_0, eqs_1) = eqs.iter().map(|col| col.split_half()).unzip();
753 [
754 EvaluationChunk {
755 n_vars: n_vars - 1,
756 cols: cols_0,
757 eqs: eqs_0,
758 },
759 EvaluationChunk {
760 n_vars: n_vars - 1,
761 cols: cols_1,
762 eqs: eqs_1,
763 },
764 ]
765 }
766}
767
768fn map_reduce_helper<P: PackedField, T: Send>(
774 chunk: EvaluationChunk<'_, P>,
775 sub_vars: usize,
776 map: &(impl (for<'a> Fn(EvaluationChunk<'a, P>) -> T) + Sync),
777 reduce: &(impl (Fn(T, T, usize) -> T) + Sync),
778) -> T {
779 if sub_vars == chunk.n_vars {
780 return map(chunk);
781 }
782
783 let level = chunk.n_vars - 1;
785 let [chunk_0, chunk_1] = chunk.split_half();
786 let (ret_0, ret_1) = rayon::join(
787 move || map_reduce_helper(chunk_0, sub_vars, map, reduce),
788 move || map_reduce_helper(chunk_1, sub_vars, map, reduce),
789 );
790 reduce(ret_0, ret_1, level)
791}
792
793fn map_reduce_with_fold_helper<P: PackedField, T: Send>(
794 chunk: PreFoldEvaluationChunk<'_, P>,
795 sub_vars: usize,
796 map: &(impl (for<'a> Fn(EvaluationChunk<'a, P>) -> T) + Sync),
797 reduce: &(impl (Fn(T, T, usize) -> T) + Sync),
798) -> T {
799 if sub_vars == chunk.n_vars {
800 return map(chunk.fold());
801 }
802
803 let level = chunk.n_vars - 1;
805 let [chunk_0, chunk_1] = chunk.split_half();
806 let (ret_0, ret_1) = rayon::join(
807 move || map_reduce_with_fold_helper(chunk_0, sub_vars, map, reduce),
808 move || map_reduce_with_fold_helper(chunk_1, sub_vars, map, reduce),
809 );
810 reduce(ret_0, ret_1, level)
811}
812
813#[cfg(test)]
814mod tests {
815 use binius_compute::GlobalAllocator;
816 use binius_field::{Field, FieldOps, PackedField};
817 use binius_math::test_utils::{Packed128b, random_field_buffer, random_scalars};
818 use itertools::Itertools;
819 use rand::{SeedableRng, rngs::StdRng};
820
821 use super::*;
822
823 fn chunk_aggregate<P: PackedField>(
826 chunk: &EvaluationChunk<'_, P>,
827 col_ids: &[ColId],
828 eq_ids: &[EqId],
829 ) -> P::Scalar {
830 let mut acc = P::Scalar::ZERO;
831 for (i, &col_id) in col_ids.iter().enumerate() {
832 let col = chunk.col(col_id);
833 let eq = chunk.eq(eq_ids[i % eq_ids.len()]);
834 for j in 0..col.lo.len() {
835 acc += eq.get(j) * col.lo.get(j) * col.hi.get(j);
836 }
837 }
838 acc
839 }
840
841 #[test]
844 fn column_matches_column_slices() {
845 type P = Packed128b;
846 type F = <P as FieldOps>::Scalar;
847
848 let n_vars = 5;
849 let mut rng = StdRng::seed_from_u64(2);
850 let alloc = GlobalAllocator;
851
852 let borrowed = random_field_buffer::<P>(&mut rng, n_vars);
855 let mut store = MleStore::<GlobalAllocator, P>::new(n_vars, &alloc);
856 let mut col_ids = vec![store.push(borrowed.as_view())];
857 col_ids.push(store.push_owned(random_field_buffer::<P>(&mut rng, n_vars)));
858 col_ids.extend(store.push_split_half(random_field_buffer::<P>(&mut rng, n_vars + 1)));
859 col_ids.push(store.push_owned(random_field_buffer::<P>(&mut rng, n_vars)));
860
861 let challenges = random_scalars::<F>(&mut rng, n_vars);
862 for (round, &challenge) in challenges.iter().enumerate() {
863 for (&id, expected) in iter::zip(&col_ids, store.column_slices()) {
864 let got = store.column(id);
865 assert_eq!(got.log_len(), expected.log_len(), "length mismatch in round {round}");
866 for i in 0..expected.len() {
867 assert_eq!(got.get(i), expected.get(i), "scalar {i} mismatch in round {round}");
868 }
869 }
870 store.fold(challenge);
871 }
872 }
873
874 #[test]
875 fn map_reduce_pairs_on_highest_variable() {
876 type P = Packed128b;
877 type F = <P as FieldOps>::Scalar;
878
879 let n_vars = 7;
880 let mut rng = StdRng::seed_from_u64(0);
881 let alloc = GlobalAllocator;
882
883 let borrowed = [
885 random_field_buffer::<P>(&mut rng, n_vars),
886 random_field_buffer::<P>(&mut rng, n_vars),
887 ];
888 let mut store = MleStore::<GlobalAllocator, P>::new(n_vars, &alloc);
889 let mut col_ids = borrowed
890 .iter()
891 .map(|col| store.push(col.as_view()))
892 .collect::<Vec<_>>();
893 col_ids.push(store.push_owned(random_field_buffer::<P>(&mut rng, n_vars)));
894 col_ids.extend(store.push_split_half(random_field_buffer::<P>(&mut rng, n_vars + 1)));
895
896 let eq_ids = (0..2)
897 .map(|_| store.register_eq_tracker(&random_scalars::<F>(&mut rng, n_vars)))
898 .collect::<Vec<_>>();
899
900 let cols = store.column_slices();
904 let eqs = store.eq_expansions();
905 let half = 1usize << (n_vars - 1);
906 let mut expected = F::ZERO;
907 for (i, col) in cols.iter().enumerate() {
908 let eq = eqs[i % eqs.len()];
909 for j in 0..half {
910 expected += eq.get(j) * col.get(j) * col.get(half + j);
911 }
912 }
913
914 for chunk_vars in 0..n_vars {
915 let got = store.map_reduce(
916 chunk_vars,
917 |chunk| chunk_aggregate(&chunk, &col_ids, &eq_ids),
918 |lhs, rhs, _level| lhs + rhs,
919 );
920 assert_eq!(got, expected, "mismatch at chunk_vars = {chunk_vars}");
921 }
922 }
923
924 #[test]
925 fn map_reduce_with_fold_matches_fold_then_map_reduce() {
926 type P = Packed128b;
927 type F = <P as FieldOps>::Scalar;
928
929 let n_vars = 8;
930 let mut rng = StdRng::seed_from_u64(1);
931
932 let borrowed = [
936 random_field_buffer::<P>(&mut rng, n_vars),
937 random_field_buffer::<P>(&mut rng, n_vars),
938 ];
939 let owned = random_field_buffer::<P>(&mut rng, n_vars);
940 let split = random_field_buffer::<P>(&mut rng, n_vars + 1);
941 let eq_points = [
942 random_scalars::<F>(&mut rng, n_vars),
943 random_scalars::<F>(&mut rng, n_vars),
944 ];
945 let challenge = random_scalars::<F>(&mut rng, 1)[0];
946 let alloc = GlobalAllocator;
947
948 let build = || {
949 let mut store = MleStore::<GlobalAllocator, P>::new(n_vars, &alloc);
950 let mut col_ids = borrowed
951 .iter()
952 .map(|col| store.push(col.as_view()))
953 .collect::<Vec<_>>();
954 col_ids.push(store.push_owned(owned.clone()));
955 col_ids.extend(store.push_split_half(split.clone()));
956 let eq_ids = eq_points
957 .iter()
958 .map(|point| store.register_eq_tracker(point))
959 .collect::<Vec<_>>();
960 (store, col_ids, eq_ids)
961 };
962
963 let scalars =
965 |slice: &FieldSlice<'_, P>| (0..slice.len()).map(|i| slice.get(i)).collect_vec();
966 let state = |store: &MleStore<'_, GlobalAllocator, P>| {
967 let cols = store.column_slices().iter().flat_map(scalars).collect_vec();
968 let eqs = store
969 .eq_expansions()
970 .iter()
971 .flat_map(|eq| scalars(&eq.as_view()))
972 .collect_vec();
973 (store.n_vars(), cols, eqs)
974 };
975
976 for chunk_vars in 0..n_vars - 1 {
979 let (mut fold_first, col_ids, eq_ids) = build();
980 fold_first.fold(challenge);
981 let expected = fold_first.map_reduce(
982 chunk_vars,
983 |chunk| chunk_aggregate(&chunk, &col_ids, &eq_ids),
984 |lhs, rhs, _level| lhs + rhs,
985 );
986
987 let (mut fused, col_ids, eq_ids) = build();
988 let got = fused.map_reduce_with_fold(
989 chunk_vars,
990 challenge,
991 |chunk| chunk_aggregate(&chunk, &col_ids, &eq_ids),
992 |lhs, rhs, _level| lhs + rhs,
993 );
994
995 assert_eq!(got, expected, "result mismatch at chunk_vars = {chunk_vars}");
996 assert_eq!(
997 state(&fold_first),
998 state(&fused),
999 "folded-state mismatch at chunk_vars = {chunk_vars}"
1000 );
1001 }
1002
1003 let (mut fold_first, fold_col_ids, fold_eq_ids) = build();
1006 let (mut fused, fused_col_ids, fused_eq_ids) = build();
1007 let challenges = random_scalars::<F>(&mut rng, n_vars);
1008 for (round, &challenge) in challenges.iter().take(n_vars - 1).enumerate() {
1009 let n = fused.n_vars();
1010 let chunk_vars = (n - 2).min(3);
1011
1012 fold_first.fold(challenge);
1013 let expected = fold_first.map_reduce(
1014 chunk_vars,
1015 |chunk| chunk_aggregate(&chunk, &fold_col_ids, &fold_eq_ids),
1016 |lhs, rhs, _level| lhs + rhs,
1017 );
1018 let got = fused.map_reduce_with_fold(
1019 chunk_vars,
1020 challenge,
1021 |chunk| chunk_aggregate(&chunk, &fused_col_ids, &fused_eq_ids),
1022 |lhs, rhs, _level| lhs + rhs,
1023 );
1024
1025 assert_eq!(got, expected, "result mismatch in round {round}");
1026 assert_eq!(state(&fold_first), state(&fused), "folded-state mismatch in round {round}");
1027 }
1028 }
1029}