1use 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#[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#[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#[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 inner_product_packed(a.log_len(), a.iter_packed().copied(), b.iter_packed().copied())
76}
77
78#[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 let n = a.len();
99
100 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#[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)); assert_eq!(b.len(), 1 << log_n.saturating_sub(P::LOG_WIDTH)); 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 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 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 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 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 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 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 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 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}