binius_circuits/bignum/
mod_divide.rs1use binius_core::Word;
5use binius_frontend::{CircuitBuilder, Wire, hints::Hint};
6
7use super::num_biguint_from_u64_limbs;
8
9pub struct ModDivideHint;
16
17impl ModDivideHint {
18 pub const fn new() -> Self {
19 Self
20 }
21
22 pub fn call(
34 builder: &CircuitBuilder,
35 dividend: &[Wire],
36 divisor: &[Wire],
37 modulus: &[Wire],
38 ) -> (Vec<Wire>, Vec<Wire>) {
39 let inputs: Vec<Wire> = dividend
40 .iter()
41 .chain(divisor)
42 .chain(modulus)
43 .copied()
44 .collect();
45 let mut out = builder.call_hint(
46 Self::new(),
47 &[dividend.len(), divisor.len(), modulus.len()],
48 &inputs,
49 );
50 let slope = out.split_off(modulus.len());
51 (out, slope)
52 }
53}
54
55impl Default for ModDivideHint {
56 fn default() -> Self {
57 Self::new()
58 }
59}
60
61impl Hint for ModDivideHint {
62 const NAME: &'static str = "binius.mod_divide";
63
64 fn shape(&self, dimensions: &[usize]) -> (usize, usize) {
65 let [dividend_limbs, divisor_limbs, mod_limbs] = dimensions else {
66 panic!("ModDivide requires 3 dimensions");
67 };
68 (*dividend_limbs + *divisor_limbs + *mod_limbs, 2 * *mod_limbs)
69 }
70
71 fn execute(&self, dimensions: &[usize], inputs: &[Word], outputs: &mut [Word]) {
72 let [n_dividend, n_divisor, n_mod] = dimensions else {
73 panic!("ModDivide requires 3 dimensions");
74 };
75
76 let dividend_limbs = &inputs[0..*n_dividend];
77 let divisor_limbs = &inputs[*n_dividend..*n_dividend + *n_divisor];
78 let mod_limbs = &inputs[*n_dividend + *n_divisor..];
79
80 let dividend = num_biguint_from_u64_limbs(dividend_limbs.iter().map(|w| w.as_u64()));
81 let divisor = num_biguint_from_u64_limbs(divisor_limbs.iter().map(|w| w.as_u64()));
82 let modulus = num_biguint_from_u64_limbs(mod_limbs.iter().map(|w| w.as_u64()));
83
84 let zero = num_bigint::BigUint::ZERO;
85 let (quotient, slope) = if let Some(inverse) = divisor.modinv(&modulus) {
86 let slope = (÷nd * &inverse) % &modulus;
87 let numerator = &slope * &divisor;
88 let quotient = if numerator >= dividend {
94 (numerator - ÷nd) / &modulus
95 } else {
96 zero
97 };
98 (quotient, slope)
99 } else {
100 (zero.clone(), zero)
101 };
102
103 assert_eq!(outputs.len(), 2 * *n_mod);
104 let (quotient_words, slope_words) = outputs.split_at_mut(*n_mod);
105
106 for (i, limb) in quotient.iter_u64_digits().enumerate() {
109 if i < *n_mod {
110 quotient_words[i] = Word::from_u64(limb);
111 }
112 }
113 for i in quotient.iter_u64_digits().len()..*n_mod {
114 quotient_words[i] = Word::ZERO;
115 }
116
117 for (i, limb) in slope.iter_u64_digits().enumerate() {
119 if i < *n_mod {
120 slope_words[i] = Word::from_u64(limb);
121 }
122 }
123 for i in slope.iter_u64_digits().len()..*n_mod {
124 slope_words[i] = Word::ZERO;
125 }
126 }
127}