binius_hash_prover/
parallel_compression.rs1use std::{array, fmt::Debug, mem::MaybeUninit};
5
6use binius_hash::CompressionFunction;
7use binius_utils::rayon::prelude::*;
8
9pub trait ParallelPseudoCompression<T, const N: usize> {
19 type Compression: CompressionFunction<T, N>;
21
22 fn compression(&self) -> &Self::Compression;
24
25 fn parallel_compress(&self, inputs: &[T], out: &mut [MaybeUninit<T>]);
46}
47
48#[derive(Debug, Clone, Default)]
53pub struct ParallelCompressionAdaptor<C> {
54 compression: C,
55}
56
57impl<C> ParallelCompressionAdaptor<C> {
58 pub const fn new(compression: C) -> Self {
60 Self { compression }
61 }
62}
63
64impl<T, C, const ARITY: usize> ParallelPseudoCompression<T, ARITY> for ParallelCompressionAdaptor<C>
65where
66 T: Clone + Send + Sync,
67 C: CompressionFunction<T, ARITY> + Sync,
68{
69 type Compression = C;
70
71 fn compression(&self) -> &Self::Compression {
72 &self.compression
73 }
74
75 fn parallel_compress(&self, inputs: &[T], out: &mut [MaybeUninit<T>]) {
76 assert_eq!(inputs.len(), ARITY * out.len(), "Input length must be N * output length");
77
78 inputs
79 .par_chunks_exact(ARITY)
80 .zip(out.par_iter_mut())
81 .for_each(|(chunk, output)| {
82 let chunk_array: [T; ARITY] = array::from_fn(|j| chunk[j].clone());
84 let compressed = self.compression.compress(chunk_array);
85 output.write(compressed);
86 });
87 }
88}
89
90#[cfg(test)]
91mod tests {
92 use std::mem::MaybeUninit;
93
94 use rand::prelude::*;
95
96 use super::*;
97
98 #[derive(Clone, Debug)]
100 struct XorCompression;
101
102 impl CompressionFunction<u64, 3> for XorCompression {
103 fn compress(&self, input: [u64; 3]) -> u64 {
104 input[0] ^ input[1] ^ input[2]
105 }
106 }
107
108 #[test]
109 fn test_parallel_compression_adaptor() {
110 let mut rng = StdRng::seed_from_u64(0);
111 let compression = XorCompression;
112 let adaptor = ParallelCompressionAdaptor::new(compression.clone());
113
114 const N: usize = 3;
116 const NUM_CHUNKS: usize = 4;
117 let inputs: Vec<u64> = (0..N * NUM_CHUNKS).map(|_| rng.random()).collect();
118
119 let mut adaptor_output = [MaybeUninit::<u64>::uninit(); NUM_CHUNKS];
121 adaptor.parallel_compress(&inputs, &mut adaptor_output);
122 let adaptor_results: Vec<u64> = adaptor_output
123 .into_iter()
124 .map(|x| unsafe { x.assume_init() })
125 .collect();
126
127 let mut manual_results = Vec::new();
129 for chunk_idx in 0..NUM_CHUNKS {
130 let start = chunk_idx * N;
131 let chunk = [inputs[start], inputs[start + 1], inputs[start + 2]];
132 manual_results.push(compression.compress(chunk));
133 }
134
135 assert_eq!(adaptor_results, manual_results);
137 }
138
139 #[test]
140 #[should_panic(expected = "Input length must be N * output length")]
141 fn test_mismatched_input_length() {
142 let compression = XorCompression;
143 let adaptor = ParallelCompressionAdaptor::new(compression);
144
145 let inputs = vec![1u64, 2, 3, 4]; let mut output = [MaybeUninit::<u64>::uninit(); 2]; adaptor.parallel_compress(&inputs, &mut output);
149 }
150}