Skip to main content

binius_math/
inner_product.rs

1// Copyright 2024-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Inner products `sum_i a_i * b_i` over sequences of field elements.
5
6use std::{iter, ops::Deref};
7
8use binius_field::{ExtensionField, Field, FieldOps, PackedField, WideMul};
9use binius_utils::rayon::{
10	prelude::*,
11	task_size::{IndexedParallelIteratorExt, WorkPerItem},
12};
13
14use crate::FieldBuffer;
15
16/// Computes the inner product `sum_i a_i * b_i`.
17///
18/// Arithmetic is element-wise.
19/// A packed element therefore yields one independent inner product per lane, not one scalar.
20///
21/// # Panics
22///
23/// Panics if the two sequences have different lengths.
24#[inline]
25pub fn inner_product<F: FieldOps>(
26	a: impl IntoIterator<Item = F>,
27	b: impl IntoIterator<Item = F>,
28) -> F {
29	itertools::zip_eq(a, b).map(|(a_i, b_i)| b_i * a_i).sum()
30}
31
32/// Computes the inner product `sum_i a_i * b_i`.
33///
34/// The left side lies in a subfield of the right side's field.
35///
36/// # Panics
37///
38/// Panics if the two sequences have different lengths.
39#[inline]
40pub fn inner_product_subfield<F, FSub>(
41	a: impl IntoIterator<Item = FSub>,
42	b: impl IntoIterator<Item = F>,
43) -> F
44where
45	F: Field + ExtensionField<FSub>,
46	FSub: Field,
47{
48	itertools::zip_eq(a, b).map(|(a_i, b_i)| b_i * a_i).sum()
49}
50
51/// Sums the coefficient-by-coefficient products of two multilinears.
52///
53/// ```text
54/// result = sum_i a_i * b_i
55/// ```
56///
57/// This is not the product polynomial; it is a single field element.
58/// Pairing a polynomial with an equality indicator expansion evaluates it at that point.
59///
60/// ## Preconditions
61///
62/// * the two buffers must hold the same number of coefficients
63#[inline]
64pub fn inner_product_buffers<F, P, DataA, DataB>(
65	a: &FieldBuffer<P, DataA>,
66	b: &FieldBuffer<P, DataB>,
67) -> F
68where
69	F: Field,
70	P: PackedField<Scalar = F>,
71	DataA: Deref<Target = [P]>,
72	DataB: Deref<Target = [P]>,
73{
74	// Both buffers hold the same number of packed words, so the packed form applies directly.
75	inner_product_packed(a.log_len(), a.iter_packed().copied(), b.iter_packed().copied())
76}
77
78/// Sums the coefficient-by-coefficient products of two multilinears, across threads.
79///
80/// The value matches the single-threaded pairing.
81/// The coefficient range is split into tasks, each summing its own partial products.
82///
83/// ## Preconditions
84///
85/// * the two buffers must hold the same number of coefficients
86#[inline]
87pub fn inner_product_par<F, P, DataA, DataB>(
88	a: &FieldBuffer<P, DataA>,
89	b: &FieldBuffer<P, DataB>,
90) -> F
91where
92	F: Field,
93	P: PackedField<Scalar = F>,
94	DataA: Deref<Target = [P]>,
95	DataB: Deref<Target = [P]>,
96{
97	// Below one packed word the buffer still occupies a whole word, so only the live lanes count.
98	let n = a.len();
99
100	// Accumulate the products in unreduced (wide) form and reduce a single time at the end.
101	// For packed `GF(2^128)` fields this amortizes the reduction cost across all products.
102	// For every other field the widening multiply is the trivial one, so this is a plain
103	// product-then-sum.
104	let wide_sum = a
105		.as_ref()
106		.par_iter()
107		.zip_eq(b.as_ref().par_iter())
108		.with_min_task(WorkPerItem::FieldMuls)
109		.map(|(&a_i, &b_i)| P::wide_mul(a_i, b_i))
110		.sum::<<P as WideMul>::Output>();
111	P::reduce(wide_sum).into_iter().take(n).sum()
112}
113
114/// Computes the inner product of two scalar sequences generated by iterators of packed elements.
115///
116/// ## Preconditions
117///
118/// * `a` and `b` have length `1 << log_n.saturating_sub(P::LOG_WIDTH)`
119#[inline]
120pub fn inner_product_packed<F, P>(
121	log_n: usize,
122	a: impl ExactSizeIterator<Item = P>,
123	b: impl ExactSizeIterator<Item = P>,
124) -> F
125where
126	F: Field,
127	P: PackedField<Scalar = F>,
128{
129	assert_eq!(a.len(), 1 << log_n.saturating_sub(P::LOG_WIDTH)); // pre-condition
130	assert_eq!(b.len(), 1 << log_n.saturating_sub(P::LOG_WIDTH)); // pre-condition
131
132	// Accumulate the products in unreduced (wide) form and reduce a single time at the end.
133	// For packed `GF(2^128)` fields this amortizes the reduction cost across all products.
134	// For every other field the widening multiply is the trivial one, so this is a plain
135	// product-then-sum.
136	let wide_sum = iter::zip(a, b)
137		.map(|(a_i, b_i)| P::wide_mul(a_i, b_i))
138		.sum::<<P as WideMul>::Output>();
139	P::reduce(wide_sum).into_iter().take(1 << log_n).sum()
140}
141
142#[cfg(test)]
143mod tests {
144	use binius_field::{Divisible, Ghash128b, PackedField, PackedGhash4x128b, Random};
145	use proptest::prelude::*;
146	use rand::{SeedableRng, rngs::StdRng};
147
148	use super::*;
149	use crate::test_utils::random_field_buffer;
150
151	type P = PackedGhash4x128b;
152	type F = Ghash128b;
153
154	// The packing width is four scalars, so this range straddles it in both directions.
155	const MAX_VARS: usize = 8;
156
157	proptest! {
158		#[test]
159		fn the_two_buffer_inner_products_agree_with_the_scalar_reference(
160			n_vars in 0..=MAX_VARS,
161			seed: u64,
162		) {
163			let mut rng = StdRng::seed_from_u64(seed);
164			let a = random_field_buffer::<P>(&mut rng, n_vars);
165			let b = random_field_buffer::<P>(&mut rng, n_vars);
166
167			// Both forms defer the field reduction, so a scalar sum is the reference.
168			let reference = (0..a.len()).map(|i| a.get(i) * b.get(i)).sum::<F>();
169
170			prop_assert_eq!(inner_product_buffers(&a, &b), reference);
171			prop_assert_eq!(inner_product_par(&a, &b), reference);
172		}
173	}
174
175	#[test]
176	fn test_inner_product_over_packed_is_lane_wise() {
177		let mut rng = StdRng::seed_from_u64(7);
178
179		// Invariant: a packed element is a vector of independent lanes.
180		// Multiplication and addition never cross a lane boundary.
181		// So lane `j` of the result depends only on lane `j` of the two inputs.
182		//
183		// Fixture state: 8 packed elements per side, 4 lanes each.
184		// Every lane therefore carries its own sequence of 8 scalars.
185		//
186		//     a: [ a_0[0] | a_0[1] | a_0[2] | a_0[3] ], [ a_1[0] | ... ], ...
187		//     b: [ b_0[0] | b_0[1] | b_0[2] | b_0[3] ], [ b_1[0] | ... ], ...
188		//     -> lane j = sum_i a_i[j] * b_i[j]
189		let a = (0..8).map(|_| P::random(&mut rng)).collect::<Vec<P>>();
190		let b = (0..8).map(|_| P::random(&mut rng)).collect::<Vec<P>>();
191
192		let packed = inner_product(a.iter().copied(), b.iter().copied());
193
194		// Rebuild each lane's own scalar inner product and compare it against that lane.
195		for lane in 0..P::WIDTH {
196			let expected = iter::zip(&a, &b)
197				.map(|(a_i, b_i)| a_i.get(lane) * b_i.get(lane))
198				.sum::<F>();
199			assert_eq!(packed.get(lane), expected, "mismatch in lane {lane}");
200		}
201	}
202
203	#[test]
204	fn test_inner_product_matches_subfield_at_the_full_field() {
205		let mut rng = StdRng::seed_from_u64(11);
206
207		// Invariant: every field is a degree-1 extension of itself.
208		// The same-field form and the subfield form therefore agree on that overlap.
209		//
210		// Fixture state: 16 random elements per side, drawn from the same field.
211		let a = (0..16).map(|_| F::random(&mut rng)).collect::<Vec<F>>();
212		let b = (0..16).map(|_| F::random(&mut rng)).collect::<Vec<F>>();
213
214		// The turbofish forces the subfield to be the full field, which is the overlap.
215		// Pin the two forms equal there so their separate bodies cannot drift apart.
216		assert_eq!(
217			inner_product(a.iter().copied(), b.iter().copied()),
218			inner_product_subfield::<F, F>(a.iter().copied(), b.iter().copied()),
219		);
220	}
221
222	#[test]
223	fn test_inner_product_packed_matches_naive() {
224		let mut rng = StdRng::seed_from_u64(42);
225
226		// The packed inner product defers the `GF(2^128)` reduction (widening multiply).
227		// Check the reduced result against a naive scalar-by-scalar reference.
228		for log_n in [4, 8, 12] {
229			let n = 1 << log_n;
230			let a = (0..n / P::WIDTH)
231				.map(|_| P::random(&mut rng))
232				.collect::<Vec<P>>();
233			let b = (0..n / P::WIDTH)
234				.map(|_| P::random(&mut rng))
235				.collect::<Vec<P>>();
236
237			let naive = iter::zip(&a, &b)
238				.flat_map(|(a_i, b_i)| iter::zip(a_i.iter(), b_i.iter()))
239				.map(|(a_j, b_j)| a_j * b_j)
240				.sum::<F>();
241
242			let packed: F = inner_product_packed(log_n, a.iter().copied(), b.iter().copied());
243			assert_eq!(packed, naive, "mismatch at log_n={log_n}");
244		}
245	}
246
247	#[test]
248	fn test_inner_product_packed_below_the_packing_width() {
249		let mut rng = StdRng::seed_from_u64(0);
250
251		// Below the packing width the trailing scalars of the one word must not contribute.
252		let a = P::random(&mut rng);
253		let b = P::random(&mut rng);
254
255		let packed: F = inner_product_packed(0, iter::once(a), iter::once(b));
256		assert_eq!(packed, a.iter().next().unwrap() * b.iter().next().unwrap());
257	}
258}