Skip to main content

binius_hash_prover/
parallel_compression.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{array, fmt::Debug, mem::MaybeUninit};
5
6use binius_hash::CompressionFunction;
7use binius_utils::rayon::prelude::*;
8
9/// A trait for parallel application of N-to-1 compression functions.
10///
11/// This trait enables efficient batch compression operations where multiple N-element
12/// chunks are compressed in parallel. It's particularly useful for constructing hash trees
13/// and Merkle trees where many compression operations need to be performed simultaneously.
14///
15/// The trait is parameterized by:
16/// - `T`: The type of values being compressed (typically hash digests)
17/// - `N`: The arity of the compression function (number of inputs per compression)
18pub trait ParallelPseudoCompression<T, const N: usize> {
19	/// The underlying compression function that performs N-to-1 compression.
20	type Compression: CompressionFunction<T, N>;
21
22	/// Returns a reference to the underlying compression function.
23	fn compression(&self) -> &Self::Compression;
24
25	/// Compresses multiple N-element chunks in parallel.
26	///
27	/// # Arguments
28	/// * `inputs` - A slice containing the values to compress. Must have length `N * out.len()`.
29	/// * `out` - Output buffer where compressed values will be written.
30	///
31	/// # Behavior
32	/// For each index `i` in `0..out.len()`, this method computes:
33	/// ```text
34	/// out[i] = Compression::compress([inputs[i*N], inputs[i*N+1], ..., inputs[i*N+N-1]])
35	/// ```
36	///
37	/// All compressions are performed in parallel for efficiency.
38	///
39	/// # Post-conditions
40	/// After this method returns, all elements in `out` will be initialized with the
41	/// compressed values from their corresponding N-element chunks in `inputs`.
42	///
43	/// # Panics
44	/// Panics if `inputs.len() != N * out.len()`.
45	fn parallel_compress(&self, inputs: &[T], out: &mut [MaybeUninit<T>]);
46}
47
48/// A simple adapter that wraps any `CompressionFunction` to implement `ParallelCompression`.
49///
50/// This adapter provides a straightforward way to use existing compression functions
51/// in parallel contexts by applying them sequentially to each N-element chunk.
52#[derive(Debug, Clone, Default)]
53pub struct ParallelCompressionAdaptor<C> {
54	compression: C,
55}
56
57impl<C> ParallelCompressionAdaptor<C> {
58	/// Creates a new adapter wrapping the given compression function.
59	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				// Convert slice to array for compression function
83				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	// Simple test compression function that XORs all inputs
99	#[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		// Test with 4 chunks of 3 elements each
115		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		// Use the adaptor
120		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		// Manually compress each chunk
128		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		// Results should be identical
136		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]; // 4 elements
146		let mut output = [MaybeUninit::<u64>::uninit(); 2]; // Expecting 6 elements (2 * 3)
147
148		adaptor.parallel_compress(&inputs, &mut output);
149	}
150}