Skip to main content

binius_circuits/bignum/
mul.rs

1// Copyright 2025 Irreducible Inc.
2use 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
12/// Multiply two arbitrary-sized `BigUint`s using textbook algorithm.
13///
14/// Produces `O(a.len() * b.len())` constraints.
15///
16/// Computes `a * b` where both inputs are `BigUint`s. The result will have
17/// `a.limbs.len() + b.limbs.len()` limbs to accommodate the full product
18/// without overflow.
19///
20/// # Arguments
21/// * `builder` - Circuit builder for constraint generation
22/// * `a` - First operand `BigUint`
23/// * `b` - Second operand `BigUint`
24///
25/// # Returns
26/// Product `BigUint` with `a.limbs.len() + b.limbs.len()` limbs
27pub fn textbook_mul(builder: &CircuitBuilder, a: &BigUint, b: &BigUint) -> BigUint {
28	// Multiply argument's limbs pairwise.
29	//
30	// The accumulator has exactly a.limbs.len() + b.limbs.len() slots to hold
31	// all partial products
32	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
44/// Square an arbitrary-sized `BigUint` using textbook algorithm.
45///
46/// Computes `a * a` using an optimized algorithm that takes advantage of the symmetry
47/// in squaring (each cross-product appears twice). This is roughly twice more efficient
48/// than using `textbook_mul`, though still quadratic in the number of constraints.
49///
50/// # Arguments
51/// * `builder` - Circuit builder for constraint generation
52/// * `a` - The `BigUint` to be squared
53///
54/// # Returns
55/// The square of `a` as a `BigUint` with `2 * a.limbs.len()` limbs
56pub 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				// Off-diagonal elements appear twice
65				accumulator[i + j].push(lo);
66				accumulator[i + j + 1].push(hi);
67			}
68		}
69	}
70	compute_stack_adds(builder, &accumulator)
71}
72
73/// BigUint size at which Karatsuba becomes better than textbook multiplication.
74const KARATSUBA_LIMBS_THRESHOLD: usize = 8;
75
76/// Multiply two arbitrary-sized `BigUint`s using textbook algorithm.
77///
78/// This method attempts to pick the most efficient multiplication algorithm.
79///
80/// Computes `a * b` where both inputs are `BigUint`s. The result will have
81/// `a.limbs.len() + b.limbs.len()` limbs to accommodate the full product
82/// without overflow.
83///
84/// # Arguments
85/// * `builder` - Circuit builder for constraint generation
86/// * `a` - First operand `BigUint`
87/// * `b` - Second operand `BigUint`
88///
89/// # Returns
90/// Product `BigUint` with `a.limbs.len() + b.limbs.len()` limbs
91pub 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
100/// Square an arbitrary-sized `BigUint`.
101///
102/// This method attempts to pick the most efficient multiplication algorithm.
103///
104/// # Arguments
105/// * `builder` - Circuit builder for constraint generation
106/// * `a` - The `BigUint` to be squared
107///
108/// # Returns
109/// The square of `a` as a `BigUint` with `2 * a.limbs.len()` limbs
110pub 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
119/// Multiply two `BigUint`s with po2 number of limbs using Karatsuba (aka Toom-22).
120///
121/// Whereas `textbook_mul` and `textbook_square` require $O(n^2)$ constraints for $n$ limbs,
122/// this method is asymptotically more efficient with $O(n^{log_2 3}) = O(n^{1.58})$,
123/// however due to larger constant factor it's beneficial for longer `BigUint`s only.
124///
125/// Computes `a * b` where both inputs are `BigUint`s. The result will have
126/// `a.limbs.len() + b.limbs.len()` limbs to accommodate the full product
127/// without overflow.
128///
129/// # Arguments
130/// * `builder` - Circuit builder for constraint generation
131/// * `a` - First operand `BigUint`
132/// * `b` - Second operand `BigUint`
133///
134/// # Returns
135/// Product `BigUint` with `a.limbs.len() + b.limbs.len()` limbs
136pub fn karatsuba_mul(builder: &CircuitBuilder, a: &BigUint, b: &BigUint) -> BigUint {
137	let n = a.limbs.len();
138
139	// Preconditions
140	assert!(n.is_power_of_two());
141	assert_eq!(b.limbs.len(), n);
142
143	if n == 1 {
144		// Base case
145		let (hi, lo) = builder.imul(a.limbs[0], b.limbs[0]);
146		return BigUint {
147			limbs: vec![lo, hi],
148		};
149	}
150
151	// a(t) = a_0 + t * a_∞
152	// b(t) = b_0 + t * b_∞
153	// for t = 2^(Word::BITS*n/2), a = a(t), b = b(t)
154	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	// Recursively multiply (a_0, b_0) and (a_∞, b_∞)
159	let ab_0 = karatsuba_mul(builder, &a_0, &b_0);
160	let ab_inf = karatsuba_mul(builder, &a_inf, &b_inf);
161
162	// Compute a(1) = a_0 + a_∞, b(1) = b_0 + b_∞, split into 2*n limbs and carry out bit
163	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	// Recursively multiply a(1) and b(1), _without_ carry out bits
167	let ab_1 = karatsuba_mul(builder, &a_1, &b_1);
168
169	// Multiply a(1) and b(1) carry out bits only
170	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	// Addition chain for overflow_sum = ab_0 + ab_inf * 2^n + ab_1 * 2^n/2
174	let mut product_stacks = vec![Vec::new(); 2 * n];
175	// 0 ... n/2 ... n ... 3n/2 ... 2n
176	// [    ab_0     )
177	//               [    ab_inf     )
178	//       [    ab_1      )
179	//               [ a_1  )     -- if b_1 carry out not zero
180	//               [ b_1  )     -- if a_1 carry out not zero
181	//                      ab_1_carry_out
182	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	// overflow_sum needs 2n+1 limbs, highest limb is 0 or 1, so we can bit-or them together
197	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	// now we need to subtract (ab_0 + ab_inf) * 2^n/2, separate the lower limbs that won't get
206	// affected
207	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	// The extra limb of overflow_limb should be zero, because product fits into 2n limbs
216	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}