binius_circuits/bignum/
mod_inverse.rs1use binius_core::Word;
5use binius_frontend::{CircuitBuilder, Wire, hints::Hint};
6
7use super::num_biguint_from_u64_limbs;
8
9pub struct ModInverseHint;
11
12impl ModInverseHint {
13 pub const fn new() -> Self {
14 Self
15 }
16
17 pub fn call(
27 builder: &CircuitBuilder,
28 base: &[Wire],
29 modulus: &[Wire],
30 ) -> (Vec<Wire>, Vec<Wire>) {
31 let inputs: Vec<Wire> = base.iter().chain(modulus).copied().collect();
32 let mut out = builder.call_hint(Self::new(), &[base.len(), modulus.len()], &inputs);
33 let inverse = out.split_off(modulus.len());
34 (out, inverse)
35 }
36}
37
38impl Default for ModInverseHint {
39 fn default() -> Self {
40 Self::new()
41 }
42}
43
44impl Hint for ModInverseHint {
45 const NAME: &'static str = "binius.mod_inverse";
46
47 fn shape(&self, dimensions: &[usize]) -> (usize, usize) {
48 let [base_limbs, mod_limbs] = dimensions else {
49 panic!("ModInverse requires 2 dimensions");
50 };
51 (*base_limbs + *mod_limbs, 2 * *mod_limbs)
52 }
53
54 fn execute(&self, dimensions: &[usize], inputs: &[Word], outputs: &mut [Word]) {
55 let [n_base, n_mod] = dimensions else {
56 panic!("ModInverse requires 2 dimensions");
57 };
58
59 let base_limbs = &inputs[0..*n_base];
60 let mod_limbs = &inputs[*n_base..];
61
62 let base = num_biguint_from_u64_limbs(base_limbs.iter().map(|w| w.as_u64()));
63 let modulus = num_biguint_from_u64_limbs(mod_limbs.iter().map(|w| w.as_u64()));
64
65 let zero = num_bigint::BigUint::ZERO;
66 let (quotient, inverse) = base.modinv(&modulus).map_or_else(
67 || (zero.clone(), zero),
68 |inverse| {
69 let quotient = (base * &inverse - num_bigint::BigUint::from(1usize)) / &modulus;
70 (quotient, inverse)
71 },
72 );
73
74 assert_eq!(outputs.len(), 2 * *n_mod);
75 let (quotient_words, inverse_words) = outputs.split_at_mut(*n_mod);
76
77 for (i, limb) in quotient.iter_u64_digits().enumerate() {
79 quotient_words[i] = Word::from_u64(limb);
80 }
81
82 for i in quotient.iter_u64_digits().len()..*n_mod {
84 quotient_words[i] = Word::ZERO;
85 }
86
87 for (i, limb) in inverse.iter_u64_digits().enumerate() {
89 inverse_words[i] = Word::from_u64(limb);
90 }
91 for i in inverse.iter_u64_digits().len()..*n_mod {
93 inverse_words[i] = Word::ZERO;
94 }
95 }
96}
97
98#[cfg(test)]
99mod tests {
100 use super::*;
101
102 #[test]
103 fn test_mod_inverse_hint() {
104 let builder = CircuitBuilder::new();
105
106 let b = builder.add_constant_64(0x123456789abcdef0);
107
108 let m0 = builder.add_constant_64(u64::MAX);
110 let m1 = builder.add_constant_64((1u64 << 63) - 1);
111
112 let (quotient, inverse) = ModInverseHint::call(&builder, &[b], &[m0, m1]);
113
114 for &wire in quotient.iter().chain(&inverse) {
117 builder.mark_inout(wire);
118 }
119
120 let circuit = builder.build();
121 let mut w = circuit.new_witness_filler();
122 circuit.populate_wire_witness(&mut w).unwrap();
123
124 assert_eq!(inverse.len(), 2);
125 assert_eq!(w[inverse[0]], Word(0xe5a542e11f99750a));
126 assert_eq!(w[inverse[1]], Word(0x1849faf75fbb9752));
127
128 assert_eq!(quotient.len(), 2);
129 assert_eq!(w[quotient[0]], Word(0x37455c1554b9aa1));
130 assert_eq!(w[quotient[1]], Word::ZERO);
131 }
132
133 #[test]
134 fn test_mod_inverse_hint_non_coprime() {
135 let builder = CircuitBuilder::new();
136
137 let b = builder.add_constant_64((1 << 19) - 1);
138
139 let m0 = builder.add_constant_64(0xfffffffffff80001);
141 let m1 = builder.add_constant_64(0x3ffff7ffffffffff);
142
143 let (quotient, inverse) = ModInverseHint::call(&builder, &[b], &[m0, m1]);
144
145 for &wire in quotient.iter().chain(&inverse) {
148 builder.mark_inout(wire);
149 }
150
151 let circuit = builder.build();
152 let mut w = circuit.new_witness_filler();
153 circuit.populate_wire_witness(&mut w).unwrap();
154
155 assert_eq!(inverse.len(), 2);
156 assert_eq!(w[inverse[0]], Word::ZERO);
157 assert_eq!(w[inverse[1]], Word::ZERO);
158
159 assert_eq!(quotient.len(), 2);
160 assert_eq!(w[quotient[0]], Word::ZERO);
161 assert_eq!(w[quotient[1]], Word::ZERO);
162 }
163}