Skip to main content

orx_parallel/extendable/par_extend_impl/
btree_map.rs

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