Skip to main content

orx_parallel/extendable/par_extend_impl/
btree_set.rs

1use crate::extendable::par_extend_core::ParExtendCore;
2use crate::extendable::par_extend_impl::utils::{ColAndPos, IdxLen, NextN};
3use alloc::collections::BTreeSet;
4use alloc::{vec, vec::Vec};
5use orx_priority_queue::{BinaryHeap, PriorityQueue};
6
7impl<T: Ord + Send> ParExtendCore<T> for BTreeSet<T> {
8    type ThreadValues = Self;
9
10    type OrderedThreadValues = ColAndPos<Self>;
11
12    fn new_thread_values() -> Self::ThreadValues {
13        Default::default()
14    }
15
16    fn new_ordered_thread_values() -> Self::OrderedThreadValues {
17        Default::default()
18    }
19
20    // thread collect
21
22    fn add_thread_value(collected: &mut Self::ThreadValues, value: T) {
23        _ = collected.insert(value);
24    }
25
26    fn add_thread_values(collected: &mut Self::ThreadValues, values: impl IntoIterator<Item = T>) {
27        collected.extend(values)
28    }
29
30    fn add_ordered_thread_value(collected: &mut Self::OrderedThreadValues, idx: usize, value: T) {
31        let inserted = collected.values.insert(value);
32        if inserted {
33            collected.positions.push(IdxLen { idx, len: 1 });
34        }
35    }
36
37    fn add_ordered_thread_values(
38        collected: &mut Self::OrderedThreadValues,
39        idx: usize,
40        values: impl IntoIterator<Item = T>,
41    ) {
42        let len_before = collected.values.len();
43        collected.values.extend(values);
44
45        let len = collected.values.len() - len_before;
46        if len > 0 {
47            collected.positions.push(IdxLen { idx, len });
48        }
49    }
50
51    // opt: thread collect
52
53    fn add_ordered_thread_optionals(
54        collected: &mut Self::OrderedThreadValues,
55        idx: usize,
56        values: impl IntoIterator<Item = Option<T>>,
57    ) -> Option<()> {
58        let len_begin = collected.values.len();
59        for value in values {
60            _ = collected.values.insert(value?);
61        }
62
63        let len = collected.values.len() - len_begin;
64        if len > 0 {
65            collected.positions.push(IdxLen { idx, len });
66        }
67
68        Some(())
69    }
70
71    // res: thread collect
72
73    fn add_ordered_thread_fallibles<E>(
74        collected: &mut Self::OrderedThreadValues,
75        idx: usize,
76        values: impl IntoIterator<Item = Result<T, E>>,
77    ) -> Result<(), E> {
78        let len_begin = collected.values.len();
79        for value in values {
80            collected.values.insert(value?);
81        }
82
83        let len = collected.values.len() - len_begin;
84        if len > 0 {
85            collected.positions.push(IdxLen { idx, len });
86        }
87
88        Ok(())
89    }
90
91    // add
92
93    fn add_one(&mut self, value: T) {
94        _ = self.insert(value);
95    }
96
97    // extend - merge
98
99    fn extend_merge_infallibles(&mut self, results: Vec<Self::ThreadValues>) {
100        for result in results {
101            self.extend(result);
102        }
103    }
104
105    fn extend_merge_ordered_infallibles(&mut self, results: Vec<Self::OrderedThreadValues>) {
106        let outer_len = results.len();
107        let mut all_values = Vec::with_capacity(outer_len);
108        let mut all_positions = Vec::with_capacity(outer_len);
109
110        for x in results {
111            all_values.push(x.values.into_iter());
112            all_positions.push(x.positions);
113        }
114
115        let mut queue = BinaryHeap::with_capacity(outer_len);
116        let mut pos_indices = vec![0; outer_len];
117        for (th, positions) in all_positions.iter().enumerate() {
118            if let Some(pos) = positions.first() {
119                let node = ThLen::new(th, pos.len);
120                queue.push(node, pos.idx);
121            }
122        }
123        let mut curr_t = queue.pop_node();
124
125        while let Some(ThLen { th, len }) = curr_t {
126            let chunk = NextN::new(&mut all_values[th], len);
127            self.extend(chunk);
128
129            pos_indices[th] += 1;
130            curr_t = match all_positions[th].get(pos_indices[th]) {
131                Some(pos) => {
132                    let node = ThLen::new(th, pos.len);
133                    Some(queue.push_then_pop(node, pos.idx).0)
134                }
135                None => queue.pop_node(),
136            };
137        }
138    }
139}
140
141// merge helpers
142
143#[derive(Clone)]
144struct ThLen {
145    th: usize,
146    len: usize,
147}
148
149impl ThLen {
150    #[inline(always)]
151    fn new(th: usize, len: usize) -> Self {
152        Self { th, len }
153    }
154}