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}