Skip to main content

binius_ip_prover/sumcheck/
mle_store.rs

1// Copyright 2026 The Binius Developers
2
3//! Shared multilinear column store for sumcheck round evaluators.
4//!
5//! An [`MleStore`] owns the equal-length multilinear columns that a group of
6//! [round evaluators](super::round_evaluator) reads, along with the deduplicated
7//! equality-indicator trackers for MLE-check evaluation points. Columns enter the
8//! store either borrowed ([`MleStore::push`]) or owned ([`MleStore::push_owned`]) and are
9//! addressed by the returned [`ColId`], so several evaluators can read — and the store can fold —
10//! one shared column exactly once per challenge.
11//!
12//! # Invariant
13//!
14//! The store folds — columns and eq trackers both; evaluators only read. Every column and every
15//! registered tracker advances exactly once per [`MleStore::fold`] call, no matter how many
16//! evaluators reference it.
17//!
18//! Folding is eager: [`MleStore::fold`] advances every column immediately, and the round pass
19//! over the columns is a plain read. A deferred-fold variant that fuses the fold into the next
20//! round's read pass can replace the internals without changing this interface.
21
22use 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/// Identifier of a column held by an [`MleStore`].
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
38pub struct ColId(usize);
39
40impl ColId {
41	/// Returns the position of the column in the store, which indexes the
42	/// [`MleStore::final_evals`] output.
43	pub const fn index(self) -> usize {
44		self.0
45	}
46}
47
48/// Identifier of an equality-indicator tracker held by an [`MleStore`].
49#[derive(Debug, Clone, Copy, PartialEq, Eq)]
50pub struct EqId(usize);
51
52impl EqId {
53	/// Returns the registration position of the tracker in the store.
54	pub const fn index(self) -> usize {
55		self.0
56	}
57}
58
59/// One physical entry in the store, holding one or two logical columns.
60///
61/// A `Borrowed` or `Owned` entry is a single column. A `SplitHalf` entry holds two adjacent
62/// columns — the low and high halves of one parent buffer — in a single allocation, so no copy is
63/// made to separate them.
64enum Column<'a, A: Allocator, P: PackedField> {
65	Borrowed(FieldSlice<'a, P>),
66	Owned(FieldVec<P, A>),
67	/// A parent buffer whose low and high halves are two adjacent columns.
68	///
69	/// Pushed by [`MleStore::push_split_half`]. The buffer keeps its original length for the life
70	/// of the store; each [`MleStore::fold`] advances both halves in place within it, and the two
71	/// columns are read as the front `2^n_vars` scalars of the low and high halves. This shares one
72	/// allocation between the sibling columns with no copy at any point.
73	SplitHalf(FieldVec<P, A>),
74}
75
76/// A store of equal-length multilinear columns shared by a group of round evaluators.
77///
78/// See the [module documentation](self) for the folding invariant.
79pub struct MleStore<'a, A: Allocator, P: PackedField> {
80	n_vars: usize,
81	columns: Vec<Column<'a, A, P>>,
82	/// Number of logical columns, counting each [`Column::SplitHalf`] entry as two. This is the
83	/// number of assigned [`ColId`]s and the length of the [`Self::final_evals`] output.
84	n_cols: usize,
85	eq_trackers: Vec<EqTracker<P>>,
86	/// Allocator the owned/promoted columns are drawn from; borrowed for `'a`.
87	alloc: &'a A,
88}
89
90impl<'a, A: Allocator, F: Field, P: PackedField<Scalar = F>> MleStore<'a, A, P> {
91	/// Creates an empty store over columns with `n_vars` variables.
92	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	/// Returns the number of variables remaining in the columns.
103	///
104	/// Decrements with each [`Self::fold`] call.
105	pub const fn n_vars(&self) -> usize {
106		self.n_vars
107	}
108
109	/// Pushes a borrowed column and returns its identifier.
110	///
111	/// The column is not copied; the first [`Self::fold`] writes into a fresh half-size buffer.
112	pub fn push(&mut self, column: FieldSlice<'a, P>) -> ColId {
113		// precondition
114		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	/// Pushes an owned column and returns its identifier.
124	pub fn push_owned(&mut self, column: FieldVec<P, A>) -> ColId {
125		// precondition
126		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	/// Allocates the identifier for one newly pushed logical column.
136	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	/// Pushes the low and high halves of `buffer` as two columns, returning their ids `[low,
143	/// high]`.
144	///
145	/// The halves are not copied: the store takes ownership of `buffer` and holds both columns in
146	/// it as a single split-half entry, so no up-front copy of the full buffer is made.
147	/// Each [`Self::fold`] advances both halves in place within the buffer. `buffer` splits on its
148	/// highest variable, so its low half fixes that variable to 0 and its high half to 1 —
149	/// matching the store's high-to-low fold order.
150	pub fn push_split_half(&mut self, buffer: FieldVec<P, A>) -> [ColId; 2] {
151		// precondition
152		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	/// Registers an equality-indicator tracker for an MLE-check evaluation point.
165	///
166	/// Trackers are deduplicated: evaluators registering the same evaluation point share one
167	/// tracker, which the store folds once per challenge.
168	pub fn register_eq_tracker(&mut self, eval_point: &[F]) -> EqId {
169		// precondition
170		assert_eq!(
171			eval_point.len(),
172			self.n_vars,
173			"evaluation point length must equal the store's number of variables"
174		);
175		// Trackers fold in lockstep with the store, so the remaining coordinates of an existing
176		// tracker are the prefix of its original evaluation point.
177		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	/// Returns the equality-indicator expansion of every registered tracker, in [`EqId`] order.
189	///
190	/// The driving prover slices each expansion per chunk once per round; the returned order
191	/// matches [`EqId::index`], so an evaluator's tracker id indexes the resulting per-chunk
192	/// slices.
193	pub fn eq_expansions(&self) -> Vec<&FieldBuffer<P>> {
194		self.eq_trackers
195			.iter()
196			.map(|tracker| tracker.expansion())
197			.collect()
198	}
199
200	/// Returns the read-only view of the current round that the evaluators interpolate against.
201	///
202	/// Valid until the next [`Self::fold`], which is when the values it exposes advance.
203	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	/// Folds every column and every eq tracker with a verifier challenge.
211	///
212	/// Columns fold on the highest variable, matching the high-to-low binding order of the
213	/// sumcheck provers this store backs.
214	pub fn fold(&mut self, challenge: F) {
215		// precondition
216		assert!(self.n_vars > 0, "fold requires at least one remaining variable");
217
218		// The number of live variables in each column before this fold; a split-half buffer keeps
219		// its full length, so its halves must be truncated to this before folding.
220		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					// The first fold of a borrowed column writes into a fresh half-size owned
227					// buffer, avoiding an up-front copy of the full column.
228					*column = Column::Owned(fold_highest_var(alloc, slice, challenge));
229				}
230				Column::SplitHalf(buffer) => {
231					// Fold each half on its own highest variable in place. The two halves are the
232					// two columns, so folding the whole buffer's highest variable would instead
233					// combine them; splitting first binds each column's variable independently. The
234					// buffer keeps its length — the folded columns are the (now shorter) fronts of
235					// its halves — so no copy is made.
236					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	/// Returns the borrowed view of one logical column, addressed by its [`ColId`].
252	///
253	/// The slice spans the column's `2^n_vars()` live scalars, so it tracks the store's folding.
254	pub fn column(&self, id: ColId) -> FieldSlice<'_, P> {
255		// A split-half entry carries two logical columns, so walk the physical entries counting
256		// down the logical index rather than indexing `columns` directly.
257		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					// The buffer holds the two columns as its low and high halves; each column is
264					// the front `2^n_vars` scalars of one half, so read it as that half's chunk 0.
265					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	/// Expands the store into one borrowed slice per logical column, in [`ColId`] order.
276	///
277	/// A split-half entry expands into the front `2^n_vars` scalars of its low and high
278	/// halves, so the returned length is the logical column count — larger than the physical entry
279	/// count whenever a split-half column is present.
280	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					// The buffer holds the two columns as its low and high halves; each column is
288					// the front `2^n_vars` scalars of one half, so read it as that half's
289					// chunk 0.
290					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	/// Returns the evaluation of every column at the challenge point, indexed by [`ColId`].
300	///
301	/// Each column's evaluation is computed once, no matter how many claims read the column.
302	pub fn final_evals(&self) -> Vec<F> {
303		// precondition
304		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	/// Maps every chunk of the halved hypercube through `map` and combines the results with
313	/// `reduce`, driven by a recursive [`rayon::join`] tree.
314	///
315	/// The store's columns are expanded once — a split-half column becomes its two halves — and
316	/// each column is split on the round's highest variable into its low and high halves. The
317	/// recursion peels the remaining variables off both halves, and off every eq-indicator
318	/// expansion, together, so each leaf hands `map` one [`EvaluationChunk`]: the paired column
319	/// halves and the matching eq chunk, at `2^chunk_vars` scalars per half.
320	///
321	/// `reduce(low, high, level)` combines the two sub-results of a bisection, where `level` is the
322	/// index of the variable that was bisected to produce them (`low` fixes it to 0, `high` to 1).
323	/// A plain sum ignores `level`; an eq-weighted reduction uses it to pick the round's
324	/// coordinate.
325	///
326	/// `chunk_vars` is capped at `n_vars() - 1`, so leaves never exceed the halved hypercube.
327	///
328	/// ## Preconditions
329	///
330	/// * `n_vars()` must be greater than 0.
331	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	/// Folds the store with `challenge` and, in the same pass, maps and reduces the resulting
358	/// halved hypercube — equivalent to [`Self::fold`] followed by [`Self::map_reduce`], but
359	/// folding each column and eq expansion into the map's read of it so they are touched once
360	/// instead of twice.
361	///
362	/// `chunk_vars` is capped at `n_vars() - 2` (the folded store's leaf size). For a chunk size
363	/// below `P::LOG_WIDTH` the columns are already cache-resident and the fused pass cannot split
364	/// sub-packing-width leaves, so this falls back to a plain [`Self::fold`] then
365	/// [`Self::map_reduce`].
366	///
367	/// ## Preconditions
368	///
369	/// * `n_vars()` must be greater than 1.
370	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		// Decrement n_vars to reflect the fold.
380		let n_vars = self.n_vars - 1;
381		let chunk_vars = chunk_vars.min(n_vars - 1);
382
383		// Small rounds are cache-resident, so fusing buys nothing, and the raw-slice split cannot
384		// express a sub-packing-width leaf; fold and map-reduce in two clean passes instead.
385		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		// Fresh destination buffers for the borrowed columns, held outside the column borrow so
393		// they can be moved into the store once the fold has written them.
394		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		// Build one deferred-fold producer per logical column: its low and high halves paired on
405		// the round's highest variable, folding in place (owned/split-half) or into a fresh `dst`
406		// (borrowed).
407		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		// Split each producer into the `[low, high]` pair whose outputs are the two halves of the
452		// folded column.
453		let cols = cols.into_iter().map(|col| col.split_half()).collect();
454
455		// Carry each eq expansion as an in-place producer over its low and high halves. The
456		// recursion contracts it into its front half via `fold_eq`; `truncate_one_var` below then
457		// advances each tracker's bookkeeping over the folded-out variable.
458		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		// The fold wrote each column's folded data into the front of its buffer (or into `dst`);
479		// persist it so the store matches a plain `fold`.
480		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
501/// The deferred fold of one column half or one eq expansion.
502///
503/// A column half folds with [`Self::fold`], interpolating `seg_0` and `seg_1` on the round's
504/// highest variable — in place over `seg_0`, or into a fresh `dst` for a borrowed column. An eq
505/// expansion folds with [`Self::fold_eq`], which contracts (sums) the halves instead of
506/// interpolating them.
507enum 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	/// Bisects the producer's output on its highest variable, splitting each segment in half.
521	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	/// Combines the two segments with `combine` and returns the output slice.
558	///
559	/// The in-place form reads its low half out of the destination it overwrites.
560	/// The out-of-place form reads both halves from the borrowed source.
561	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	/// Folds a column half, interpolating its two segments on the round's variable.
579	fn fold(self, challenge_broadcast: &P) -> &'a [P] {
580		self.fold_with(|lo, hi| extrapolate_line(lo, hi, *challenge_broadcast))
581	}
582
583	/// Contracts an eq expansion by summing its two segments.
584	///
585	/// Summing marginalises the bound variable out, which is part (3) of the [Gruen24] split.
586	///
587	/// [Gruen24]: <https://eprint.iacr.org/2024/108>
588	fn fold_eq(self) -> &'a [P] {
589		self.fold_with(|lo, hi| lo + hi)
590	}
591}
592
593/// A range of the halved hypercube whose values have not yet been folded, the deferred-fold
594/// counterpart of [`EvaluationChunk`]. Each column is a `[low, high]` pair of fold producers and
595/// each eq expansion is a single producer; [`Self::fold`] runs them all to produce an
596/// [`EvaluationChunk`] at a leaf.
597struct 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	/// Bisects the range on its highest remaining variable, matching
606	/// [`EvaluationChunk::split_half`].
607	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	/// Folds every column into its low and high halves, producing the leaf [`EvaluationChunk`].
647	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
669/// The scalar state of one round, as an evaluator may read it while interpolating.
670///
671/// Holds no columns and no allocator, so an evaluator is generic over neither.
672/// The store has not folded this round yet, so every value here is the current round's.
673pub 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	/// Returns the number of variables not yet bound, this round's included.
680	pub const fn n_vars(&self) -> usize {
681		self.n_vars
682	}
683
684	/// Returns the coordinate of the variable this round binds, for a registered eq tracker.
685	///
686	/// This is the round's `alpha`: the highest coordinate of that tracker's point still unbound.
687	pub fn eq_alpha(&self, id: EqId) -> F {
688		self.eq_trackers[id.index()].next_coordinate()
689	}
690
691	/// Returns the equality prefix of a registered eq tracker.
692	///
693	/// This is the product of the equality terms over all coordinates bound so far.
694	/// The [Gruen24] technique multiplies it into each round polynomial.
695	///
696	/// [Gruen24]: <https://eprint.iacr.org/2024/108>
697	pub const fn eq_prefix(&self, id: EqId) -> F {
698		self.eq_trackers[id.index()].eq_prefix_eval()
699	}
700}
701
702/// One column's low and high halves within an [`EvaluationChunk`].
703///
704/// The column is split on the round's highest variable: `lo` fixes that variable to 0, `hi` to 1.
705/// Both range over the chunk's scalars.
706pub struct ColumnChunk<'c, P: PackedField> {
707	pub lo: FieldSlice<'c, P>,
708	pub hi: FieldSlice<'c, P>,
709}
710
711/// A range of the halved hypercube, prepared for the round evaluators.
712///
713/// Holds, over `n_vars` variables, the paired low/high halves of every logical column and the
714/// eq-indicator expansion of every tracker. Each column was split on the round's highest variable,
715/// so a [`ColumnChunk`]'s `lo` and `hi` differ only in that variable. The range is bisected — the
716/// highest remaining variable peeled off both halves of every column and off every eq expansion —
717/// down to the leaves that [`MleStore::map_reduce`] hands to its `map` callback. A
718/// column read by several evaluators is split a single time. Evaluators read their columns by
719/// [`ColId`] and their eq trackers by [`EqId`].
720pub 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	/// Returns the low and high halves of a column at this chunk.
728	pub fn col(&self, id: ColId) -> &ColumnChunk<'c, P> {
729		&self.cols[id.index()]
730	}
731
732	/// Returns the equality-indicator expansion of a registered tracker at this chunk.
733	///
734	/// The expansion ranges over the halved hypercube, so it is chunked with the same chunk index
735	/// as the column halves.
736	pub fn eq(&self, id: EqId) -> &FieldSlice<'c, P> {
737		&self.eqs[id.index()]
738	}
739
740	/// Bisects the range into its two halves on the highest remaining variable, splitting both
741	/// halves of every column and every eq expansion. Each returned chunk has one fewer variable.
742	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
768/// Recursively maps and reduces an [`EvaluationChunk`] for [`MleStore::map_reduce`].
769///
770/// Once the chunk has been narrowed to `sub_vars` variables it is handed to `map`; otherwise it is
771/// bisected with [`EvaluationChunk::split_half`] and the two halves are mapped in parallel and
772/// combined with `reduce`.
773fn 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	// The bisection binds the highest remaining variable; its index is the reduction level.
784	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	// The bisection binds the highest remaining variable; its index is the reduction level.
804	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	// A per-chunk aggregate that is sensitive to both the low/high pairing within a column and the
824	// alignment of each eq expansion with its column, so a wrong recursion pairing changes the sum.
825	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	// `column` and `column_slices` walk the entries independently, so pin them to each other over
842	// every entry kind and over the folds that shrink a split-half column inside its parent buffer.
843	#[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		// A split-half entry sits between single-column entries, so the walk has to skip two
853		// logical columns in one step to reach the last id.
854		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		// A mix of column kinds so `chunk` exercises borrowed, owned, and split-half entries.
884		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		// Independent reference: the aggregate over the whole halved hypercube, pairing each
901		// logical column's front half (highest variable = 0) with its back half (= 1) at the same
902		// index. This is the pairing `map_reduce` must reproduce, whatever the chunking.
903		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		// Source data. Borrowed columns are read but never mutated by either path, so both stores
933		// can share them; owned and split-half buffers are folded in place, so each store clones
934		// its own.
935		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		// The store's folded state: remaining variable count plus every column and eq scalar.
964		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		// chunk_vars below P::LOG_WIDTH takes the fallback path; at or above it takes the fused
977		// path.
978		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		// Fold both stores round by round in lockstep, exercising split-half columns once the store
1004		// has shrunk below the parent buffer's length.
1005		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}