Skip to main content

orx_parallel/extendable/par_extend_impl/
vec.rs

1use crate::extendable::par_extend_core::ParExtendCore;
2use crate::extendable::par_extend_impl::utils::{ColAndPos, IdxLen};
3use alloc::{vec, vec::Vec};
4use orx_priority_queue::{BinaryHeap, PriorityQueue};
5
6impl<T: Send> ParExtendCore<T> for Vec<T> {
7    type ThreadValues = Self;
8
9    type OrderedThreadValues = ColAndPos<Self>;
10
11    fn new_thread_values() -> Self::ThreadValues {
12        Default::default()
13    }
14
15    fn new_ordered_thread_values() -> Self::OrderedThreadValues {
16        Default::default()
17    }
18
19    // thread collect
20
21    fn add_thread_value(collected: &mut Self::ThreadValues, value: T) {
22        collected.push(value);
23    }
24
25    fn add_thread_values(collected: &mut Self::ThreadValues, values: impl IntoIterator<Item = T>) {
26        collected.extend(values)
27    }
28
29    fn add_ordered_thread_value(collected: &mut Self::OrderedThreadValues, idx: usize, value: T) {
30        collected.values.push(value);
31        collected.positions.push(IdxLen { idx, len: 1 });
32    }
33
34    fn add_ordered_thread_values(
35        collected: &mut Self::OrderedThreadValues,
36        idx: usize,
37        values: impl IntoIterator<Item = T>,
38    ) {
39        let len_begin = collected.values.len();
40        collected.values.extend(values);
41
42        let len = collected.values.len() - len_begin;
43        if len > 0 {
44            collected.positions.push(IdxLen { idx, len });
45        }
46    }
47
48    // opt: thread collect
49
50    fn add_ordered_thread_optionals(
51        collected: &mut Self::OrderedThreadValues,
52        idx: usize,
53        values: impl IntoIterator<Item = Option<T>>,
54    ) -> Option<()> {
55        let len_begin = collected.values.len();
56        for value in values {
57            collected.values.push(value?);
58        }
59
60        let len = collected.values.len() - len_begin;
61        if len > 0 {
62            collected.positions.push(IdxLen { idx, len });
63        }
64
65        Some(())
66    }
67
68    // res: thread collect
69
70    fn add_ordered_thread_fallibles<E>(
71        collected: &mut Self::OrderedThreadValues,
72        idx: usize,
73        values: impl IntoIterator<Item = Result<T, E>>,
74    ) -> Result<(), E> {
75        let len_begin = collected.values.len();
76        for value in values {
77            collected.values.push(value?);
78        }
79
80        let len = collected.values.len() - len_begin;
81        if len > 0 {
82            collected.positions.push(IdxLen { idx, len });
83        }
84
85        Ok(())
86    }
87
88    // add
89
90    #[inline(always)]
91    fn add_one(&mut self, value: T) {
92        self.push(value);
93    }
94
95    // extend - merge
96
97    fn extend_merge_infallibles(&mut self, results: Vec<Self::ThreadValues>) {
98        let collected_len: usize = results.iter().map(|x| x.len()).sum();
99        self.reserve(collected_len);
100        for result in results {
101            self.extend(result);
102        }
103    }
104
105    fn extend_merge_ordered_infallibles(&mut self, mut results: Vec<Self::OrderedThreadValues>) {
106        let collected_len: usize = results.iter().map(|x| x.values.len()).sum();
107        self.reserve(collected_len);
108        let initial_len = self.len();
109        let total_len = initial_len + collected_len;
110
111        let mut queue = BinaryHeap::with_capacity(results.len());
112        let mut pos_indices = vec![0; results.len()];
113
114        for (t, vec) in results.iter().enumerate() {
115            if let Some(pos) = vec.positions.first() {
116                let node = ThBegLen::new(t, 0, pos.len);
117                queue.push(node, pos.idx);
118            }
119        }
120        let mut curr_t = queue.pop_node();
121        let mut ptr_dst = unsafe { self.as_mut_ptr().add(initial_len) };
122
123        while let Some(ThBegLen { th, beg, len }) = curr_t {
124            let ptr_src = unsafe { results[th].values.as_ptr().add(beg) };
125            unsafe { ptr_dst.copy_from_nonoverlapping(ptr_src, len) };
126
127            pos_indices[th] += 1;
128            curr_t = match results[th].positions.get(pos_indices[th]) {
129                Some(pos) => {
130                    let beg = beg + len;
131                    let node = ThBegLen::new(th, beg, pos.len);
132                    Some(queue.push_then_pop(node, pos.idx).0)
133                }
134                None => queue.pop_node(),
135            };
136
137            ptr_dst = unsafe { ptr_dst.add(len) };
138        }
139
140        for vec in results.iter_mut() {
141            // SAFETY: this prevents to drop the elements which are already moved to pinned_vec
142            // allocation within vec.capacity() will still be reclaimed; however, as uninitialized memory
143            unsafe { vec.values.set_len(0) };
144        }
145
146        unsafe { self.set_len(total_len) };
147    }
148}
149
150// merge helpers
151
152/// Merge segment metadata for ordered thread values.
153#[derive(Clone)]
154pub struct ThBegLen {
155    /// Thread index.
156    pub th: usize,
157    /// Start offset within the thread buffer.
158    pub beg: usize,
159    /// Segment length.
160    pub len: usize,
161}
162
163impl ThBegLen {
164    /// Creates a new merge segment.
165    #[inline(always)]
166    pub fn new(th: usize, beg: usize, len: usize) -> Self {
167        Self { th, beg, len }
168    }
169}