1use std::{
4 cmp::{max, min},
5 iter,
6 ops::Range,
7 slice::from_raw_parts_mut,
8};
9
10use binius_field::{BinaryField, PackedField};
11use binius_utils::rayon::{
12 iter::{IndexedParallelIterator, IntoParallelIterator, ParallelIterator},
13 slice::ParallelSliceMut,
14};
15use itertools::izip;
16
17use super::{
18 AdditiveNTT, DomainContext,
19 reference::{NeighborsLastReference, input_check},
20};
21use crate::field_buffer::FieldSliceMut;
22
23const DEFAULT_LOG_BASE_LEN: usize = 10;
27
28fn forward_depth_first<P: PackedField>(
59 domain_context: &impl DomainContext<Field = P::Scalar>,
60 data: &mut [P],
61 log_d: usize,
62 layer: usize,
63 block: usize,
64 mut layer_range: Range<usize>,
65 log_base_len: usize,
66) {
67 debug_assert!(P::LOG_WIDTH < log_d);
69 debug_assert_eq!(data.len(), 1 << (log_d - P::LOG_WIDTH));
70 debug_assert!(layer_range.end <= domain_context.log_domain_size());
71 debug_assert!(layer <= layer_range.start);
72 debug_assert!(log_base_len > P::LOG_WIDTH);
73
74 let log_n = log_d + layer;
75 debug_assert!(layer_range.end <= log_n);
76
77 if layer >= layer_range.end {
78 return;
79 }
80
81 if log_d <= log_base_len {
83 forward_breadth_first(domain_context, data, log_d, layer, block, layer_range);
84 return;
85 }
86
87 let block_size_half = 1 << (log_d - 1 - P::LOG_WIDTH);
88 if layer >= layer_range.start {
89 let (block0, block1) = data.split_at_mut(block_size_half);
91 if block == 0 {
92 for (u, v) in iter::zip(block0, block1) {
95 *v += *u;
96 }
97 } else {
98 let twiddle = domain_context.twiddle(layer, block);
99 let packed_twiddle = P::broadcast(twiddle);
100 for (u, v) in iter::zip(block0, block1) {
101 *u += *v * packed_twiddle;
103 *v += *u;
104 }
105 }
106
107 layer_range.start += 1;
108 }
109
110 forward_depth_first(
112 domain_context,
113 &mut data[..block_size_half],
114 log_d - 1,
115 layer + 1,
116 block << 1,
117 layer_range.clone(),
118 log_base_len,
119 );
120 forward_depth_first(
121 domain_context,
122 &mut data[block_size_half..],
123 log_d - 1,
124 layer + 1,
125 (block << 1) + 1,
126 layer_range,
127 log_base_len,
128 );
129}
130
131fn forward_breadth_first<P: PackedField>(
141 domain_context: &impl DomainContext<Field = P::Scalar>,
142 data: &mut [P],
143 log_d: usize,
144 base_layer: usize,
145 base_block: usize,
146 layer_range: Range<usize>,
147) {
148 debug_assert!(P::LOG_WIDTH < log_d);
150 debug_assert_eq!(data.len(), 1 << (log_d - P::LOG_WIDTH));
151 debug_assert!(layer_range.end <= domain_context.log_domain_size());
152 debug_assert!(base_layer <= layer_range.start);
153
154 let log_n = log_d + base_layer;
155 debug_assert!(layer_range.end <= log_n);
156
157 let packed_cutoff = (log_n - P::LOG_WIDTH).clamp(layer_range.start, layer_range.end);
158
159 for layer in layer_range.start..packed_cutoff {
162 let log_block_size = log_n - P::LOG_WIDTH - layer;
164 let log_half_block_size = log_block_size - 1;
165
166 let log_blocks = layer - base_layer;
168 let mut layer_twiddles = domain_context
169 .iter_twiddles(layer, 0)
170 .skip(base_block << log_blocks)
171 .take(1 << log_blocks);
172 let mut blocks = data.chunks_exact_mut(1 << log_block_size);
173
174 if base_block == 0 {
178 layer_twiddles.next();
179 if let Some(block) = blocks.next() {
180 let (block0, block1) = block.split_at_mut(1 << log_half_block_size);
181 for (u, v) in iter::zip(block0, block1) {
182 *v += *u;
183 }
184 }
185 }
186
187 for (block, twiddle) in iter::zip(blocks, layer_twiddles) {
188 let packed_twiddle = P::broadcast(twiddle);
189 let (block0, block1) = block.split_at_mut(1 << log_half_block_size);
190 for (u, v) in iter::zip(block0, block1) {
191 *u += *v * packed_twiddle;
193 *v += *u;
194 }
195 }
196 }
197
198 for layer in packed_cutoff..layer_range.end {
202 let log_block_size = log_n - layer;
204 let log_half_block_size = log_block_size - 1;
205 let log_blocks_per_packed = P::LOG_WIDTH - log_block_size;
206 let log_half_blocks_per_packed = log_blocks_per_packed + 1;
207
208 let mut packed_twiddle_offset = P::zero();
210 for block in 0..1 << log_blocks_per_packed {
211 let twiddle0 = domain_context.twiddle(layer, block);
212 let twiddle1 = domain_context.twiddle(layer, (1 << log_blocks_per_packed) | block);
213
214 let block_start = block << log_block_size;
215 for j in 0..1 << log_half_block_size {
216 packed_twiddle_offset.set(block_start | j, twiddle0);
217 packed_twiddle_offset.set(block_start | j | (1 << log_half_block_size), twiddle1);
218 }
219 }
220
221 let log_packed_pairs = log_d - P::LOG_WIDTH - 1;
224 let layer_twiddles = domain_context
225 .iter_twiddles(layer, log_half_blocks_per_packed)
226 .skip(base_block << log_packed_pairs)
227 .take(1 << log_packed_pairs);
228
229 let (data_pairs, rest) = data.as_chunks_mut::<2>();
230 debug_assert!(
231 rest.is_empty(),
232 "data_packed length is a power of two; \
233 data_packed length is greater than 1 (checked at beginning of method)"
234 );
235 debug_assert_eq!(data_pairs.len(), 1 << log_packed_pairs);
236
237 for ([packed0, packed1], first_twiddle) in iter::zip(data_pairs, layer_twiddles) {
238 let packed_twiddle = P::broadcast(first_twiddle) + packed_twiddle_offset;
239
240 let (mut u, mut v) = (*packed0).interleave(*packed1, log_half_block_size);
241 u += v * packed_twiddle;
242 v += u;
243 (*packed0, *packed1) = u.interleave(v, log_half_block_size);
244 }
245 }
246}
247
248fn forward_shared_layer<P: PackedField>(
260 domain_context: &(impl DomainContext<Field = P::Scalar> + Sync),
261 data: &mut [P],
262 log_d: usize,
263 layer: usize,
264 log_num_shares: usize,
265) {
266 debug_assert_eq!(data.len() * (1 << P::LOG_WIDTH), 1 << log_d);
268 debug_assert!(1 << (log_num_shares + 1) <= data.len());
269 debug_assert!(layer < domain_context.log_domain_size());
270
271 let log_num_chunks = log_num_shares + 1;
272 let log_d_chunk = log_d - log_num_chunks;
273 let data_ptr = data.as_mut_ptr();
274 let shift = log_num_shares - layer;
275 let tasks: Vec<_> = (0..1 << log_num_shares)
276 .map(|k| {
277 let (chunk0, chunk1) = with_middle_bit(k, shift);
278 let block = chunk0 >> (log_num_chunks - layer);
279 assert!(P::LOG_WIDTH <= log_d_chunk);
280 let log_chunk_len = log_d_chunk - P::LOG_WIDTH;
281 let chunk0 = unsafe {
282 from_raw_parts_mut(data_ptr.add(chunk0 << log_chunk_len), 1 << log_chunk_len)
283 };
284 let chunk1 = unsafe {
285 from_raw_parts_mut(data_ptr.add(chunk1 << log_chunk_len), 1 << log_chunk_len)
286 };
287 let twiddle = (block != 0).then(|| P::broadcast(domain_context.twiddle(layer, block)));
290 (chunk0, chunk1, twiddle)
291 })
292 .collect();
293
294 tasks
295 .into_par_iter()
296 .for_each(|(chunk0, chunk1, twiddle)| match twiddle {
297 Some(twiddle) => {
298 for (u, v) in iter::zip(chunk0, chunk1) {
299 butterfly(u, v, twiddle);
300 }
301 }
302 None => {
303 for (u, v) in iter::zip(chunk0, chunk1) {
304 *v += *u;
305 }
306 }
307 });
308}
309
310#[inline(always)]
312fn butterfly<P: PackedField>(u: &mut P, v: &mut P, twiddle: P) {
313 *u += *v * twiddle;
314 *v += *u;
315}
316
317fn fused_pair<P: PackedField>(
325 planes: [&mut [P]; 4],
326 twiddle_0: P,
327 twiddle_1_even: P,
328 twiddle_1_odd: P,
329) {
330 let [plane_0, plane_1, plane_2, plane_3] = planes;
331
332 for (x_0, x_1, x_2, x_3) in izip!(plane_0, plane_1, plane_2, plane_3) {
333 butterfly(x_0, x_2, twiddle_0);
334 butterfly(x_1, x_3, twiddle_0);
335 butterfly(x_0, x_1, twiddle_1_even);
336 butterfly(x_2, x_3, twiddle_1_odd);
337 }
338}
339
340fn fused_pair_zero_block<P: PackedField>(planes: [&mut [P]; 4], twiddle_1_odd: P) {
346 let [plane_0, plane_1, plane_2, plane_3] = planes;
347
348 for (x_0, x_1, x_2, x_3) in izip!(plane_0, plane_1, plane_2, plane_3) {
349 *x_2 += *x_0;
350 *x_3 += *x_1;
351 *x_1 += *x_0;
352 butterfly(x_2, x_3, twiddle_1_odd);
353 }
354}
355
356fn forward_shared_layer_pair<P: PackedField>(
387 domain_context: &(impl DomainContext<Field = P::Scalar> + Sync),
388 data: &mut [P],
389 log_d: usize,
390 first_layer: usize,
391 log_num_shares: usize,
392) {
393 debug_assert_eq!(data.len() << P::LOG_WIDTH, 1 << log_d);
395 debug_assert!(first_layer + 2 <= log_num_shares);
396 debug_assert!(first_layer + 2 + P::LOG_WIDTH < log_d);
397 debug_assert!(first_layer + 2 <= domain_context.log_domain_size());
398
399 let log_plane_len = log_d - first_layer - 2 - P::LOG_WIDTH;
400 let plane_len = 1 << log_plane_len;
401
402 let log_run_len = log_plane_len.saturating_sub(log_num_shares - first_layer);
405 let run_len = 1 << log_run_len;
406
407 let tasks = data
409 .chunks_exact_mut(plane_len << 2)
410 .enumerate()
411 .flat_map(|(h, super_block)| {
412 let twiddle = |layer, block| P::broadcast(domain_context.twiddle(layer, block));
413 let twiddles = (
414 twiddle(first_layer, h),
415 twiddle(first_layer + 1, h << 1),
416 twiddle(first_layer + 1, (h << 1) | 1),
417 );
418
419 let (halves_0_1, halves_2_3) = super_block.split_at_mut(plane_len << 1);
420 let (plane_0, plane_1) = halves_0_1.split_at_mut(plane_len);
421 let (plane_2, plane_3) = halves_2_3.split_at_mut(plane_len);
422
423 izip!(
424 plane_0.chunks_exact_mut(run_len),
425 plane_1.chunks_exact_mut(run_len),
426 plane_2.chunks_exact_mut(run_len),
427 plane_3.chunks_exact_mut(run_len),
428 )
429 .map(move |(run_0, run_1, run_2, run_3)| ([run_0, run_1, run_2, run_3], h, twiddles))
430 })
431 .collect::<Vec<_>>();
432
433 tasks
434 .into_par_iter()
435 .for_each(|(planes, h, (twiddle_0, twiddle_1_even, twiddle_1_odd))| {
436 if h == 0 {
439 fused_pair_zero_block(planes, twiddle_1_odd);
440 } else {
441 fused_pair(planes, twiddle_0, twiddle_1_even, twiddle_1_odd);
442 }
443 });
444}
445
446fn forward_shared_layers<P: PackedField>(
457 domain_context: &(impl DomainContext<Field = P::Scalar> + Sync),
458 data: &mut [P],
459 log_d: usize,
460 layers: Range<usize>,
461 log_num_shares: usize,
462) {
463 debug_assert!(layers.end <= log_num_shares);
464
465 let mut layer = layers.start;
466 while layer < layers.end {
467 if layers.end - layer >= 2 && layer + 2 + P::LOG_WIDTH < log_d {
469 forward_shared_layer_pair(domain_context, data, log_d, layer, log_num_shares);
470 layer += 2;
471 } else {
472 forward_shared_layer(domain_context, data, log_d, layer, log_num_shares);
473 layer += 1;
474 }
475 }
476}
477
478fn with_middle_bit(k: usize, shift: usize) -> (usize, usize) {
487 assert!(shift >= 1);
488
489 let ms = k >> (shift - 1);
491 let ls = k & ((1 << shift) - 1);
492
493 let k0 = ls | ((ms & !1) << shift);
494 let k1 = ls | ((ms | 1) << shift);
495
496 (k0, k1)
497}
498
499#[derive(Debug)]
500pub struct NeighborsLastBreadthFirst<DC> {
501 pub domain_context: DC,
503}
504
505impl<F, DC> AdditiveNTT for NeighborsLastBreadthFirst<DC>
506where
507 F: BinaryField,
508 DC: DomainContext<Field = F>,
509{
510 type Field = F;
511
512 fn forward_transform<P: PackedField<Scalar = F>>(
513 &self,
514 mut data: FieldSliceMut<'_, P>,
515 skip_early: usize,
516 skip_late: usize,
517 ) {
518 let log_d = data.log_len();
519 if log_d <= P::LOG_WIDTH {
520 let fallback_ntt = NeighborsLastReference {
521 domain_context: &self.domain_context,
522 };
523 return fallback_ntt.forward_transform(data, skip_early, skip_late);
524 }
525
526 input_check(&self.domain_context, log_d, skip_early, skip_late);
527
528 forward_breadth_first(
529 self.domain_context(),
530 data.as_mut(),
531 log_d,
532 0,
533 0,
534 skip_early..(log_d - skip_late),
535 );
536 }
537
538 fn inverse_transform<P: PackedField<Scalar = F>>(
539 &self,
540 _data: FieldSliceMut<'_, P>,
541 _skip_early: usize,
542 _skip_late: usize,
543 ) {
544 todo!()
545 }
546
547 fn domain_context(&self) -> &impl DomainContext<Field = F> {
548 &self.domain_context
549 }
550}
551
552#[derive(Debug)]
563pub struct NeighborsLastSingleThread<DC> {
564 pub domain_context: DC,
566 pub log_base_len: usize,
568}
569
570impl<DC> NeighborsLastSingleThread<DC> {
571 pub const fn new(domain_context: DC) -> Self {
573 Self {
574 domain_context,
575 log_base_len: DEFAULT_LOG_BASE_LEN,
576 }
577 }
578}
579
580impl<DC: DomainContext> AdditiveNTT for NeighborsLastSingleThread<DC> {
581 type Field = DC::Field;
582
583 fn forward_transform<P: PackedField<Scalar = Self::Field>>(
584 &self,
585 mut data: FieldSliceMut<'_, P>,
586 skip_early: usize,
587 skip_late: usize,
588 ) {
589 let log_d = data.log_len();
590 if log_d <= P::LOG_WIDTH {
591 let fallback_ntt = NeighborsLastReference {
592 domain_context: &self.domain_context,
593 };
594 return fallback_ntt.forward_transform(data, skip_early, skip_late);
595 }
596
597 input_check(&self.domain_context, log_d, skip_early, skip_late);
598
599 forward_depth_first(
600 &self.domain_context,
601 data.as_mut(),
602 log_d,
603 0,
604 0,
605 skip_early..(log_d - skip_late),
606 self.log_base_len.max(P::LOG_WIDTH + 1),
608 );
609 }
610
611 fn inverse_transform<P: PackedField<Scalar = Self::Field>>(
612 &self,
613 _data_orig: FieldSliceMut<'_, P>,
614 _skip_early: usize,
615 _skip_late: usize,
616 ) {
617 unimplemented!()
618 }
619
620 fn domain_context(&self) -> &impl DomainContext<Field = DC::Field> {
621 &self.domain_context
622 }
623}
624
625#[derive(Debug)]
636pub struct NeighborsLastMultiThread<DC> {
637 pub domain_context: DC,
639 pub log_base_len: usize,
641 pub log_num_shares: usize,
645}
646
647impl<DC> NeighborsLastMultiThread<DC> {
648 pub const fn new(domain_context: DC, log_num_shares: usize) -> Self {
650 Self {
651 domain_context,
652 log_base_len: DEFAULT_LOG_BASE_LEN,
653 log_num_shares,
654 }
655 }
656}
657
658impl<DC: DomainContext + Sync> AdditiveNTT for NeighborsLastMultiThread<DC> {
659 type Field = DC::Field;
660
661 fn forward_transform<P: PackedField<Scalar = Self::Field>>(
662 &self,
663 mut data: FieldSliceMut<'_, P>,
664 skip_early: usize,
665 skip_late: usize,
666 ) {
667 let log_d = data.log_len();
668 if log_d <= P::LOG_WIDTH {
669 let fallback_ntt = NeighborsLastReference {
670 domain_context: &self.domain_context,
671 };
672 return fallback_ntt.forward_transform(data, skip_early, skip_late);
673 }
674
675 input_check(&self.domain_context, log_d, skip_early, skip_late);
676
677 let maximum_log_num_shares = log_d - P::LOG_WIDTH - 1;
686 let actual_log_num_shares = min(self.log_num_shares, maximum_log_num_shares);
687 let first_independent_layer = actual_log_num_shares;
688
689 let last_layer = log_d - skip_late;
690 let shared_layers = skip_early..min(first_independent_layer, last_layer);
691 let independent_layers = max(first_independent_layer, skip_early)..last_layer;
692
693 forward_shared_layers(
694 &self.domain_context,
695 data.as_mut(),
696 log_d,
697 shared_layers,
698 actual_log_num_shares,
699 );
700
701 let layer = min(independent_layers.start, maximum_log_num_shares);
706 let log_d_chunk = log_d - layer;
707 data.as_mut()
708 .par_chunks_exact_mut(1 << (log_d_chunk - P::LOG_WIDTH))
709 .enumerate()
710 .for_each(|(block, chunk)| {
711 forward_depth_first(
712 &self.domain_context,
713 chunk,
714 log_d_chunk,
715 layer,
716 block,
717 independent_layers.clone(),
718 self.log_base_len,
719 );
720 });
721 }
722
723 fn inverse_transform<P: PackedField<Scalar = Self::Field>>(
724 &self,
725 _data_orig: FieldSliceMut<'_, P>,
726 _skip_early: usize,
727 _skip_late: usize,
728 ) {
729 unimplemented!()
730 }
731
732 fn domain_context(&self) -> &impl DomainContext<Field = DC::Field> {
733 &self.domain_context
734 }
735}
736
737#[cfg(test)]
738mod tests {
739 use binius_field::{PackedGhash1x128b, PackedGhash2x128b, PackedGhash4x128b};
740 use proptest::prelude::*;
741 use rand::{SeedableRng, rngs::StdRng};
742
743 use super::*;
744 use crate::{ntt::domain_context::GaoMateerPreExpanded, test_utils::random_field_buffer};
745
746 fn shared_window<P: PackedField>(
751 log_d: usize,
752 skip_early: usize,
753 skip_late: usize,
754 log_num_shares: usize,
755 ) -> (Range<usize>, usize) {
756 let actual_log_num_shares = min(log_num_shares, log_d - P::LOG_WIDTH - 1);
758 let layers = skip_early..min(actual_log_num_shares, log_d - skip_late);
759 (layers, actual_log_num_shares)
760 }
761
762 fn check_fused_matches_per_layer<P: PackedField<Scalar: BinaryField>>(
764 log_d: usize,
765 skip_early: usize,
766 skip_late: usize,
767 log_num_shares: usize,
768 seed: u64,
769 ) {
770 let (layers, actual_log_num_shares) =
771 shared_window::<P>(log_d, skip_early, skip_late, log_num_shares);
772 if layers.is_empty() {
773 return;
774 }
775
776 let domain_context = GaoMateerPreExpanded::<P::Scalar>::generate(log_d);
777 let mut rng = StdRng::seed_from_u64(seed);
778
779 let mut fused = random_field_buffer::<P>(&mut rng, log_d);
781 let mut per_layer = fused.clone();
782
783 forward_shared_layers(
784 &domain_context,
785 fused.as_mut(),
786 log_d,
787 layers.clone(),
788 actual_log_num_shares,
789 );
790 for layer in layers {
791 forward_shared_layer(
792 &domain_context,
793 per_layer.as_mut(),
794 log_d,
795 layer,
796 actual_log_num_shares,
797 );
798 }
799
800 assert_eq!(fused, per_layer);
801 }
802
803 fn check_transform_matches_reference<P: PackedField<Scalar: BinaryField>>(
805 log_d: usize,
806 skip_early: usize,
807 skip_late: usize,
808 log_num_shares: usize,
809 seed: u64,
810 ) {
811 let domain_context = GaoMateerPreExpanded::<P::Scalar>::generate(log_d);
812 let mut rng = StdRng::seed_from_u64(seed);
813
814 let mut actual = random_field_buffer::<P>(&mut rng, log_d);
815 let mut expected = actual.clone();
816
817 let multi_thread = NeighborsLastMultiThread {
818 domain_context: &domain_context,
819 log_base_len: 3,
820 log_num_shares,
821 };
822 let reference = NeighborsLastReference {
823 domain_context: &domain_context,
824 };
825
826 multi_thread.forward_transform(actual.as_mut_view(), skip_early, skip_late);
827 reference.forward_transform(expected.as_mut_view(), skip_early, skip_late);
828
829 assert_eq!(actual, expected);
830 }
831
832 fn sweep_window_widths<P: PackedField<Scalar: BinaryField>>(log_d: usize, seed: u64) {
837 for skip_early in 0..4 {
838 for width in 1..=5 {
839 if skip_early + width > log_d {
841 continue;
842 }
843 check_fused_matches_per_layer::<P>(log_d, skip_early, 0, skip_early + width, seed);
845 }
846 }
847 }
848
849 #[test]
850 fn fused_shared_layers_match_per_layer_over_window_widths() {
851 for log_d in 6..13 {
857 sweep_window_widths::<PackedGhash1x128b>(log_d, 0);
858 sweep_window_widths::<PackedGhash2x128b>(log_d, 1);
859 sweep_window_widths::<PackedGhash4x128b>(log_d, 2);
860 }
861 }
862
863 #[test]
864 fn fused_shared_layers_match_per_layer_at_the_narrowest_plane() {
865 check_fused_matches_per_layer::<PackedGhash4x128b>(8, 1, 0, 5, 7);
870
871 check_fused_matches_per_layer::<PackedGhash4x128b>(8, 0, 0, 5, 8);
876 }
877
878 #[test]
879 fn multi_thread_transform_matches_the_reference() {
880 for log_d in [6, 9, 12] {
882 for skip_early in [0, 1, 2, 4] {
883 for skip_late in [0, 1, 3] {
884 if skip_early + skip_late > log_d {
885 continue;
886 }
887 for log_num_shares in [0, 1, 2, 3, 5, 1000] {
888 check_transform_matches_reference::<PackedGhash1x128b>(
889 log_d,
890 skip_early,
891 skip_late,
892 log_num_shares,
893 0,
894 );
895 check_transform_matches_reference::<PackedGhash4x128b>(
896 log_d,
897 skip_early,
898 skip_late,
899 log_num_shares,
900 1,
901 );
902 }
903 }
904 }
905 }
906 }
907
908 proptest! {
909 #[test]
910 fn prop_fused_shared_layers_match_per_layer(
911 log_d in 6..13usize,
912 skip_early in 0..5usize,
913 skip_late in 0..4usize,
914 log_num_shares in 0..8usize,
915 seed: u64,
916 ) {
917 prop_assume!(skip_early + skip_late <= log_d);
920
921 check_fused_matches_per_layer::<PackedGhash1x128b>(
922 log_d, skip_early, skip_late, log_num_shares, seed,
923 );
924 check_fused_matches_per_layer::<PackedGhash2x128b>(
925 log_d, skip_early, skip_late, log_num_shares, seed,
926 );
927 check_fused_matches_per_layer::<PackedGhash4x128b>(
928 log_d, skip_early, skip_late, log_num_shares, seed,
929 );
930 }
931
932 #[test]
933 fn prop_multi_thread_transform_matches_the_reference(
934 log_d in 1..11usize,
935 skip_early in 0..5usize,
936 skip_late in 0..4usize,
937 log_num_shares in 0..8usize,
938 seed: u64,
939 ) {
940 prop_assume!(skip_early + skip_late <= log_d);
943
944 check_transform_matches_reference::<PackedGhash1x128b>(
945 log_d, skip_early, skip_late, log_num_shares, seed,
946 );
947 check_transform_matches_reference::<PackedGhash4x128b>(
948 log_d, skip_early, skip_late, log_num_shares, seed,
949 );
950 }
951 }
952
953 #[test]
954 fn test_with_middle_bit() {
955 assert_eq!(with_middle_bit(0b000, 1), (0b0000, 0b0010));
956 assert_eq!(with_middle_bit(0b000, 2), (0b0000, 0b0100));
957 assert_eq!(with_middle_bit(0b000, 3), (0b0000, 0b1000));
958
959 assert_eq!(with_middle_bit(0b111, 1), (0b1101, 0b1111));
960 assert_eq!(with_middle_bit(0b111, 2), (0b1011, 0b1111));
961 assert_eq!(with_middle_bit(0b111, 3), (0b0111, 0b1111));
962
963 assert_eq!(with_middle_bit(0b1110110, 2), (0b11101010, 0b11101110));
964 }
965}