binius_circuits/bignum/
mul.rs1use std::iter;
3
4use binius_core::word::Word;
5use binius_frontend::{CircuitBuilder, Wire};
6
7use super::{
8 addsub::{add_with_carry_out, compute_stack_adds, compute_stack_adds_with_carry_outs, sub},
9 biguint::BigUint,
10};
11
12pub fn textbook_mul(builder: &CircuitBuilder, a: &BigUint, b: &BigUint) -> BigUint {
28 let mut accumulator = vec![vec![]; a.limbs.len() + b.limbs.len()];
33 for (i, &ai) in a.limbs.iter().enumerate() {
34 for (j, &bj) in b.limbs.iter().enumerate() {
35 let (hi, lo) = builder.imul(ai, bj);
36 let k = i + j;
37 accumulator[k].push(lo);
38 accumulator[k + 1].push(hi);
39 }
40 }
41 compute_stack_adds(builder, &accumulator)
42}
43
44pub fn textbook_square(builder: &CircuitBuilder, a: &BigUint) -> BigUint {
57 let mut accumulator = vec![vec![]; a.limbs.len() + a.limbs.len()];
58 for (i, &ai) in a.limbs.iter().enumerate() {
59 for (j, &aj) in a.limbs.iter().enumerate().skip(i) {
60 let (hi, lo) = builder.imul(ai, aj);
61 accumulator[i + j].push(lo);
62 accumulator[i + j + 1].push(hi);
63 if i != j {
64 accumulator[i + j].push(lo);
66 accumulator[i + j + 1].push(hi);
67 }
68 }
69 }
70 compute_stack_adds(builder, &accumulator)
71}
72
73const KARATSUBA_LIMBS_THRESHOLD: usize = 8;
75
76pub fn optimal_mul(builder: &CircuitBuilder, a: &BigUint, b: &BigUint) -> BigUint {
92 let n = a.limbs.len();
93 if n == b.limbs.len() && n.is_power_of_two() && n >= KARATSUBA_LIMBS_THRESHOLD {
94 karatsuba_mul(builder, a, b)
95 } else {
96 textbook_mul(builder, a, b)
97 }
98}
99
100pub fn optimal_sqr(builder: &CircuitBuilder, a: &BigUint) -> BigUint {
111 let n = a.limbs.len();
112 if n.is_power_of_two() && n >= KARATSUBA_LIMBS_THRESHOLD {
113 karatsuba_mul(builder, a, a)
114 } else {
115 textbook_square(builder, a)
116 }
117}
118
119pub fn karatsuba_mul(builder: &CircuitBuilder, a: &BigUint, b: &BigUint) -> BigUint {
137 let n = a.limbs.len();
138
139 assert!(n.is_power_of_two());
141 assert_eq!(b.limbs.len(), n);
142
143 if n == 1 {
144 let (hi, lo) = builder.imul(a.limbs[0], b.limbs[0]);
146 return BigUint {
147 limbs: vec![lo, hi],
148 };
149 }
150
151 let n_half = n >> 1;
155 let (a_0, a_inf) = a.clone().split_at_limbs(n_half);
156 let (b_0, b_inf) = b.clone().split_at_limbs(n_half);
157
158 let ab_0 = karatsuba_mul(builder, &a_0, &b_0);
160 let ab_inf = karatsuba_mul(builder, &a_inf, &b_inf);
161
162 let (a_1, a_1_carry_out) = add_with_carry_out(builder, &a_0, &a_inf);
164 let (b_1, b_1_carry_out) = add_with_carry_out(builder, &b_0, &b_inf);
165
166 let ab_1 = karatsuba_mul(builder, &a_1, &b_1);
168
169 let ab_1_carry_out =
171 builder.shr(builder.band(a_1_carry_out, b_1_carry_out), (Word::BITS - 1) as u32);
172
173 let mut product_stacks = vec![Vec::new(); 2 * n];
175 add_to_stacks(&mut product_stacks[0..], &ab_0);
183 add_to_stacks(&mut product_stacks[n..], &ab_inf);
184 add_to_stacks(&mut product_stacks[n_half..], &ab_1);
185 add_to_stacks(&mut product_stacks[n..], &a_1.zero_unless(builder, b_1_carry_out));
186 add_to_stacks(&mut product_stacks[n..], &b_1.zero_unless(builder, a_1_carry_out));
187 add_to_stacks(
188 &mut product_stacks[n + n_half..],
189 &BigUint {
190 limbs: [ab_1_carry_out].to_vec(),
191 },
192 );
193 let (mut overflow_sum, carry_outs) =
194 compute_stack_adds_with_carry_outs(builder, &product_stacks);
195
196 let carry_out = carry_outs
198 .into_iter()
199 .reduce(|lhs, rhs| builder.bor(lhs, rhs))
200 .expect("at least one carry out");
201 overflow_sum
202 .limbs
203 .push(builder.shr(carry_out, (Word::BITS - 1) as u32));
204
205 let (ab_lo_invariant, ab_hi_overflow) = overflow_sum.split_at_limbs(n_half);
208 let remainder_len = ab_hi_overflow.limbs.len();
209 let ab_0_padded = ab_0.zero_extend(builder, remainder_len);
210 let ab_inf_padded = ab_inf.zero_extend(builder, remainder_len);
211
212 let ab_hi_overflow = sub(builder, &sub(builder, &ab_hi_overflow, &ab_0_padded), &ab_inf_padded);
213 let (ab_hi, overflow_limb) = ab_hi_overflow.split_at_limbs(2 * n - n_half);
214
215 assert_eq!(overflow_limb.limbs.len(), 1);
217 builder.assert_zero("karatsuba_mul_overflow_limb", overflow_limb.limbs[0]);
218
219 ab_lo_invariant.concat_limbs(&ab_hi)
220}
221
222fn add_to_stacks(limb_stacks: &mut [Vec<Wire>], a: &BigUint) {
223 assert!(limb_stacks.len() >= a.limbs.len());
224 for (limb_stack, &limb) in iter::zip(limb_stacks, &a.limbs) {
225 limb_stack.push(limb);
226 }
227}