binius_circuits/bignum/
reduce.rs1use binius_core::word::Word;
3use binius_frontend::{CircuitBuilder, Wire};
4
5use super::{
6 addsub::{add, sub},
7 biguint::{BigUint, assert_eq, assert_eq_cond},
8 mul::{optimal_mul, textbook_mul},
9};
10
11pub struct ModReduce {
17 pub a: BigUint,
18 pub modulus: BigUint,
19 pub quotient: BigUint,
20 pub remainder: BigUint,
21}
22
23impl ModReduce {
24 pub fn new(
36 builder: &CircuitBuilder,
37 a: BigUint,
38 modulus: BigUint,
39 quotient: BigUint,
40 remainder: BigUint,
41 ) -> Self {
42 let zero = builder.add_constant(Word::ZERO);
43
44 let product = optimal_mul(builder, "ient, &modulus);
45
46 let remainder_padded = remainder.pad_limbs_to(product.limbs.len(), zero);
47 let reconstructed = add(builder, &product, &remainder_padded);
48
49 let n_limbs = reconstructed.limbs.len().max(a.limbs.len());
50 assert_eq(
51 builder,
52 "modreduce_a_eq_reconstructed",
53 &reconstructed.pad_limbs_to(n_limbs, zero),
54 &a.pad_limbs_to(n_limbs, zero),
55 );
56
57 ModReduce {
58 a,
59 modulus,
60 quotient,
61 remainder,
62 }
63 }
64}
65
66pub struct PseudoMersenneModReduce {
78 lhs: BigUint,
79 rhs: BigUint,
80}
81
82impl PseudoMersenneModReduce {
83 #[must_use]
103 pub fn new(
104 builder: &CircuitBuilder,
105 a: &BigUint,
106 modulus_po2: usize,
107 modulus_subtrahend: &BigUint,
108 quotient: &BigUint,
109 remainder: &BigUint,
110 ) -> Self {
111 assert!(modulus_po2.is_multiple_of(Word::BITS));
118 assert!(modulus_subtrahend.limbs.len() * Word::BITS <= modulus_po2);
119 assert!(remainder.limbs.len() * Word::BITS <= modulus_po2);
120
121 let zero = builder.add_constant(Word::ZERO);
122
123 let n_lo_limbs = modulus_po2 / Word::BITS;
124
125 let (a_lo, a_hi) = a.pad_limbs_to(n_lo_limbs, zero).split_at_limbs(n_lo_limbs);
126
127 let rhs_hi = sub(builder, quotient, &a_hi.pad_limbs_to(quotient.limbs.len(), zero));
128 let rhs = remainder
129 .pad_limbs_to(n_lo_limbs, zero)
130 .concat_limbs(&rhs_hi);
131
132 let quotient_modulus_subtrahend = textbook_mul(builder, quotient, modulus_subtrahend);
133 let lhs_rhs_len = [
134 rhs.limbs.len(),
135 quotient_modulus_subtrahend.limbs.len() + 1,
136 a_lo.limbs.len() + 1,
137 ]
138 .into_iter()
139 .max()
140 .expect("exactly 3 elements");
141
142 let lhs = add(
143 builder,
144 &a_lo.pad_limbs_to(lhs_rhs_len, zero),
145 "ient_modulus_subtrahend.pad_limbs_to(lhs_rhs_len, zero),
146 );
147
148 let rhs = rhs.pad_limbs_to(lhs_rhs_len, zero);
149
150 Self { lhs, rhs }
151 }
152
153 pub fn constrain(self, builder: &CircuitBuilder) {
155 assert_eq(builder, "modreduce_pseudo_mersenne", &self.lhs, &self.rhs);
156 }
157
158 pub fn constrain_cond(self, builder: &CircuitBuilder, cond: Wire) {
160 assert_eq_cond(builder, "modred_pseudo_mersenne", &self.lhs, &self.rhs, cond);
161 }
162}