Skip to main content

binius_frontend/eval_form/
batch.rs

1// Copyright 2025-2026 The Binius Developers
2//! Batched execution context for circuit evaluation.
3//!
4//! This is the structure-of-arrays counterpart to the single-instance context.
5//! It evaluates the same bytecode over many independent instances of one circuit at once.
6//! The opcode dispatch is shared through the executor and the execution-context trait.
7//!
8//! The value vector is transposed into a 2D array:
9//! - rows are value-vector indices (wires).
10//! - columns are instances.
11//!
12//! An instruction applies its scalar operation across a whole row: every instance in one pass.
13//! This is the memory order the batch prover wants downstream.
14//!
15//! ```text
16//!                  instance 0   instance 1   ...   instance n-1
17//!   value index 0 [   w        |   w        | ... |   w        ]   <- one row
18//!   value index 1 [   w        |   w        | ... |   w        ]
19//!         ...
20//! ```
21
22use std::array;
23
24use binius_core::Word;
25use binius_utils::strided_array::StridedArray2DViewMut;
26
27use super::{
28	assertion::{MAX_ASSERTION_FAILURES, render_path},
29	exec::EvalContext,
30};
31use crate::{
32	artifact::witness::{AssertionFailure, PopulateError},
33	ir::path::{PathSpec, PathSpecTree},
34};
35
36/// A single assertion failure recorded for the current lowest-failing instance.
37struct InstanceAssertionFailure {
38	path_spec: PathSpec,
39	message: String,
40}
41
42/// The failure of batch witness population, attributed to a single instance.
43///
44/// Serial batched evaluation reports the lowest-indexed failing instance. Parallel batched
45/// evaluation runs over independent stripes, and may report the first failing stripe observed.
46///
47/// The inner [`PopulateError`] is the error's source, so the pair renders as one chain.
48#[derive(Debug, thiserror::Error)]
49#[error("instance {instance} is not satisfied: {source}")]
50#[non_exhaustive]
51pub struct BatchPopulateError {
52	/// The index of the reported failing instance.
53	pub instance: usize,
54	/// The assertion failures recorded for that instance.
55	#[source]
56	pub source: PopulateError,
57}
58
59/// Execution context holding the transposed value array during batch evaluation.
60pub struct BatchExecutionContext<'a, 'v> {
61	/// Rows are value-vector indices; columns are instances.
62	values: &'a mut StridedArray2DViewMut<'v, Word>,
63	/// Failures recorded for [`Self::min_failing_instance`], capped by [`MAX_ASSERTION_FAILURES`].
64	///
65	/// Cleared whenever a strictly lower failing instance is found.
66	/// So a higher-numbered instance's failures never survive to be reported.
67	failures: Vec<InstanceAssertionFailure>,
68	/// Every violation recorded for [`Self::min_failing_instance`], capped or not.
69	///
70	/// Reset alongside `failures` whenever a strictly lower failing instance is discovered.
71	min_failure_count: usize,
72	/// The lowest-indexed instance that has failed an assertion so far.
73	min_failing_instance: Option<usize>,
74}
75
76impl<'a, 'v> BatchExecutionContext<'a, 'v> {
77	pub const fn new(values: &'a mut StridedArray2DViewMut<'v, Word>) -> Self {
78		Self {
79			values,
80			failures: Vec::new(),
81			min_failure_count: 0,
82			min_failing_instance: None,
83		}
84	}
85
86	/// Turn recorded failures into an error attributed to the lowest-failing instance.
87	pub fn check_assertions(
88		self,
89		path_spec_tree: Option<&PathSpecTree>,
90	) -> Result<(), BatchPopulateError> {
91		let Some(instance) = self.min_failing_instance else {
92			return Ok(());
93		};
94
95		// `failures` already holds only records for the reported instance.
96		// Resolve each one's path against the tree.
97		let failures: Vec<AssertionFailure> = self
98			.failures
99			.into_iter()
100			.map(|f| AssertionFailure {
101				path: render_path(path_spec_tree, f.path_spec),
102				detail: f.message,
103			})
104			.collect();
105
106		Err(BatchPopulateError {
107			instance,
108			source: PopulateError {
109				total: self.min_failure_count,
110				failures,
111			},
112		})
113	}
114}
115
116impl EvalContext for BatchExecutionContext<'_, '_> {
117	fn n_instances(&self) -> usize {
118		self.values.width()
119	}
120
121	#[inline]
122	fn load(&self, reg: u32, instance: usize) -> Word {
123		self.values[(reg as usize, instance)]
124	}
125
126	#[inline]
127	fn store(&mut self, reg: u32, instance: usize, value: Word) {
128		self.values[(reg as usize, instance)] = value;
129	}
130
131	/// Runs `op` over whole rows, one contiguous run of instances per register.
132	fn map<const D: usize, const S: usize, F>(&mut self, dsts: [u32; D], srcs: [u32; S], op: F)
133	where
134		F: Fn([Word; S]) -> [Word; D],
135	{
136		let n_instances = self.values.width();
137		let (mut dst_rows, src_rows) = self
138			.values
139			.rows_mut(dsts.map(|reg| reg as usize), srcs.map(|reg| reg as usize));
140		for i in 0..n_instances {
141			let out = op(array::from_fn(|k| src_rows[k][i]));
142			for (row, value) in dst_rows.iter_mut().zip(out) {
143				row[i] = value;
144			}
145		}
146	}
147
148	/// Folds one whole row into another.
149	fn update<F>(&mut self, dst: u32, src: u32, op: F)
150	where
151		F: Fn(Word, Word) -> Word,
152	{
153		let ([dst_row], [src_row]) = self.values.rows_mut([dst as usize], [src as usize]);
154		for (acc, &word) in dst_row.iter_mut().zip(src_row) {
155			*acc = op(*acc, word);
156		}
157	}
158
159	/// Scans whole rows for a violation before walking instances to record any.
160	///
161	/// A satisfied circuit never leaves the scan.
162	fn check<const S: usize, P, M>(
163		&mut self,
164		srcs: [u32; S],
165		path_spec: PathSpec,
166		fails: P,
167		message: M,
168	) where
169		P: Fn([Word; S]) -> bool,
170		M: Fn([Word; S]) -> String,
171	{
172		let n_instances = self.values.width();
173		let any = {
174			let rows = srcs.map(|reg| self.values.row(reg as usize));
175			(0..n_instances).any(|i| fails(array::from_fn(|k| rows[k][i])))
176		};
177		if !any {
178			return;
179		}
180		for i in 0..n_instances {
181			let words = srcs.map(|reg| self.load(reg, i));
182			if fails(words) {
183				self.note_assertion_failure(i, path_spec, message(words));
184			}
185		}
186	}
187
188	/// Record an assertion failure for one instance.
189	///
190	/// A failure for a higher instance than the current lowest-failing one is dropped.
191	/// It can never become the reported instance.
192	/// A failure for a new, strictly lower instance clears every record kept so far.
193	/// Those records belonged to an instance that turned out not to be the one reported.
194	#[cold]
195	fn note_assertion_failure(&mut self, instance: usize, path_spec: PathSpec, message: String) {
196		match self.min_failing_instance {
197			Some(min) if instance > min => return,
198			Some(min) if instance < min => {
199				self.min_failing_instance = Some(instance);
200				self.failures.clear();
201				self.min_failure_count = 0;
202			}
203			// Either the first failure ever seen, or another failure of the current minimum.
204			_ => self.min_failing_instance = Some(instance),
205		}
206
207		self.min_failure_count += 1;
208		if self.failures.len() < MAX_ASSERTION_FAILURES {
209			self.failures
210				.push(InstanceAssertionFailure { path_spec, message });
211		}
212	}
213}
214
215#[cfg(test)]
216mod tests {
217	use binius_core::Word;
218	use binius_utils::strided_array::StridedArray2DViewMut;
219
220	use crate::{
221		MAX_ASSERTION_FAILURES, Wire, artifact::witness::AssertionFailure, builder::CircuitBuilder,
222	};
223
224	// The batched interpreter must reproduce, for every instance, exactly what the single-instance
225	// interpreter produces for the same inputs. This is the core equivalence guarantee.
226	#[test]
227	fn batched_matches_scalar_per_instance() {
228		// A circuit that exercises every row-wise opcode plus a constant, with only witness
229		// inputs and force-committed outputs (no inout wires — the M4 setting).
230		let builder = CircuitBuilder::new();
231		let a = builder.add_witness();
232		let b = builder.add_witness();
233		// Two inputs the caller pins, so every assertion below holds for every instance.
234		let zero = builder.add_witness();
235		let msb = builder.add_witness();
236		let k = builder.add_constant_64(0x0123_4567_89ab_cdef);
237		let c = builder.band(a, b);
238		let d = builder.bxor(a, k);
239		let (sum, cout) = builder.iadd(a, b);
240		let e = builder.rotr(b, 7);
241		let f = builder.bor(c, e);
242		let g = builder.fax(a, b, k);
243		let xs = builder.bxor_multi(&[a, b, k, e, f]);
244		// A condition that differs between instances, so both arms of the select are taken.
245		let lt = builder.icmp_ult(a, b);
246		let sel = builder.select(lt, d, e);
247		let (diff, bout) = builder.isub_bin_bout(a, b, lt);
248		let (mul_hi, mul_lo) = builder.imul(a, b);
249		let (gmul_lo, gmul_hi) = builder.bmul(a, b, k, e);
250		let lanes = builder.iadd_32(a, b);
251		let (lanes_sum, lanes_cout) = builder.iadd32_cin_cout(a, b, lt);
252		// One shift per variant.
253		let shifts = [
254			builder.shl(a, 13),
255			builder.shr(a, 13),
256			builder.sar(a, 13),
257			builder.rotr(a, 13),
258			builder.sll32(a, 13),
259			builder.srl32(a, 13),
260			builder.sra32(a, 13),
261			builder.rotr32(a, 13),
262		];
263		// Assertions that hold for every instance, so each scan runs but records nothing.
264		builder.assert_eq("xor_roundtrip", builder.bxor(d, k), a);
265		builder.assert_eq_cond("xor_roundtrip_cond", builder.bxor(d, k), a, lt);
266		builder.assert_zero("zero_is_zero", zero);
267		builder.assert_false("zero_msb_false", zero);
268		builder.assert_non_zero("msb_is_non_zero", msb);
269		builder.assert_true("msb_is_true", msb);
270		for wire in [
271			c, d, sum, cout, e, f, g, xs, lt, sel, diff, bout, mul_hi, mul_lo, gmul_lo, gmul_hi,
272			lanes, lanes_sum, lanes_cout,
273		]
274		.into_iter()
275		.chain(shifts)
276		{
277			builder.force_commit(wire);
278		}
279		let circuit = builder.build();
280
281		let layout = circuit.value_vec_layout().clone();
282		assert_eq!(layout.n_inout, 0, "fixture should have no inout wires");
283		let combined = layout.combined_len();
284		let full_len = combined + layout.n_scratch;
285		// A large instance count, so the equivalence holds well beyond a handful of instances.
286		// This call goes through the plain batched evaluator, not the tiled parallel dispatch.
287		let n = 1024usize;
288
289		// Distinct inputs per instance.
290		let inputs: Vec<(u64, u64)> = (0..n)
291			.map(|i| {
292				let i = i as u64;
293				(i.wrapping_mul(0x9e37_79b9_7f4a_7c15), i ^ 0x0000_0000_dead_beef)
294			})
295			.collect();
296
297		// Single-instance reference: populate each instance on its own.
298		let scalar: Vec<Vec<Word>> = inputs
299			.iter()
300			.map(|&(x, y)| {
301				let mut filler = circuit.new_witness_filler();
302				filler[a] = Word(x);
303				filler[b] = Word(y);
304				filler[zero] = Word::ZERO;
305				filler[msb] = Word::MSB_ONE;
306				circuit.populate_wire_witness(&mut filler).unwrap();
307				filler.value_vec().combined_witness().to_vec()
308			})
309			.collect();
310
311		// Batched: fill the input rows for every instance, then evaluate all at once.
312		let a_row = circuit.witness_row(a);
313		let b_row = circuit.witness_row(b);
314		let zero_row = circuit.witness_row(zero);
315		let msb_row = circuit.witness_row(msb);
316		let mut data = vec![Word::ZERO; full_len * n];
317		let mut view = StridedArray2DViewMut::without_stride(&mut data, full_len, n).unwrap();
318		for (instance, &(x, y)) in inputs.iter().enumerate() {
319			view[(a_row, instance)] = Word(x);
320			view[(b_row, instance)] = Word(y);
321			view[(zero_row, instance)] = Word::ZERO;
322			view[(msb_row, instance)] = Word::MSB_ONE;
323		}
324		circuit.populate_wire_witness_batched(&mut view).unwrap();
325
326		// Every instance's committed prefix must equal the single-instance witness.
327		for instance in 0..n {
328			for row in 0..combined {
329				assert_eq!(
330					view[(row, instance)],
331					scalar[instance][row],
332					"mismatch at row {row}, instance {instance}"
333				);
334			}
335		}
336	}
337
338	// A batched run must flag the lowest-indexed instance whose inputs violate an assertion.
339	#[test]
340	fn batched_reports_lowest_failing_instance() {
341		// Assert a == b; instances where a != b fail.
342		let builder = CircuitBuilder::new();
343		let a = builder.add_witness();
344		let b = builder.add_witness();
345		builder.assert_eq("a_eq_b", a, b);
346		let circuit = builder.build();
347
348		let layout = circuit.value_vec_layout().clone();
349		let full_len = layout.combined_len() + layout.n_scratch;
350		let n = 4usize;
351
352		// Instances 2 and 3 violate a == b; instance 2 is the lowest.
353		let inputs = [(1u64, 1u64), (7, 7), (4, 5), (9, 8)];
354		let a_row = circuit.witness_row(a);
355		let b_row = circuit.witness_row(b);
356		let mut data = vec![Word::ZERO; full_len * n];
357		let mut view = StridedArray2DViewMut::without_stride(&mut data, full_len, n).unwrap();
358		for (instance, &(x, y)) in inputs.iter().enumerate() {
359			view[(a_row, instance)] = Word(x);
360			view[(b_row, instance)] = Word(y);
361		}
362
363		let err = circuit
364			.populate_wire_witness_batched(&mut view)
365			.expect_err("instances 2 and 3 violate a == b");
366		assert_eq!(err.instance, 2);
367		assert_eq!(err.source.total, 1);
368		assert_eq!(
369			err.source.failures,
370			vec![AssertionFailure {
371				path: ".a_eq_b".to_string(),
372				detail: "Word(0x0000000000000004) != Word(0x0000000000000005)".to_string(),
373			}]
374		);
375		// thiserror supplies the instance prefix and chains the inner error as the source.
376		let rendered = err.to_string();
377		assert!(rendered.starts_with("instance 2 is not satisfied: "), "{rendered}");
378		assert!(rendered.contains(".a_eq_b: Word(0x0000000000000004)"), "{rendered}");
379		// An error message must not carry its own trailing newline.
380		assert!(!rendered.ends_with('\n'), "{rendered}");
381		assert!(std::error::Error::source(&err).is_some(), "the inner error must be the source");
382	}
383
384	// Invariant: the lowest-failing instance is reported, however full the cap got beforehand.
385	//
386	// Fixture state, two assertions in program order:
387	//   assertion one fails for every instance from 50 up, 150 instances in total.
388	//   assertion two fails only for instance 10.
389	//
390	// The cap on retained failures is 100.
391	// Assertion one alone fills it, with records for instances 50..149, before assertion two runs.
392	// Instance 10 is the true minimum, so it must still be the reported instance.
393	// Its own record must survive, unevicted by the unrelated higher-numbered records.
394	#[test]
395	fn batch_min_failing_instance_survives_a_full_cap_of_higher_instances() {
396		let builder = CircuitBuilder::new();
397		let a = builder.add_witness();
398		let b = builder.add_witness();
399		let c = builder.add_witness();
400		let d = builder.add_witness();
401		// Assertion 1: fails for instance >= 50 when the inputs are driven that way below.
402		builder.assert_eq("assertion_one", a, b);
403		// Assertion 2, recorded after assertion 1 for every instance: fails only for instance 10.
404		builder.assert_eq("assertion_two", c, d);
405		let circuit = builder.build();
406
407		let layout = circuit.value_vec_layout().clone();
408		let full_len = layout.combined_len() + layout.n_scratch;
409		// 150 instances fail assertion 1 (50..199).
410		// That overflows the cap of 100 well before assertion 2 ever runs.
411		let n = 200usize;
412
413		let a_row = circuit.witness_row(a);
414		let b_row = circuit.witness_row(b);
415		let c_row = circuit.witness_row(c);
416		let d_row = circuit.witness_row(d);
417		let mut data = vec![Word::ZERO; full_len * n];
418		let mut view = StridedArray2DViewMut::without_stride(&mut data, full_len, n).unwrap();
419		for instance in 0..n {
420			// a == b everywhere except instance >= 50, where assertion 1 fails.
421			let a_val = Word(1);
422			let b_val = Word(if instance >= 50 { 2 } else { 1 });
423			// c == d everywhere except instance 10, where assertion 2 fails.
424			let c_val = Word(3);
425			let d_val = Word(if instance == 10 { 4 } else { 3 });
426			view[(a_row, instance)] = a_val;
427			view[(b_row, instance)] = b_val;
428			view[(c_row, instance)] = c_val;
429			view[(d_row, instance)] = d_val;
430		}
431
432		let err = circuit
433			.populate_wire_witness_batched(&mut view)
434			.expect_err("instance 10 and instances 50..199 all violate an assertion");
435
436		// Instance 10 is the true minimum: 10 < 50.
437		assert_eq!(err.instance, 10);
438		// Exactly one assertion fails for instance 10: assertion 2.
439		assert_eq!(err.source.total, 1);
440		assert_eq!(
441			err.source.failures,
442			vec![AssertionFailure {
443				path: ".assertion_two".to_string(),
444				detail: "Word(0x0000000000000003) != Word(0x0000000000000004)".to_string(),
445			}]
446		);
447	}
448
449	// Invariant: `total` counts every violation the reported instance racked up.
450	// So it can exceed `failures.len()` once the cap kicks in.
451	#[test]
452	fn batch_total_counts_every_violation_past_the_cap_for_the_reported_instance() {
453		// Half again the cap, so the cap alone cannot retain them all.
454		const N: usize = MAX_ASSERTION_FAILURES + 50;
455
456		let builder = CircuitBuilder::new();
457		let x: [Wire; N] = core::array::from_fn(|_| builder.add_witness());
458		let y: [Wire; N] = core::array::from_fn(|_| builder.add_witness());
459		builder.assert_eq_v("pairs", x, y);
460		let circuit = builder.build();
461
462		let layout = circuit.value_vec_layout().clone();
463		let full_len = layout.combined_len() + layout.n_scratch;
464		let n_instances = 4usize;
465
466		let x_rows: Vec<usize> = x.iter().map(|&w| circuit.witness_row(w)).collect();
467		let y_rows: Vec<usize> = y.iter().map(|&w| circuit.witness_row(w)).collect();
468		let mut data = vec![Word::ZERO; full_len * n_instances];
469		let mut view =
470			StridedArray2DViewMut::without_stride(&mut data, full_len, n_instances).unwrap();
471		for instance in 0..n_instances {
472			for i in 0..N {
473				// x == y everywhere except instance 0, where every one of the N pairs disagrees.
474				let x_val = Word(i as u64);
475				let y_val = if instance == 0 {
476					Word(i as u64 + 1)
477				} else {
478					x_val
479				};
480				view[(x_rows[i], instance)] = x_val;
481				view[(y_rows[i], instance)] = y_val;
482			}
483		}
484
485		let err = circuit
486			.populate_wire_witness_batched(&mut view)
487			.expect_err("instance 0 violates every one of the N pairwise assertions");
488
489		assert_eq!(err.instance, 0);
490		// Every one of the N assertions failed for instance 0.
491		assert_eq!(err.source.total, N);
492		// Only the cap's worth of records survives, so `total` exceeds `failures.len()`.
493		assert_eq!(err.source.failures.len(), MAX_ASSERTION_FAILURES);
494		assert!(err.source.total > err.source.failures.len());
495	}
496}