1use std::ops::{Deref, DerefMut};
16
17use binius_compute::{Allocator, BufferData, CollectIntoAllocVec};
18use binius_field::{Field, PackedField};
19use binius_utils::{
20 random_access_sequence::RandomAccessSequence,
21 rayon::{
22 prelude::*,
23 task_size::{IndexedParallelIteratorExt, WorkPerItem},
24 },
25};
26
27use crate::{FieldBuffer, FieldVec, line::extrapolate_line};
28
29pub fn fold_highest_var_inplace<P: PackedField, Data: BufferData<P>>(
42 values: &mut FieldBuffer<P, Data>,
43 scalar: P::Scalar,
44) {
45 let broadcast_scalar = P::broadcast(scalar);
48 {
49 let mut split = values.split_half_mut();
51 let (mut lo, mut hi) = split.halves();
52 (lo.as_mut(), hi.as_mut())
54 .into_par_iter()
55 .with_min_task(WorkPerItem::FieldMuls)
56 .for_each(|(lo_i, hi_i)| {
57 *lo_i = extrapolate_line(*lo_i, *hi_i, broadcast_scalar);
58 });
59 }
60
61 values.truncate(values.log_len() - 1);
63}
64
65pub fn fold_highest_var<A: Allocator, P: PackedField, Data: Deref<Target = [P]>>(
81 alloc: &A,
82 values: &FieldBuffer<P, Data>,
83 scalar: P::Scalar,
84) -> FieldVec<P, A> {
85 assert!(values.log_len() > 0, "precondition: buffer must have at least one variable");
86
87 let broadcast_scalar = P::broadcast(scalar);
89 let (lo, hi) = values.split_half();
90
91 let data = (lo.as_ref(), hi.as_ref())
94 .into_par_iter()
95 .with_min_task(WorkPerItem::FieldMuls)
96 .map(|(&lo_i, &hi_i)| extrapolate_line(lo_i, hi_i, broadcast_scalar))
97 .collect_into_alloc_vec(alloc);
98 FieldBuffer::new(values.log_len() - 1, data)
99}
100
101pub fn binary_fold_high<P, DataOut, DataIn>(
114 values: &mut FieldBuffer<P, DataOut>,
115 tensor: &FieldBuffer<P, DataIn>,
116 bits: &(impl RandomAccessSequence<bool> + Sync),
117) where
118 P: PackedField,
119 DataOut: DerefMut<Target = [P]>,
120 DataIn: Deref<Target = [P]>,
121{
122 assert!(bits.len().is_power_of_two(), "precondition: bits length must be a power of two");
123
124 let values_log_len = values.log_len();
125 let width = P::WIDTH.min(values.len());
127
128 assert_eq!(
129 1 << (values_log_len + tensor.log_len()),
130 bits.len(),
131 "precondition: bits length must equal values length times tensor length"
132 );
133
134 values
135 .iter_packed_mut()
136 .enumerate()
137 .for_each(|(i, packed)| {
138 *packed = P::from_scalars((0..width).map(|j| {
139 let scalar_index = i << P::LOG_WIDTH | j;
141 let mut acc = P::Scalar::ZERO;
142
143 for (k, tensor_packed) in tensor.iter_packed().enumerate() {
146 for (l, tensor_scalar) in tensor_packed.iter().take(tensor.len()).enumerate() {
147 let tensor_scalar_index = k << P::LOG_WIDTH | l;
148 if bits.get(tensor_scalar_index << values_log_len | scalar_index) {
149 acc += tensor_scalar;
150 }
151 }
152 }
153
154 acc
155 }));
156 });
157}
158
159#[cfg(test)]
160mod tests {
161 use std::iter::repeat_with;
162
163 use binius_compute::GlobalAllocator;
164 use binius_utils::rayon::task_size::min_len_for_work;
165 use proptest::prelude::*;
166 use rand::prelude::*;
167
168 use super::*;
169 use crate::{
170 multilinear::eq::eq_ind_partial_eval,
171 test_utils::{B128, Packed128b, random_field_buffer, random_scalars},
172 };
173
174 type P = Packed128b;
175 type F = B128;
176
177 const MAX_VARS: usize = 8;
179
180 #[test]
181 fn fold_splits_above_the_task_threshold() {
182 let mut rng = StdRng::seed_from_u64(0);
183
184 let min_len = min_len_for_work(WorkPerItem::FieldMuls);
192 let n_vars = (2 * min_len * P::WIDTH).next_power_of_two().ilog2() as usize + 1;
193 let half = 1 << (n_vars - 1);
194 let original = random_field_buffer::<P>(&mut rng, n_vars);
195 let challenge = random_scalars::<F>(&mut rng, 1)[0];
196
197 let mut folded = original.clone();
198 fold_highest_var_inplace(&mut folded, challenge);
199 assert_eq!(folded.log_len(), n_vars - 1);
200
201 for i in 0..half {
203 let expected = extrapolate_line(original.get(i), original.get(i | half), challenge);
204 assert_eq!(folded.get(i), expected, "mismatch at index {i}");
205 }
206 }
207
208 #[test]
209 fn folding_a_slice_backed_buffer_matches_folding_an_owned_one() {
210 let mut rng = StdRng::seed_from_u64(0);
211
212 let original = random_field_buffer::<P>(&mut rng, 5);
221 let scalar = random_scalars::<F>(&mut rng, 1)[0];
222
223 let mut expected = original.clone();
224 fold_highest_var_inplace(&mut expected, scalar);
225
226 let mut owned = original;
227 let mut slice = owned.as_mut_view();
228 fold_highest_var_inplace(&mut slice, scalar);
229
230 assert_eq!(slice.log_len(), 4);
231 assert_eq!(slice, expected.as_mut_view());
232 }
233
234 proptest! {
235 #[test]
236 fn the_two_folds_agree(
237 n_vars in 1..=MAX_VARS,
238 seed: u64,
239 ) {
240 let mut rng = StdRng::seed_from_u64(seed);
241 let original = random_field_buffer::<P>(&mut rng, n_vars);
242 let scalar = random_scalars::<F>(&mut rng, 1)[0];
243
244 let out_of_place = fold_highest_var(&GlobalAllocator, &original, scalar);
246
247 let mut in_place = original;
249 fold_highest_var_inplace(&mut in_place, scalar);
250
251 prop_assert_eq!(out_of_place.log_len(), n_vars - 1);
252 prop_assert_eq!(out_of_place, in_place);
253 }
254
255 #[test]
256 fn the_binary_fold_matches_folding_the_widened_bits(
257 dest_vars in 0..=6usize,
258 tensor_vars in 0..=4usize,
259 seed: u64,
260 ) {
261 let mut rng = StdRng::seed_from_u64(seed);
262 let point = random_scalars::<F>(&mut rng, tensor_vars);
263 let tensor = eq_ind_partial_eval::<P>(&point);
264
265 let bits = repeat_with(|| rng.random())
267 .take(1 << (dest_vars + tensor_vars))
268 .collect::<Vec<bool>>();
269
270 let mut folded = FieldBuffer::<P>::zeros(dest_vars);
271 binary_fold_high(&mut folded, &tensor, &bits.as_slice());
272
273 let scalars = bits
275 .iter()
276 .map(|&bit| if bit { F::ONE } else { F::ZERO })
277 .collect::<Vec<F>>();
278 let mut reference = FieldBuffer::<P>::from_values(&scalars);
279 for &coord in point.iter().rev() {
280 fold_highest_var_inplace(&mut reference, coord);
281 }
282
283 prop_assert_eq!(folded, reference);
284 }
285 }
286}