1use binius_compute::{Allocator, BufferData, VecLike};
14use binius_field::{PackedField, field::FieldOps};
15
16use super::hypercube::{self, Hypercube, OneCube};
17use crate::{FieldBuffer, FieldVec};
18
19pub fn tensor_prod_eq_ind<P: PackedField>(
24 values: FieldBuffer<P, Vec<P>>,
25 extra_query_coordinates: &[P::Scalar],
26) -> FieldBuffer<P, Vec<P>> {
27 hypercube::tensor_prod_eq_ind::<OneCube, P>(values, extra_query_coordinates)
28}
29
30pub fn eq_ind_partial_eval<P: PackedField>(point: &[P::Scalar]) -> FieldBuffer<P> {
40 hypercube::eq_ind_partial_eval::<OneCube, P>(point)
41}
42
43pub fn eq_ind_partial_eval_in<A: Allocator, P: PackedField>(
47 alloc: &A,
48 point: &[P::Scalar],
49) -> FieldVec<P, A> {
50 hypercube::eq_ind_partial_eval_in::<OneCube, A, P>(alloc, point)
51}
52
53pub fn scaled_eq_ind_partial_eval<P: PackedField>(
63 point: &[P::Scalar],
64 scale: P::Scalar,
65) -> FieldBuffer<P> {
66 hypercube::scaled_eq_ind_partial_eval::<OneCube, P>(point, scale)
67}
68
69pub fn scaled_eq_ind_partial_eval_into<P: PackedField, Data: VecLike<P>>(
79 point: &[P::Scalar],
80 scale: P::Scalar,
81 buffer: Data,
82) -> FieldBuffer<P, Data> {
83 hypercube::scaled_eq_ind_partial_eval_into::<OneCube, P, Data>(point, scale, buffer)
84}
85
86pub fn eq_ind_truncate_low_inplace<P: PackedField, Data: BufferData<P>>(
98 values: &mut FieldBuffer<P, Data>,
99 truncated_log_len: usize,
100) {
101 hypercube::eq_ind_truncate_low_inplace::<OneCube, _, _>(values, truncated_log_len);
102}
103
104#[inline(always)]
116pub fn eq_one_var<F: FieldOps>(x: F, y: F) -> F {
117 OneCube::eq_one_var(x, y)
118}
119
120pub fn eq_ind<F: FieldOps>(x: &[F], y: &[F]) -> F {
128 hypercube::eq_ind::<OneCube, F>(x, y)
129}
130
131pub fn eq_ind_zero<F: FieldOps>(point: &[F]) -> F {
139 hypercube::eq_ind_zero::<OneCube, F>(point)
140}
141
142pub fn eq_ind_partial_eval_scalars<F: FieldOps>(point: &[F]) -> Vec<F> {
148 hypercube::eq_ind_partial_eval_scalars::<OneCube, F>(point)
149}
150
151pub fn scaled_eq_ind_partial_eval_scalars<F: FieldOps>(point: &[F], scale: F) -> Vec<F> {
156 hypercube::scaled_eq_ind_partial_eval_scalars::<OneCube, F>(point, scale)
157}
158
159#[cfg(test)]
160mod tests {
161 use binius_compute::GlobalAllocator;
162 use binius_field::Field;
163 use rand::prelude::*;
164
165 use super::*;
166 use crate::{
167 bit_reverse::bit_reverse_packed,
168 test_utils::{B128, Packed128b, index_to_hypercube_point, random_scalars},
169 };
170
171 type P = Packed128b;
172 type F = B128;
173
174 #[test]
175 fn expansion_holds_the_indicator_at_every_vertex() {
176 let mut rng = StdRng::seed_from_u64(0);
177
178 let n_vars = 5;
181 let point = random_scalars(&mut rng, n_vars);
182 let expansion = eq_ind_partial_eval::<P>(&point);
183
184 for index in 0..1 << n_vars {
185 let vertex = index_to_hypercube_point(n_vars, index);
186 assert_eq!(expansion.get(index), eq_ind::<F>(&point, &vertex));
187 }
188 }
189
190 #[test]
191 fn expansion_of_the_empty_point() {
192 let result = eq_ind_partial_eval::<P>(&[]);
194 assert_eq!(result.log_len(), 0);
195 assert_eq!(result.len(), 1);
196 assert_eq!(result.get(0), F::ONE);
197 }
198
199 #[test]
200 fn expansion_of_one_coordinate_is_the_basis() {
201 let r0 = F::new(2);
203 let result = eq_ind_partial_eval::<P>(&[r0]);
204 assert_eq!(result.log_len(), 1);
205 assert_eq!(result.len(), 2);
206 assert_eq!(result.get(0), F::ONE - r0);
207 assert_eq!(result.get(1), r0);
208 }
209
210 #[test]
211 fn expansion_of_two_coordinates() {
212 let r0 = F::new(2);
214 let r1 = F::new(3);
215 let result = eq_ind_partial_eval::<P>(&[r0, r1]);
216 assert_eq!(result.log_len(), 2);
217 assert_eq!(result.len(), 4);
218
219 let expected = vec![
221 (F::ONE - r0) * (F::ONE - r1),
222 r0 * (F::ONE - r1),
223 (F::ONE - r0) * r1,
224 r0 * r1,
225 ];
226 assert_eq!(result.iter_scalars().collect::<Vec<F>>(), expected);
227 }
228
229 #[test]
230 fn expansion_of_three_coordinates_fills_one_packed_word() {
231 let r0 = F::new(2);
233 let r1 = F::new(3);
234 let r2 = F::new(5);
235 let result = eq_ind_partial_eval::<P>(&[r0, r1, r2]);
236 assert_eq!(result.log_len(), 3);
237 assert_eq!(result.len(), 8);
238
239 let expected = vec![
240 (F::ONE - r0) * (F::ONE - r1) * (F::ONE - r2),
241 r0 * (F::ONE - r1) * (F::ONE - r2),
242 (F::ONE - r0) * r1 * (F::ONE - r2),
243 r0 * r1 * (F::ONE - r2),
244 (F::ONE - r0) * (F::ONE - r1) * r2,
245 r0 * (F::ONE - r1) * r2,
246 (F::ONE - r0) * r1 * r2,
247 r0 * r1 * r2,
248 ];
249 assert_eq!(result.iter_scalars().collect::<Vec<F>>(), expected);
250 }
251
252 #[test]
253 fn eq_ind_zero_is_the_product_of_complements() {
254 let mut rng = StdRng::seed_from_u64(0);
255
256 for n_vars in 0..5 {
258 let point = random_scalars::<F>(&mut rng, n_vars);
259 let expected: F = point.iter().map(|&r| F::ONE - r).product();
260 assert_eq!(eq_ind_zero(&point), expected);
261
262 assert_eq!(eq_ind_zero(&point), eq_ind(&vec![F::ZERO; n_vars], &point));
264 }
265 }
266
267 #[test]
268 fn every_storage_form_holds_the_same_values() {
269 let mut rng = StdRng::seed_from_u64(0);
270
271 for log_n in [0, 1, 2, 5, 8] {
277 let point = random_scalars::<F>(&mut rng, log_n);
278 let reference = eq_ind_partial_eval::<P>(&point);
279
280 let pooled = eq_ind_partial_eval_in::<_, P>(&GlobalAllocator, &point);
281 assert!(pooled.iter_scalars().eq(reference.iter_scalars()), "pool at log_n={log_n}");
282
283 let capacity = 1 << log_n.saturating_sub(P::LOG_WIDTH);
284 let supplied = scaled_eq_ind_partial_eval_into::<P, _>(
285 &point,
286 F::ONE,
287 Vec::with_capacity(capacity),
288 );
289 assert_eq!(supplied, reference, "supplied store at log_n={log_n}");
290
291 let scalars = eq_ind_partial_eval_scalars(&point);
292 assert!(reference.iter_scalars().eq(scalars), "scalars at log_n={log_n}");
293 }
294 }
295
296 #[test]
297 fn the_scale_applies_to_every_storage_form_alike() {
298 let mut rng = StdRng::seed_from_u64(1);
299
300 for log_n in [0, 1, 2, 5, 8] {
303 let point = random_scalars::<F>(&mut rng, log_n);
304 let scale = random_scalars::<F>(&mut rng, 1)[0];
305 let unscaled = eq_ind_partial_eval::<P>(&point);
306
307 let scaled = scaled_eq_ind_partial_eval::<P>(&point, scale);
308 for (got, base) in scaled.iter_scalars().zip(unscaled.iter_scalars()) {
309 assert_eq!(got, scale * base, "fresh store at log_n={log_n}");
310 }
311
312 let scalars = scaled_eq_ind_partial_eval_scalars(&point, scale);
313 assert!(scaled.iter_scalars().eq(scalars), "scalars at log_n={log_n}");
314 }
315 }
316
317 #[test]
318 fn a_scale_of_one_is_the_identity() {
319 let mut rng = StdRng::seed_from_u64(2);
320
321 for log_n in [0, 1, 2, 5, 8] {
324 let point = random_scalars::<F>(&mut rng, log_n);
325 assert_eq!(
326 scaled_eq_ind_partial_eval::<P>(&point, F::ONE),
327 eq_ind_partial_eval::<P>(&point),
328 "mismatch at log_n={log_n}"
329 );
330 }
331 }
332
333 #[test]
334 fn a_scale_of_zero_gives_all_zeros() {
335 let mut rng = StdRng::seed_from_u64(3);
336
337 for log_n in [0, 1, 2, 5] {
339 let point = random_scalars::<F>(&mut rng, log_n);
340 let scaled = scaled_eq_ind_partial_eval::<P>(&point, F::ZERO);
341 assert!(scaled.iter_scalars().all(|v| v == F::ZERO), "nonzero at log_n={log_n}");
342 }
343 }
344
345 #[test]
346 fn a_caller_reserved_store_matches_the_allocating_form() {
347 let mut rng = StdRng::seed_from_u64(5);
348
349 for log_n in [0, 1, 2, 5, 8] {
352 let point = random_scalars::<F>(&mut rng, log_n);
353 let scale = random_scalars::<F>(&mut rng, 1)[0];
354
355 let capacity = 1 << log_n.saturating_sub(P::LOG_WIDTH);
356 let result = scaled_eq_ind_partial_eval_into::<P, _>(
357 &point,
358 scale,
359 Vec::with_capacity(capacity),
360 );
361
362 assert_eq!(result.log_len(), log_n, "wrong length at log_n={log_n}");
363 assert_eq!(
364 result,
365 scaled_eq_ind_partial_eval::<P>(&point, scale),
366 "mismatch at log_n={log_n}"
367 );
368 }
369 }
370
371 #[test]
372 fn appending_onto_a_one_value_store_builds_from_scratch() {
373 let mut rng = StdRng::seed_from_u64(6);
374
375 let point = random_scalars::<F>(&mut rng, 5);
378 let seed = FieldBuffer::<P, _>::scalar_with_capacity(F::ONE, point.len());
379
380 assert_eq!(tensor_prod_eq_ind::<P>(seed, &point), eq_ind_partial_eval::<P>(&point));
381 }
382
383 #[test]
384 fn appending_in_batches_matches_one_full_expansion() {
385 let mut rng = StdRng::seed_from_u64(7);
386
387 let batches = 4;
391 let max_n_vars = batches * (batches + 1) / 2;
392 let mut coords = Vec::with_capacity(max_n_vars);
393 let mut eq_expansion = FieldBuffer::<P, _>::scalar_with_capacity(F::ONE, max_n_vars);
394
395 for batch_len in 1..=batches {
396 let extra = random_scalars(&mut rng, batch_len);
397
398 eq_expansion = tensor_prod_eq_ind::<P>(eq_expansion, &extra);
399 coords.extend(&extra);
400
401 assert_eq!(eq_expansion.log_len(), coords.len());
403 for i in 0..eq_expansion.len() {
404 let vertex = index_to_hypercube_point(coords.len(), i);
405 assert_eq!(eq_expansion.get(i), eq_ind(&vertex, &coords));
406 }
407 }
408 }
409
410 #[test]
411 fn prepending_via_bit_reverse_matches_one_full_expansion() {
412 let mut rng = StdRng::seed_from_u64(8);
413
414 let n_vars = 10;
421 let point = random_scalars::<F>(&mut rng, n_vars);
422
423 let mut tensor = FieldBuffer::<P>::from_values(&[F::ONE]);
424 for &r in point.iter().rev() {
425 bit_reverse_packed(tensor.as_mut_view());
426 tensor = tensor_prod_eq_ind::<P>(tensor, &[r]);
427 bit_reverse_packed(tensor.as_mut_view());
428 }
429
430 assert_eq!(tensor, eq_ind_partial_eval::<P>(&point));
431 }
432
433 #[test]
434 fn repeated_truncation_matches_expansion_of_the_prefix() {
435 let mut rng = StdRng::seed_from_u64(0);
436
437 let reductions = 4;
441 let n_vars = reductions * (reductions + 1) / 2;
442 let point = random_scalars(&mut rng, n_vars);
443
444 let mut eq_ind = eq_ind_partial_eval::<P>(&point);
445 let mut log_n_values = n_vars;
446
447 for reduction in (0..=reductions).rev() {
448 let truncated_log_n_values = log_n_values - reduction;
449 eq_ind_truncate_low_inplace(&mut eq_ind, truncated_log_n_values);
450
451 let eq_ind_ref = eq_ind_partial_eval::<P>(&point[..truncated_log_n_values]);
453 assert_eq!(eq_ind_ref.len(), eq_ind.len());
454 for i in 0..eq_ind.len() {
455 assert_eq!(eq_ind.get(i), eq_ind_ref.get(i));
456 }
457
458 log_n_values = truncated_log_n_values;
459 }
460
461 assert_eq!(log_n_values, 0);
463 }
464}