Skip to main content

rten_parallel/
par_iter.rs

1//! [`ParIter`] utility to simplify implementing Rayon's parallel iterator traits.
2use rayon::iter::plumbing::{Consumer, Producer, ProducerCallback, UnindexedConsumer, bridge};
3use rayon::prelude::*;
4
5use rten_base::iter::SplitIterator;
6
7/// Wraps a splittable iterator to implement Rayon's parallel iterator traits.
8///
9/// This type makes it easy to implement Rayon's parallel iterator traits
10/// ([`ParallelIterator`], [`IndexedParallelIterator`]) for custom iterators.
11/// Adding Rayon support to an iterator using this type requires implementing
12/// the following traits for the iterator:
13///
14/// 1. [`DoubleEndedIterator`] and [`ExactSizeIterator`]. These are requirements
15///    for Rayon's [`IndexedParallelIterator`]. The [`DoubleEndedIterator`]
16///    implementation is not actually needed for most uses, so a simple but
17///    slow implementation will suffice.
18///    See <https://github.com/rayon-rs/rayon/issues/1053>.
19/// 2. [`SplitIterator`] to define how Rayon should split the iterator
20///
21/// With these traits implemented, a parallel iterator can be created using
22/// `ParIter::from(iter)`. For improved ergonomics, you can also implement
23/// Rayon's [`IntoParallelIterator`] trait using [`ParIter<I>`] as the parallel
24/// iterator associated type.
25///
26/// # Example
27///
28/// This is a minimal example showing how to add Rayon support to a custom
29/// iterator type.
30///
31/// ```
32/// use rayon::iter::{ParallelIterator, IntoParallelIterator};
33///
34/// use rten_base::iter::SplitIterator;
35/// use rten_parallel::par_iter::ParIter;
36///
37/// #[derive(Clone)]
38/// struct CustomRange {
39///     start: u32,
40///     end: u32,
41/// }
42///
43/// impl CustomRange {
44///     fn new(start: u32, end: u32) -> Self {
45///         assert!(start <= end);
46///         Self { start, end }
47///     }
48/// }
49///
50/// impl Iterator for CustomRange {
51///     type Item = u32;
52///
53///     fn next(&mut self) -> Option<u32> {
54///         if self.start < self.end {
55///             let item = self.start;
56///             self.start += 1;
57///             Some(item)
58///         } else {
59///             None
60///         }
61///     }
62///
63///     fn size_hint(&self) -> (usize, Option<usize>) {
64///         let len = self.end as usize - self.start as usize;
65///         (len, Some(len))
66///     }
67/// }
68///
69/// impl ExactSizeIterator for CustomRange {}
70///
71/// // DoubleEndedIterator is currently necessary, but usually unused. A crude
72/// // but simple implementation will suffice.
73/// impl DoubleEndedIterator for CustomRange {
74///     fn next_back(&mut self) -> Option<Self::Item> {
75///         if self.start < self.end {
76///             let item = self.end - 1;
77///             self.end -= 1;
78///             Some(item)
79///         } else {
80///             None
81///         }
82///     }
83/// }
84///
85/// impl SplitIterator for CustomRange {
86///     fn split_at(self, index: usize) -> (Self, Self) {
87///         assert!(index < self.len());
88///         let left = CustomRange { start: self.start, end: self.start + index as u32 };
89///         let right = CustomRange { start: self.start + index as u32, end: self.end };
90///         debug_assert_eq!(left.len() + right.len(), self.len());
91///         (left, right)
92///     }
93/// }
94///
95/// impl IntoParallelIterator for CustomRange {
96///     type Iter = ParIter<Self>;
97///     type Item = u32;
98///
99///     fn into_par_iter(self) -> Self::Iter {
100///         ParIter::from(self)
101///     }
102/// }
103///
104/// let range = CustomRange::new(0, 100);
105///
106/// // Process items in serial.
107/// let serial_nums: Vec<_> = range.clone().map(|x| x * x).collect();
108///
109/// // Process items in parallel.
110/// let par_nums: Vec<_> = range.into_par_iter().map(|x| x * x).collect();
111/// assert_eq!(par_nums, serial_nums);
112/// ```
113pub struct ParIter<I: SplitIterator>(I);
114
115impl<I: SplitIterator> From<I> for ParIter<I> {
116    fn from(val: I) -> Self {
117        ParIter(val)
118    }
119}
120
121impl<I: SplitIterator + DoubleEndedIterator + Send> ParallelIterator for ParIter<I>
122where
123    <I as Iterator>::Item: Send,
124{
125    type Item = I::Item;
126
127    fn drive_unindexed<C>(self, consumer: C) -> C::Result
128    where
129        C: UnindexedConsumer<Self::Item>,
130    {
131        bridge(self, consumer)
132    }
133
134    fn opt_len(&self) -> Option<usize> {
135        Some(ExactSizeIterator::len(&self.0))
136    }
137}
138
139impl<I: SplitIterator + DoubleEndedIterator + Send> IndexedParallelIterator for ParIter<I>
140where
141    <I as Iterator>::Item: Send,
142{
143    fn drive<C>(self, consumer: C) -> C::Result
144    where
145        C: Consumer<Self::Item>,
146    {
147        bridge(self, consumer)
148    }
149
150    fn len(&self) -> usize {
151        ExactSizeIterator::len(&self.0)
152    }
153
154    fn with_producer<CB>(self, callback: CB) -> CB::Output
155    where
156        CB: ProducerCallback<Self::Item>,
157    {
158        callback.callback(self)
159    }
160}
161
162impl<I: SplitIterator + DoubleEndedIterator + Send> Producer for ParIter<I> {
163    type Item = I::Item;
164
165    type IntoIter = I;
166
167    fn into_iter(self) -> Self::IntoIter {
168        self.0
169    }
170
171    fn split_at(self, index: usize) -> (Self, Self) {
172        let (left_inner, right_inner) = SplitIterator::split_at(self.0, index);
173        (Self(left_inner), Self(right_inner))
174    }
175}
176
177/// Wrapper around either a serial or parallel iterator, returned by
178/// [`MaybeParIter::maybe_par_iter`].
179pub enum MaybeParallel<PI: ParallelIterator, SI: Iterator<Item = PI::Item>> {
180    Serial(SI),
181    Parallel(PI),
182}
183
184impl<PI: ParallelIterator, SI: Iterator<Item = PI::Item>> MaybeParallel<PI, SI> {
185    pub fn for_each<F: Fn(PI::Item) + Send + Sync>(self, f: F) {
186        match self {
187            MaybeParallel::Serial(iter) => iter.for_each(f),
188            MaybeParallel::Parallel(iter) => iter.for_each(f),
189        }
190    }
191}
192
193/// Trait which allows use of Rayon parallelism to be conditionally enabled.
194///
195/// See <https://crates.io/crates/rayon-cond> for a more full-featured alternative.
196pub trait MaybeParIter {
197    type Item;
198    type ParIter: ParallelIterator<Item = Self::Item>;
199    type Iter: Iterator<Item = Self::Item>;
200
201    /// Return an iterator which executes either in serial on the current
202    /// thread, or in parallel in a Rayon thread pool if `parallel` is true.
203    fn maybe_par_iter(self, parallel: bool) -> MaybeParallel<Self::ParIter, Self::Iter>;
204}
205
206impl<Item, I: rayon::iter::IntoParallelIterator<Item = Item> + IntoIterator<Item = Item>>
207    MaybeParIter for I
208{
209    type Item = Item;
210    type ParIter = I::Iter;
211    type Iter = I::IntoIter;
212
213    fn maybe_par_iter(self, parallel: bool) -> MaybeParallel<Self::ParIter, Self::Iter> {
214        if parallel {
215            MaybeParallel::Parallel(self.into_par_iter())
216        } else {
217            MaybeParallel::Serial(self.into_iter())
218        }
219    }
220}
221
222#[cfg(test)]
223mod tests {
224    use std::sync::atomic::{AtomicU32, Ordering};
225
226    use super::MaybeParIter;
227
228    #[test]
229    fn test_maybe_par_iter() {
230        let count = AtomicU32::new(0);
231        (0..1000).maybe_par_iter(false).for_each(|_| {
232            count.fetch_add(1, Ordering::SeqCst);
233        });
234        assert_eq!(count.load(Ordering::SeqCst), 1000);
235
236        let count = AtomicU32::new(0);
237        (0..1000).maybe_par_iter(true).for_each(|_| {
238            count.fetch_add(1, Ordering::SeqCst);
239        });
240        assert_eq!(count.load(Ordering::SeqCst), 1000);
241    }
242}