orx_parallel/extendable/par_extend_impl/
btree_set.rs1use 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 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 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 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 fn add_one(&mut self, value: T) {
94 _ = self.insert(value);
95 }
96
97 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#[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}