Skip to main content

binius_spartan_frontend/
wire_elimination.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::mem;
5
6use binius_field::Field;
7
8use super::{
9	circuit_builder::{ConstraintSystemIR, WireStatus},
10	constraint_system::{MulConstraint, WireKind},
11};
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14enum MulPosition {
15	A,
16	B,
17	C,
18}
19
20#[derive(Debug, Clone, PartialEq, Eq)]
21enum UseSite {
22	Add { index: u32 },
23	Mul { index: u32, position: MulPosition },
24}
25
26fn remove_use(uses: &mut Vec<UseSite>, use_site: &UseSite) -> Option<UseSite> {
27	let use_site_idx = uses
28		.iter()
29		.position(|use_site_rhs| use_site_rhs == use_site)?;
30	Some(uses.swap_remove(use_site_idx))
31}
32
33#[derive(Debug, Clone)]
34pub struct CostModel {
35	pub wire_cost: u64,
36	pub mul_cost: u64,
37	pub ref_cost: u64,
38}
39
40impl Default for CostModel {
41	fn default() -> Self {
42		// This is just a guess, can be tuned later.
43		CostModel {
44			wire_cost: 16,
45			mul_cost: 2,
46			ref_cost: 1,
47		}
48	}
49}
50
51/// Wire elimination optimization pass on ConstraintSystemIR.
52///
53/// Eliminates private wires from zero constraints by substitution, reducing the size
54/// of the witness while maintaining constraint system validity.
55///
56/// Maintains invariant: `private_wire_uses[i]` is empty if status is not Unknown.
57/// For Unknown wires, use sites must exactly match constraints referencing that wire.
58pub struct WireEliminationPass<F: Field> {
59	cost_model: CostModel,
60	ir: ConstraintSystemIR<F>,
61	/// Vector of use sites for each private wire. Empty if wire status is not Unknown.
62	private_wire_uses: Vec<Vec<UseSite>>,
63}
64
65pub fn run_wire_elimination<F: Field>(
66	cost_model: CostModel,
67	ir: ConstraintSystemIR<F>,
68) -> ConstraintSystemIR<F> {
69	let mut pass = WireEliminationPass::new(cost_model, ir);
70	pass.run();
71	pass.finish()
72}
73
74impl<F: Field> WireEliminationPass<F> {
75	pub fn new(cost_model: CostModel, ir: ConstraintSystemIR<F>) -> Self {
76		let mut private_wire_uses = vec![Vec::new(); ir.private_wires_status.len()];
77
78		// Populate use sites for all private wires with Unknown status
79		for (i, operand) in ir.zero_constraints.iter().enumerate() {
80			for wire in operand.wires() {
81				if matches!(wire.kind, WireKind::Private)
82					&& matches!(ir.private_wires_status[wire.id as usize], WireStatus::Unknown)
83				{
84					private_wire_uses[wire.id as usize].push(UseSite::Add { index: i as u32 });
85				}
86			}
87		}
88
89		for (i, MulConstraint { a, b, c }) in ir.mul_constraints.iter().enumerate() {
90			for (position, operand) in [
91				(MulPosition::A, a),
92				(MulPosition::B, b),
93				(MulPosition::C, c),
94			] {
95				for wire in operand.wires() {
96					if matches!(wire.kind, WireKind::Private)
97						&& matches!(ir.private_wires_status[wire.id as usize], WireStatus::Unknown)
98					{
99						private_wire_uses[wire.id as usize].push(UseSite::Mul {
100							index: i as u32,
101							position,
102						});
103					}
104				}
105			}
106		}
107
108		Self {
109			cost_model,
110			ir,
111			private_wire_uses,
112		}
113	}
114
115	fn run(&mut self) {
116		// Go over each ADD constraint, in order, and find its best candidate. Eliminate it if
117		// there's a candidate.
118		for constraint_idx in 0..self.ir.zero_constraints.len() {
119			if let Some(idx) = self.pruning_candidate(constraint_idx) {
120				self.eliminate(constraint_idx, idx);
121			}
122		}
123	}
124
125	fn pruning_candidate(&self, constraint_idx: usize) -> Option<usize> {
126		let operand = &self.ir.zero_constraints[constraint_idx];
127
128		let (idx, n_uses) = (0..operand.len())
129			.filter_map(|idx| {
130				let wire = operand.wires()[idx];
131				if matches!(wire.kind, WireKind::Private)
132					&& matches!(self.ir.private_wires_status[wire.id as usize], WireStatus::Unknown)
133				{
134					Some((idx, self.private_wire_uses[wire.id as usize].len()))
135				} else {
136					None
137				}
138			})
139			.min_by_key(|(_idx, uses_len)| *uses_len)?;
140
141		let decrement = self.cost_model.wire_cost
142			+ self.cost_model.mul_cost
143			+ operand.len() as u64 * self.cost_model.ref_cost;
144
145		// For each other wire in the ADD constraint operand, we must add one reference for each
146		// constraint that the eliminated wire was used in.
147		let increment = (n_uses as u64) * (operand.len() as u64 - 1) * self.cost_model.ref_cost;
148
149		if decrement > increment {
150			Some(idx)
151		} else {
152			None
153		}
154	}
155
156	fn eliminate(&mut self, constraint_idx: usize, idx: usize) {
157		// Remove the constraint with `take`. Empty ADD constraints are dropped in `finish()`.
158		let operand = mem::take(&mut self.ir.zero_constraints[constraint_idx]);
159
160		// Remove the eliminated constraint from all uses.
161		let eliminated_use_site = UseSite::Add {
162			index: constraint_idx as u32,
163		};
164		for wire in operand.wires() {
165			if matches!(wire.kind, WireKind::Private)
166				&& matches!(self.ir.private_wires_status[wire.id as usize], WireStatus::Unknown)
167			{
168				remove_use(&mut self.private_wire_uses[wire.id as usize], &eliminated_use_site)
169					.expect("invariant: uses are kept in sync by algorithm invariant");
170			}
171		}
172
173		// Prune the wire and get its remaining uses.
174		let wire_idx = operand.wires()[idx].id as usize;
175		assert!(
176			matches!(self.ir.private_wires_status[wire_idx], WireStatus::Unknown),
177			"precondition: the referenced wire in the constraint must be prunable"
178		);
179		self.ir.private_wires_status[wire_idx] = WireStatus::Pruned;
180		let uses = mem::take(&mut self.private_wire_uses[wire_idx]);
181
182		// Replace the eliminated wire in all use sites.
183		for use_site in uses {
184			let dst_operand = match use_site {
185				UseSite::Add { index } => &mut self.ir.zero_constraints[index as usize],
186				UseSite::Mul { index, position } => {
187					let constraint = &mut self.ir.mul_constraints[index as usize];
188					match position {
189						MulPosition::A => &mut constraint.a,
190						MulPosition::B => &mut constraint.b,
191						MulPosition::C => &mut constraint.c,
192					}
193				}
194			};
195
196			let (additions, removals) = dst_operand.merge(&operand);
197			for wire in additions.wires() {
198				if matches!(wire.kind, WireKind::Private)
199					&& matches!(self.ir.private_wires_status[wire.id as usize], WireStatus::Unknown)
200				{
201					self.private_wire_uses[wire.id as usize].push(use_site.clone());
202				}
203			}
204			for wire in removals.wires() {
205				if matches!(wire.kind, WireKind::Private)
206					&& matches!(self.ir.private_wires_status[wire.id as usize], WireStatus::Unknown)
207				{
208					remove_use(&mut self.private_wire_uses[wire.id as usize], &use_site)
209						.expect("invariant: uses are kept in sync by algorithm invariant");
210				}
211			}
212		}
213	}
214
215	fn finish(self) -> ConstraintSystemIR<F> {
216		// Status is already maintained in self.ir.private_wires_status
217		self.ir
218	}
219}
220
221#[cfg(test)]
222mod tests {
223	use std::{iter, iter::successors};
224
225	use binius_field::{Field, Ghash128b as B128, PackedField};
226
227	use super::*;
228	use crate::circuit_builder::{CircuitBuilder, ConstraintBuilder, WitnessGenerator};
229
230	fn fibonacci<Builder: CircuitBuilder>(
231		builder: &mut Builder,
232		x0: Builder::Wire,
233		x1: Builder::Wire,
234		n: usize,
235	) -> Builder::Wire {
236		if n == 0 {
237			return x0;
238		}
239
240		let (_xnsub1, xn) = successors(Some((x0, x1)), |&(a, b)| {
241			let next = builder.mul(a, b);
242			Some((b, next))
243		})
244		.nth(n - 1)
245		.expect("closure always returns Some");
246
247		xn
248	}
249
250	#[test]
251	fn test_wire_elimination_fibonacci() {
252		// Build constraint system for fibonacci(20). Inputs are precommit (secret) wires so the
253		// mul chain stays private and exercises wire elimination — all-public inputs would instead
254		// be elided into derived wires, leaving nothing to eliminate.
255		let mut constraint_builder = ConstraintBuilder::new();
256		let x0 = constraint_builder.alloc_precommit();
257		let x1 = constraint_builder.alloc_precommit();
258		let xn = constraint_builder.alloc_precommit();
259		let out = fibonacci(&mut constraint_builder, x0, x1, 20);
260		constraint_builder.assert_eq(out, xn);
261		let ir = constraint_builder.build();
262
263		let ir = run_wire_elimination(CostModel::default(), ir);
264		let (optimized_cs, layout) = ir.finalize();
265
266		// Generate witness for optimized constraint system
267		let mut witness_generator = WitnessGenerator::new(&layout);
268		let x0_val = witness_generator.write_precommit(x0, B128::ONE);
269		let x1_val = witness_generator.write_precommit(x1, B128::MULTIPLICATIVE_GENERATOR);
270		let xn_val =
271			witness_generator.write_precommit(xn, B128::MULTIPLICATIVE_GENERATOR.pow(6765));
272		let out_val = fibonacci(&mut witness_generator, x0_val, x1_val, 20);
273		witness_generator.assert_eq(out_val, xn_val);
274		let witness = witness_generator.build().unwrap();
275
276		// Validate witness against optimized constraint system
277		optimized_cs.validate(&witness);
278
279		// Verify that some optimization occurred
280		// fibonacci(20) creates 20 mul outputs, plus 1 add output from assert_eq
281		// After optimization, some should be eliminated
282		let n_private_after = optimized_cs.n_private();
283		let n_private_before = 21;
284		assert!(
285			n_private_after < n_private_before,
286			"Expected some private wires to be eliminated, got {} out of {}",
287			n_private_after,
288			n_private_before
289		);
290	}
291
292	#[test]
293	fn test_chain_of_adds() {
294		fn chain_adds<Builder: CircuitBuilder>(
295			builder: &mut Builder,
296			inputs: &[Builder::Wire],
297		) -> Builder::Wire {
298			let mut acc = inputs[0];
299			for &input in &inputs[1..] {
300				acc = builder.add(acc, input);
301			}
302			acc
303		}
304
305		let mut constraint_builder = ConstraintBuilder::new();
306
307		// Create 8 input wires and 1 sum wire. Precommit (secret) inputs keep the add chain private
308		// so wire elimination has something to optimize (all-public inputs become derived wires).
309		let inputs: Vec<_> = (0..8)
310			.map(|_| constraint_builder.alloc_precommit())
311			.collect();
312		let sum_wire = constraint_builder.alloc_precommit();
313
314		// Build constraint system
315		let result = chain_adds(&mut constraint_builder, &inputs);
316		constraint_builder.assert_eq(result, sum_wire);
317		let ir = constraint_builder.build();
318
319		let ir = run_wire_elimination(CostModel::default(), ir);
320		let (optimized_cs, layout) = ir.finalize();
321
322		// Generate test values
323		let input_values: Vec<_> = (0..8).map(|i| B128::new(1u128 << i)).collect();
324		let sum_value: B128 = input_values.iter().copied().sum();
325
326		// Generate witness
327		let mut witness_generator = WitnessGenerator::new(&layout);
328		let input_wires: Vec<_> = iter::zip(&inputs, &input_values)
329			.map(|(&wire, &value)| witness_generator.write_precommit(wire, value))
330			.collect();
331
332		let result = chain_adds(&mut witness_generator, &input_wires);
333		let sum = witness_generator.write_precommit(sum_wire, sum_value);
334		witness_generator.assert_eq(result, sum);
335		let witness = witness_generator.build().unwrap();
336
337		optimized_cs.validate(&witness);
338
339		// Verify optimization occurred
340		// 7 add operations create 7 private wires, plus 1 from assert_eq = 8 total
341		// After optimization, some should be eliminated
342		let n_private_after = optimized_cs.n_private();
343		let n_private_before = 8;
344		assert!(
345			n_private_after < n_private_before,
346			"Expected some private wires to be eliminated, got {} out of {}",
347			n_private_after,
348			n_private_before
349		);
350	}
351
352	#[test]
353	fn test_long_chain_of_adds() {
354		// This circuit computes a product of of cumulative partial sums of a sequence.
355		// It is designed so that long add terms are reused in multiplication constraints.
356		fn chain_add_muls<Builder: CircuitBuilder>(
357			builder: &mut Builder,
358			inputs: &[Builder::Wire],
359		) -> Builder::Wire {
360			let mut add_acc = inputs[0];
361			let mut mul_acc = inputs[0];
362			for &input in &inputs[1..] {
363				add_acc = builder.add(add_acc, input);
364				mul_acc = builder.mul(mul_acc, add_acc);
365			}
366			mul_acc
367		}
368
369		let mut constraint_builder = ConstraintBuilder::new();
370
371		// Create 40 input wires and 1 output wire. Precommit (secret) inputs keep the add/mul chain
372		// private so wire elimination operates on it; all-public inputs would be elided to derived
373		// wires with no constraints to optimize.
374		let inputs: Vec<_> = (0..40)
375			.map(|_| constraint_builder.alloc_precommit())
376			.collect();
377		let output = constraint_builder.alloc_precommit();
378
379		// Build constraint system
380		let result = chain_add_muls(&mut constraint_builder, &inputs);
381		constraint_builder.assert_eq(result, output);
382		let ir = constraint_builder.build();
383		let original_mul_constraint_count = ir.mul_constraints.len();
384
385		let ir = run_wire_elimination(CostModel::default(), ir);
386		let (optimized_cs, layout) = ir.finalize();
387		let optimized_mul_constraint_count = optimized_cs.mul_constraints().len();
388
389		// Assert some multiplication constraints were added in place of zero constraints.
390		// This is a strict inequality because not all zero constraints should get eliminated.
391		assert!(optimized_mul_constraint_count > original_mul_constraint_count);
392
393		// Generate test values
394		let input_values: Vec<_> = (0..40).map(|i| B128::new(1u128 << i)).collect();
395		let (_, output_value) =
396			input_values
397				.iter()
398				.fold((B128::ZERO, B128::ONE), |(add_acc, mul_acc), &value| {
399					let add_acc = add_acc + value;
400					let mul_acc = mul_acc * add_acc;
401					(add_acc, mul_acc)
402				});
403
404		// Generate witness
405		let mut witness_generator = WitnessGenerator::new(&layout);
406		let input_wires: Vec<_> = iter::zip(&inputs, &input_values)
407			.map(|(&wire, &value)| witness_generator.write_precommit(wire, value))
408			.collect();
409
410		let result = chain_add_muls(&mut witness_generator, &input_wires);
411		let sum = witness_generator.write_precommit(output, output_value);
412		witness_generator.assert_eq(result, sum);
413		let witness = witness_generator.build().unwrap();
414
415		optimized_cs.validate(&witness);
416
417		let max_mul_operand_len = optimized_cs
418			.mul_constraints()
419			.iter()
420			.map(|c| c.a.len())
421			.max()
422			.unwrap();
423		// Assert equality with empirically determined value. The important thing is that it's far
424		// less than 40 (the maximum addition chain length).
425		assert_eq!(max_mul_operand_len, 12);
426
427		// Verify optimization occurred
428		// 39 add operations + 39 mul operations + 1 from assert_eq = 79 private wires before
429		// After optimization, many should be eliminated
430		let n_private_after = optimized_cs.n_private();
431		let n_private_before = 79;
432		assert!(
433			n_private_after < n_private_before,
434			"Expected some private wires to be eliminated, got {} out of {}",
435			n_private_after,
436			n_private_before
437		);
438	}
439
440	#[test]
441	fn test_two_inout_equality() {
442		fn assert_equality<Builder: CircuitBuilder>(
443			builder: &mut Builder,
444			w0: Builder::Wire,
445			w1: Builder::Wire,
446		) {
447			builder.assert_eq(w0, w1);
448		}
449
450		let mut constraint_builder = ConstraintBuilder::new();
451
452		let w0 = constraint_builder.alloc_inout();
453		let w1 = constraint_builder.alloc_inout();
454
455		// Build constraint system
456		assert_equality(&mut constraint_builder, w0, w1);
457		let ir = constraint_builder.build();
458
459		let ir = run_wire_elimination(CostModel::default(), ir);
460		let (optimized_cs, layout) = ir.finalize();
461
462		// Generate witness
463		let value = B128::new(42);
464		let mut witness_generator = WitnessGenerator::new(&layout);
465		let w0_val = witness_generator.write_inout(w0, value);
466		let w1_val = witness_generator.write_inout(w1, value);
467		assert_equality(&mut witness_generator, w0_val, w1_val);
468		let witness = witness_generator.build().unwrap();
469
470		optimized_cs.validate(&witness);
471
472		// Verify no private wires were created
473		assert_eq!(optimized_cs.n_private(), 0, "Expected no private wires");
474	}
475
476	#[test]
477	fn test_grouped_adds_into_mul() {
478		fn grouped_adds_mul<Builder: CircuitBuilder>(
479			builder: &mut Builder,
480			inputs: &[Builder::Wire; 9],
481		) {
482			// Sum first 3 to get a
483			let a = builder.add(inputs[0], inputs[1]);
484			let a = builder.add(a, inputs[2]);
485
486			// Sum second 3 to get b
487			let b = builder.add(inputs[3], inputs[4]);
488			let b = builder.add(b, inputs[5]);
489
490			// Sum third 3 to get c
491			let c = builder.add(inputs[6], inputs[7]);
492			let c = builder.add(c, inputs[8]);
493
494			// Return (a * b, c) for assertion
495			let a_times_b = builder.mul(a, b);
496			builder.assert_eq(a_times_b, c);
497		}
498
499		let mut constraint_builder = ConstraintBuilder::new();
500
501		// Create 9 input wires. Precommit (secret) inputs keep the grouped adds and the mul private
502		// so wire elimination runs (all-public inputs would be elided to derived wires).
503		let inputs: Vec<_> = (0..9)
504			.map(|_| constraint_builder.alloc_precommit())
505			.collect();
506
507		// Build constraint system: assert a * b = c where a, b, c are sums of 3 inputs each
508		let inputs_array: [_; 9] = inputs.clone().try_into().unwrap();
509		grouped_adds_mul(&mut constraint_builder, &inputs_array);
510		let ir = constraint_builder.build();
511
512		let ir = run_wire_elimination(CostModel::default(), ir);
513		let (optimized_cs, layout) = ir.finalize();
514
515		// Generate test values: a = 2, b = 3, c = 6
516		// In binary field: 2 * 3 = 6
517		let input_values = vec![
518			B128::new(2),
519			B128::ZERO,
520			B128::ZERO, // a = 2
521			B128::new(3),
522			B128::ZERO,
523			B128::ZERO, // b = 3
524			B128::new(6),
525			B128::ZERO,
526			B128::ZERO, // c = 6
527		];
528
529		// Generate witness
530		let mut witness_generator = WitnessGenerator::new(&layout);
531		let input_wires: Vec<_> = iter::zip(&inputs, &input_values)
532			.map(|(&wire, &value)| witness_generator.write_precommit(wire, value))
533			.collect();
534
535		let input_wires_array: [_; 9] = input_wires.try_into().unwrap();
536		grouped_adds_mul(&mut witness_generator, &input_wires_array);
537		let witness = witness_generator.build().unwrap();
538
539		optimized_cs.validate(&witness);
540
541		// Verify optimization occurred
542		// 6 add operations + 1 mul operation = 7 private wires before
543		// After optimization, some should be eliminated
544		let n_private_after = optimized_cs.n_private();
545		let n_private_before = 7;
546		assert!(
547			n_private_after < n_private_before,
548			"Expected some private wires to be eliminated, got {} out of {}",
549			n_private_after,
550			n_private_before
551		);
552	}
553}