1use std::{iter, slice};
24
25use binius_compute::{Allocator, BufferData, VecLike};
26use binius_field::{Field, PackedField, field::FieldOps};
27use binius_utils::rayon::{
28 prelude::*,
29 task_size::{IndexedParallelIteratorExt, WorkPerItem, min_len_for_bytes},
30};
31
32use crate::{FieldBuffer, FieldVec};
33
34pub trait Hypercube {
40 fn basis<F: FieldOps>(coord: &F) -> [F; 2];
44
45 fn expand_var<F: FieldOps>(value: &F, coord: &F) -> [F; 2];
51
52 fn contract_var<F: FieldOps>(lo: &mut F, hi: &F);
65
66 fn eq_one_var<F: FieldOps>(x: F, y: F) -> F {
74 let [x_0, x_1] = Self::basis(&x);
76 let [y_0, y_1] = Self::basis(&y);
77 x_0 * y_0 + x_1 * y_1
78 }
79}
80
81#[derive(Debug)]
86pub struct OneCube;
87
88impl Hypercube for OneCube {
89 #[inline(always)]
90 fn basis<F: FieldOps>(coord: &F) -> [F; 2] {
91 [F::one() - coord, coord.clone()]
92 }
93
94 #[inline(always)]
95 fn expand_var<F: FieldOps>(value: &F, coord: &F) -> [F; 2] {
96 let prod = value.clone() * coord;
98 [value.clone() - &prod, prod]
99 }
100
101 #[inline(always)]
102 fn contract_var<F: FieldOps>(lo: &mut F, hi: &F) {
103 *lo += hi;
105 }
106
107 #[inline(always)]
108 fn eq_one_var<F: FieldOps>(x: F, y: F) -> F {
109 if F::Scalar::CHARACTERISTIC == 2 {
115 x + y + F::one()
116 } else {
117 let one = F::one();
118 x.clone() * y.clone() + (one.clone() - x) * (one - y)
119 }
120 }
121}
122
123#[derive(Debug)]
132pub struct InfCube;
133
134impl Hypercube for InfCube {
135 #[inline(always)]
136 fn basis<F: FieldOps>(coord: &F) -> [F; 2] {
137 [F::one(), coord.clone()]
138 }
139
140 #[inline(always)]
141 fn expand_var<F: FieldOps>(value: &F, coord: &F) -> [F; 2] {
142 [value.clone(), value.clone() * coord]
144 }
145
146 #[inline(always)]
147 fn contract_var<F: FieldOps>(_lo: &mut F, _hi: &F) {
148 }
151
152 #[inline(always)]
153 fn eq_one_var<F: FieldOps>(x: F, y: F) -> F {
154 F::one() + x * y
156 }
157}
158
159pub fn tensor_prod_eq_ind<Cube: Hypercube, P: PackedField>(
186 values: FieldBuffer<P, Vec<P>>,
187 extra_query_coordinates: &[P::Scalar],
188) -> FieldBuffer<P, Vec<P>> {
189 let start_log_len = values.log_len();
190 let final_log_len = start_log_len + extra_query_coordinates.len();
191 let mut data = values.into_inner();
192
193 let final_packed_len = packed_words::<P>(final_log_len);
197 data.reserve_exact(final_packed_len.saturating_sub(data.len()));
198
199 tensor_prod_eq_ind_reserved::<Cube, P, _>(
200 FieldBuffer::new(start_log_len, data),
201 extra_query_coordinates,
202 )
203}
204
205const fn packed_words<P: PackedField>(log_len: usize) -> usize {
209 1usize << log_len.saturating_sub(P::LOG_WIDTH)
210}
211
212fn tensor_prod_eq_ind_reserved<Cube: Hypercube, P: PackedField, Data: VecLike<P>>(
222 values: FieldBuffer<P, Data>,
223 extra_query_coordinates: &[P::Scalar],
224) -> FieldBuffer<P, Data> {
225 let start_log_len = values.log_len();
226 let final_log_len = start_log_len + extra_query_coordinates.len();
227 let mut data = values.into_inner();
228
229 debug_assert!(data.capacity() >= packed_words::<P>(final_log_len));
231
232 let sub_width_count = extra_query_coordinates
237 .len()
238 .min(P::LOG_WIDTH.saturating_sub(start_log_len));
239 let (sub_width_coords, packed_coords) = extra_query_coordinates.split_at(sub_width_count);
240
241 for (i, &r_i) in sub_width_coords.iter().enumerate() {
245 let log_len = start_log_len + i;
246 let packed_r_i = P::broadcast(r_i);
247 let (lo, _) = data[0].interleave(P::zero(), log_len);
248 let [lo, hi] = Cube::expand_var(&lo, &packed_r_i);
249 data[0] = lo.interleave(hi, log_len).0;
250 }
251
252 for &r_i in packed_coords {
257 let packed_r_i = P::broadcast(r_i);
258 let old_packed = data.len();
259
260 let low_ptr = data.as_mut_ptr();
264 let high = &mut data.spare_capacity_mut()[..old_packed];
265 let low = unsafe { slice::from_raw_parts_mut(low_ptr, old_packed) };
268 (low, high)
271 .into_par_iter()
272 .with_min_task(WorkPerItem::FieldMuls)
273 .for_each(|(low_i, high_i)| {
274 let [new_low, new_high] = Cube::expand_var(low_i, &packed_r_i);
275 *low_i = new_low;
276 high_i.write(new_high);
277 });
278 unsafe { data.set_len(2 * old_packed) };
280 }
281
282 FieldBuffer::new(final_log_len, data)
283}
284
285pub fn eq_ind_partial_eval<Cube: Hypercube, P: PackedField>(point: &[P::Scalar]) -> FieldBuffer<P> {
295 scaled_eq_ind_partial_eval::<Cube, P>(point, P::Scalar::ONE)
297}
298
299pub fn eq_ind_partial_eval_in<Cube: Hypercube, A: Allocator, P: PackedField>(
303 alloc: &A,
304 point: &[P::Scalar],
305) -> FieldVec<P, A> {
306 let packed_len = packed_words::<P>(point.len());
308 scaled_eq_ind_partial_eval_into::<Cube, P, _>(
309 point,
310 P::Scalar::ONE,
311 alloc.alloc::<P>(packed_len),
312 )
313}
314
315pub fn scaled_eq_ind_partial_eval<Cube: Hypercube, P: PackedField>(
325 point: &[P::Scalar],
326 scale: P::Scalar,
327) -> FieldBuffer<P> {
328 let packed_len = packed_words::<P>(point.len());
330 scaled_eq_ind_partial_eval_into::<Cube, P, _>(point, scale, Vec::with_capacity(packed_len))
331}
332
333pub fn scaled_eq_ind_partial_eval_into<Cube: Hypercube, P: PackedField, Data: VecLike<P>>(
346 point: &[P::Scalar],
347 scale: P::Scalar,
348 mut buffer: Data,
349) -> FieldBuffer<P, Data> {
350 assert!(
351 buffer.capacity() >= packed_words::<P>(point.len()),
352 "precondition: buffer capacity must cover the packed expansion length"
353 );
354
355 buffer.clear();
358 buffer.push(P::from_scalars(iter::once(scale)));
359 let seed = FieldBuffer::new(0, buffer);
360
361 let low_len = (point.len() / 2).max(P::LOG_WIDTH).min(point.len());
371 let (low_coords, high_coords) = point.split_at(low_len);
372
373 let low = tensor_prod_eq_ind_reserved::<Cube, P, Data>(seed, low_coords);
376
377 if high_coords.is_empty() {
379 return low;
380 }
381
382 let block = low.as_ref().len();
386 let mut data = low.into_inner();
387
388 let high = eq_ind_partial_eval_scalars::<Cube, P::Scalar>(high_coords);
390 let total = block * high.len();
391 debug_assert_eq!(total, packed_words::<P>(point.len()));
392
393 let first_ptr = data.as_mut_ptr();
397 let spare = &mut data.spare_capacity_mut()[..total - block];
398 let first = unsafe { slice::from_raw_parts(first_ptr, block) };
401
402 let min_len = (min_len_for_bytes::<P>() / block).max(1);
405 spare
406 .par_chunks_mut(block)
407 .zip(high[1..].par_iter())
408 .with_min_len(min_len)
409 .for_each(|(dst, &coeff)| {
410 let coeff = P::broadcast(coeff);
411 for (dst_i, &src_i) in iter::zip(dst, first) {
412 dst_i.write(src_i * coeff);
413 }
414 });
415 unsafe { data.set_len(total) };
417
418 let coeff = P::broadcast(high[0]);
421 for word in &mut data[..block] {
422 *word *= coeff;
423 }
424
425 FieldBuffer::new(point.len(), data)
426}
427
428pub fn eq_ind_truncate_low_inplace<Cube: Hypercube, P: PackedField, Data: BufferData<P>>(
441 values: &mut FieldBuffer<P, Data>,
442 truncated_log_len: usize,
443) {
444 assert!(
445 truncated_log_len <= values.log_len(),
446 "precondition: truncated_log_len must be at most values.log_len()"
447 );
448
449 for log_len in (truncated_log_len..values.log_len()).rev() {
451 {
452 let mut split = values.split_half_mut();
453 let (mut lo, hi) = split.halves();
454 (lo.as_mut(), hi.as_ref())
457 .into_par_iter()
458 .with_min_task_bytes::<[P; 2]>()
459 .for_each(|(zero, one)| {
460 Cube::contract_var(zero, one);
461 });
462 }
463
464 values.truncate(log_len);
465 }
466}
467
468pub fn eq_ind<Cube: Hypercube, F: FieldOps>(x: &[F], y: &[F]) -> F {
476 assert_eq!(x.len(), y.len(), "pre-condition: x and y must be the same length");
477 iter::zip(x, y)
479 .map(|(x, y)| Cube::eq_one_var(x.clone(), y.clone()))
480 .product()
481}
482
483pub fn eq_ind_zero<Cube: Hypercube, F: FieldOps>(point: &[F]) -> F {
491 point
493 .iter()
494 .map(|y| {
495 let [y_0, _] = Cube::basis(y);
496 y_0
497 })
498 .product()
499}
500
501pub fn eq_ind_partial_eval_scalars<Cube: Hypercube, F: FieldOps>(point: &[F]) -> Vec<F> {
505 scaled_eq_ind_partial_eval_scalars::<Cube, F>(point, F::one())
507}
508
509pub fn scaled_eq_ind_partial_eval_scalars<Cube: Hypercube, F: FieldOps>(
514 point: &[F],
515 scale: F,
516) -> Vec<F> {
517 let mut result = Vec::with_capacity(1 << point.len());
519 result.push(scale);
521
522 for r_i in point {
523 let len = result.len();
530 for j in 0..len {
531 let [lo, hi] = Cube::expand_var(&result[j], r_i);
532 result[j] = lo;
533 result.push(hi);
534 }
535 }
536 result
537}
538
539#[cfg(test)]
540mod tests {
541 use binius_utils::rayon::task_size::{min_len_for_bytes, min_len_for_work};
542 use proptest::prelude::*;
543 use rand::prelude::*;
544
545 use super::*;
546 use crate::test_utils::{B128, Packed128b, random_scalars};
547
548 type P = Packed128b;
549 type F = B128;
550
551 #[test]
552 fn expand_var_matches_scaled_basis() {
553 let mut rng = StdRng::seed_from_u64(0);
554
555 let [value, coord] = [(); 2].map(|_| random_scalars::<F>(&mut rng, 1)[0]);
558 assert_eq!(
559 OneCube::expand_var(&value, &coord),
560 OneCube::basis(&coord).map(|b_i| b_i * value)
561 );
562 assert_eq!(
563 InfCube::expand_var(&value, &coord),
564 InfCube::basis(&coord).map(|b_i| b_i * value)
565 );
566 }
567
568 #[test]
569 fn contract_var_inverts_expand_var() {
570 let mut rng = StdRng::seed_from_u64(0);
571
572 let [value, coord] = [(); 2].map(|_| random_scalars::<F>(&mut rng, 1)[0]);
574
575 let [mut lo, hi] = OneCube::expand_var(&value, &coord);
576 OneCube::contract_var(&mut lo, &hi);
577 assert_eq!(lo, value);
578
579 let [mut lo, hi] = InfCube::expand_var(&value, &coord);
580 InfCube::contract_var(&mut lo, &hi);
581 assert_eq!(lo, value);
582 }
583
584 #[test]
585 fn eq_one_var_matches_basis_definition() {
586 let mut rng = StdRng::seed_from_u64(0);
587
588 let [x, y] = [(); 2].map(|_| random_scalars::<F>(&mut rng, 1)[0]);
591 let eq_from_basis = |[x_0, x_1]: [F; 2], [y_0, y_1]: [F; 2]| x_0 * y_0 + x_1 * y_1;
592 assert_eq!(
593 OneCube::eq_one_var(x, y),
594 eq_from_basis(OneCube::basis(&x), OneCube::basis(&y))
595 );
596 assert_eq!(
597 InfCube::eq_one_var(x, y),
598 eq_from_basis(InfCube::basis(&x), InfCube::basis(&y))
599 );
600 }
601
602 #[test]
603 fn inf_cube_eq_ind_zero_is_one() {
604 let mut rng = StdRng::seed_from_u64(0);
605
606 for n_vars in [0, 1, 5] {
608 let point = random_scalars::<F>(&mut rng, n_vars);
609 assert_eq!(eq_ind_zero::<InfCube, F>(&point), F::ONE);
610
611 assert_eq!(
613 eq_ind_zero::<InfCube, F>(&point),
614 eq_ind::<InfCube, F>(&vec![F::ZERO; n_vars], &point)
615 );
616 }
617 }
618
619 fn inf_cube_reference(point: &[F]) -> Vec<F> {
621 (0..1 << point.len())
623 .map(|index| {
624 point
625 .iter()
626 .enumerate()
627 .filter(|(i, _)| index >> i & 1 == 1)
628 .map(|(_, r_i)| *r_i)
629 .product()
630 })
631 .collect()
632 }
633
634 fn eval_monomial_basis(coeffs: &[F], point: &[F]) -> F {
636 coeffs
638 .iter()
639 .enumerate()
640 .map(|(index, coeff)| {
641 *coeff
642 * point
643 .iter()
644 .enumerate()
645 .filter(|(i, _)| index >> i & 1 == 1)
646 .map(|(_, x_i)| *x_i)
647 .product::<F>()
648 })
649 .sum()
650 }
651
652 #[test]
653 fn inf_cube_expansion_matches_the_tensor_of_bases() {
654 let mut rng = StdRng::seed_from_u64(0);
655
656 for n_vars in [0, 1, 2, 5, 8] {
658 let point = random_scalars::<F>(&mut rng, n_vars);
659 let expansion = eq_ind_partial_eval::<InfCube, P>(&point);
660 let expansion_scalars = expansion.iter_scalars().collect::<Vec<_>>();
661 assert_eq!(expansion_scalars, inf_cube_reference(&point), "mismatch at {n_vars} vars");
662 }
663 }
664
665 #[test]
666 fn inf_cube_expansion_holds_the_monomial_coefficients_of_the_indicator() {
667 let mut rng = StdRng::seed_from_u64(0);
668
669 for n_vars in [0, 1, 2, 5] {
675 let point = random_scalars::<F>(&mut rng, n_vars);
676 let coeffs = eq_ind_partial_eval_scalars::<InfCube, F>(&point);
677
678 let x = random_scalars::<F>(&mut rng, n_vars);
679 assert_eq!(eval_monomial_basis(&coeffs, &x), eq_ind::<InfCube, F>(&x, &point));
680 }
681 }
682
683 #[test]
684 fn inf_cube_expansion_is_the_evaluation_functional() {
685 let mut rng = StdRng::seed_from_u64(0);
686
687 for n_vars in [0, 1, 2, 5] {
690 let point = random_scalars::<F>(&mut rng, n_vars);
691 let coeffs = random_scalars::<F>(&mut rng, 1 << n_vars);
692
693 let expansion = eq_ind_partial_eval_scalars::<InfCube, F>(&point);
694 let inner_product = iter::zip(&coeffs, &expansion)
695 .map(|(c, e)| *c * e)
696 .sum::<F>();
697 assert_eq!(inner_product, eval_monomial_basis(&coeffs, &point));
698 }
699 }
700
701 #[test]
702 fn growth_above_the_split_threshold_matches_the_inline_path() {
703 let mut rng = StdRng::seed_from_u64(9);
704
705 let min_len = min_len_for_work(WorkPerItem::FieldMuls);
713 let n_vars = (2 * min_len * P::WIDTH).next_power_of_two().ilog2() as usize;
714 let point = random_scalars::<F>(&mut rng, n_vars);
715
716 let packed = eq_ind_partial_eval::<OneCube, P>(&point);
717 let reference = eq_ind_partial_eval_scalars::<OneCube, F>(&point);
718 assert!(packed.iter_scalars().eq(reference.iter().copied()));
719 }
720
721 #[test]
722 fn truncation_above_the_split_threshold_matches_the_inline_path() {
723 let mut rng = StdRng::seed_from_u64(0);
724
725 let min_len = min_len_for_bytes::<[P; 2]>();
733 let n_vars = (2 * min_len * P::WIDTH).next_power_of_two().ilog2() as usize;
734 let point = random_scalars::<F>(&mut rng, n_vars);
735
736 let mut truncated = eq_ind_partial_eval::<OneCube, P>(&point);
738 eq_ind_truncate_low_inplace::<OneCube, _, _>(&mut truncated, n_vars - 1);
739 assert_eq!(truncated, eq_ind_partial_eval::<OneCube, P>(&point[..n_vars - 1]));
740 }
741
742 proptest! {
743 #![proptest_config(ProptestConfig::with_cases(16))]
744
745 #[test]
746 fn the_two_engines_agree(
747 seed in any::<u64>(),
748 log_n in 0usize..=8,
749 ) {
750 let mut rng = StdRng::seed_from_u64(seed);
751 let point = random_scalars::<F>(&mut rng, log_n);
752 let scale = random_scalars::<F>(&mut rng, 1)[0];
753
754 prop_assert_eq!(
757 eq_ind_partial_eval::<OneCube, P>(&point).iter_scalars().collect::<Vec<_>>(),
758 eq_ind_partial_eval_scalars::<OneCube, F>(&point)
759 );
760 prop_assert_eq!(
761 eq_ind_partial_eval::<InfCube, P>(&point).iter_scalars().collect::<Vec<_>>(),
762 eq_ind_partial_eval_scalars::<InfCube, F>(&point)
763 );
764 prop_assert_eq!(
765 scaled_eq_ind_partial_eval::<OneCube, P>(&point, scale)
766 .iter_scalars()
767 .collect::<Vec<_>>(),
768 scaled_eq_ind_partial_eval_scalars::<OneCube, F>(&point, scale)
769 );
770 prop_assert_eq!(
771 scaled_eq_ind_partial_eval::<InfCube, P>(&point, scale)
772 .iter_scalars()
773 .collect::<Vec<_>>(),
774 scaled_eq_ind_partial_eval_scalars::<InfCube, F>(&point, scale)
775 );
776
777 prop_assert_eq!(
785 eq_ind_partial_eval::<OneCube, F>(&point).into_inner(),
786 eq_ind_partial_eval_scalars::<OneCube, F>(&point)
787 );
788 prop_assert_eq!(
789 eq_ind_partial_eval::<InfCube, F>(&point).into_inner(),
790 eq_ind_partial_eval_scalars::<InfCube, F>(&point)
791 );
792 }
793
794 #[test]
795 fn scaling_commutes_with_the_expansion(
796 seed in any::<u64>(),
797 log_n in 0usize..=8,
798 ) {
799 let mut rng = StdRng::seed_from_u64(seed);
800 let point = random_scalars::<F>(&mut rng, log_n);
801 let scale = random_scalars::<F>(&mut rng, 1)[0];
802
803 let one_scaled = scaled_eq_ind_partial_eval::<OneCube, P>(&point, scale);
805 let one_plain = eq_ind_partial_eval::<OneCube, P>(&point);
806 for (got, base) in one_scaled.iter_scalars().zip(one_plain.iter_scalars()) {
807 prop_assert_eq!(got, scale * base);
808 }
809
810 let inf_scaled = scaled_eq_ind_partial_eval::<InfCube, P>(&point, scale);
811 let inf_plain = eq_ind_partial_eval::<InfCube, P>(&point);
812 for (got, base) in inf_scaled.iter_scalars().zip(inf_plain.iter_scalars()) {
813 prop_assert_eq!(got, scale * base);
814 }
815 }
816
817 #[test]
818 fn truncation_strips_trailing_variables(
819 seed in any::<u64>(),
820 log_n in 0usize..=8,
821 ) {
822 let mut rng = StdRng::seed_from_u64(seed);
823 let point = random_scalars::<F>(&mut rng, log_n);
824
825 for truncated_log_len in 0..=log_n {
827 let mut one_cube = eq_ind_partial_eval::<OneCube, P>(&point);
828 eq_ind_truncate_low_inplace::<OneCube, _, _>(&mut one_cube, truncated_log_len);
829 prop_assert_eq!(
830 one_cube,
831 eq_ind_partial_eval::<OneCube, P>(&point[..truncated_log_len])
832 );
833
834 let mut inf_cube = eq_ind_partial_eval::<InfCube, P>(&point);
835 eq_ind_truncate_low_inplace::<InfCube, _, _>(&mut inf_cube, truncated_log_len);
836 prop_assert_eq!(
837 inf_cube,
838 eq_ind_partial_eval::<InfCube, P>(&point[..truncated_log_len])
839 );
840 }
841 }
842
843 #[test]
844 fn the_split_agrees_with_the_doubling_rounds(
845 seed in any::<u64>(),
846 n_vars in 0usize..=10,
847 ) {
848 let mut rng = StdRng::seed_from_u64(seed);
854 let point = random_scalars::<F>(&mut rng, n_vars);
855 let scale = random_scalars::<F>(&mut rng, 1)[0];
856
857 let seeded = || {
859 let mut buffer = Vec::with_capacity(1 << n_vars.saturating_sub(P::LOG_WIDTH));
860 buffer.push(P::from_scalars(iter::once(scale)));
861 FieldBuffer::new(0, buffer)
862 };
863
864 prop_assert_eq!(
865 scaled_eq_ind_partial_eval::<OneCube, P>(&point, scale),
866 tensor_prod_eq_ind_reserved::<OneCube, P, _>(seeded(), &point)
867 );
868 prop_assert_eq!(
869 scaled_eq_ind_partial_eval::<InfCube, P>(&point, scale),
870 tensor_prod_eq_ind_reserved::<InfCube, P, _>(seeded(), &point)
871 );
872 }
873 }
874
875 #[test]
876 fn the_split_path_agrees_with_the_scalar_engine() {
877 let mut rng = StdRng::seed_from_u64(0);
878
879 for n_vars in (0..=8).chain([13, 14, 19]) {
886 let point = random_scalars::<F>(&mut rng, n_vars);
887 let scale = random_scalars::<F>(&mut rng, 1)[0];
888
889 let packed = scaled_eq_ind_partial_eval::<OneCube, P>(&point, scale);
890 let scalars = scaled_eq_ind_partial_eval_scalars::<OneCube, F>(&point, scale);
891 assert!(packed.iter_scalars().eq(scalars), "one cube at {n_vars} vars");
892
893 let packed = scaled_eq_ind_partial_eval::<InfCube, P>(&point, scale);
894 let scalars = scaled_eq_ind_partial_eval_scalars::<InfCube, F>(&point, scale);
895 assert!(packed.iter_scalars().eq(scalars), "inf cube at {n_vars} vars");
896 }
897 }
898}