1use 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 CostModel {
44 wire_cost: 16,
45 mul_cost: 2,
46 ref_cost: 1,
47 }
48 }
49}
50
51pub struct WireEliminationPass<F: Field> {
59 cost_model: CostModel,
60 ir: ConstraintSystemIR<F>,
61 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 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 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 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 let operand = mem::take(&mut self.ir.zero_constraints[constraint_idx]);
159
160 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 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 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 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 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 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 optimized_cs.validate(&witness);
278
279 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 let inputs: Vec<_> = (0..8)
310 .map(|_| constraint_builder.alloc_precommit())
311 .collect();
312 let sum_wire = constraint_builder.alloc_precommit();
313
314 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 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 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 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 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 let inputs: Vec<_> = (0..40)
375 .map(|_| constraint_builder.alloc_precommit())
376 .collect();
377 let output = constraint_builder.alloc_precommit();
378
379 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!(optimized_mul_constraint_count > original_mul_constraint_count);
392
393 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 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_eq!(max_mul_operand_len, 12);
426
427 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 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 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 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 let a = builder.add(inputs[0], inputs[1]);
484 let a = builder.add(a, inputs[2]);
485
486 let b = builder.add(inputs[3], inputs[4]);
488 let b = builder.add(b, inputs[5]);
489
490 let c = builder.add(inputs[6], inputs[7]);
492 let c = builder.add(c, inputs[8]);
493
494 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 let inputs: Vec<_> = (0..9)
504 .map(|_| constraint_builder.alloc_precommit())
505 .collect();
506
507 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 let input_values = vec![
518 B128::new(2),
519 B128::ZERO,
520 B128::ZERO, B128::new(3),
522 B128::ZERO,
523 B128::ZERO, B128::new(6),
525 B128::ZERO,
526 B128::ZERO, ];
528
529 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 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}