Skip to main content

dtrees_rs/algorithms/optimal/depth2/
error_minimizer.rs

1use crate::algorithms::common::errors::{ErrorWrapper, NativeError};
2use crate::algorithms::common::types::FitError;
3use crate::algorithms::common::utils::{
4    build_labels_count_distribution_matrix, deduce_sibling_error, deduce_sibling_error_with_buffer,
5};
6use crate::algorithms::optimal::depth2::OptimalDepth2Tree;
7use crate::cover::Cover;
8use crate::globals::{float_is_null, item};
9use crate::tree::Tree;
10
11/// Depth-2 solver that minimises the error of the tree: the optimal tree of
12/// depth at most 2.
13pub struct ErrorMinimizer<E>
14where
15    E: ErrorWrapper + ?Sized,
16{
17    error_fn: Box<E>,
18}
19
20impl Default for ErrorMinimizer<NativeError> {
21    fn default() -> Self {
22        Self {
23            error_fn: Box::<NativeError>::default(),
24        }
25    }
26}
27
28impl<E> OptimalDepth2Tree for ErrorMinimizer<E>
29where
30    E: ErrorWrapper + ?Sized,
31{
32    fn find_optimal_depth_one_tree(
33        &self,
34        min_sup: usize,
35        cover: &mut Cover,
36        provided_candidates: Option<&[usize]>,
37    ) -> Result<Tree, FitError> {
38        let candidates = self.get_candidates(cover, min_sup, provided_candidates);
39
40        if candidates.is_empty() {
41            return Err(FitError::EmptyCandidates);
42        }
43
44        let parent_labels_count = cover.labels_count();
45        let mut tree = Tree::empty_tree(1);
46        let mut left_index = 0;
47        let mut right_index = 0;
48
49        if let Some(root) = tree.get_node_mut(tree.get_root_index()) {
50            left_index = root.left;
51            right_index = root.right
52        }
53
54        let mut best_error = <f64>::INFINITY;
55
56        for &candidate in candidates.iter() {
57            let _ = cover.branch_on(item(candidate, 0));
58            let left_labels_count = cover.labels_count();
59            let left_error = self.error_fn.compute(&left_labels_count);
60            cover.backtrack();
61
62            let right_labels_count = deduce_sibling_error(&parent_labels_count, &left_labels_count);
63            let right_error = self.error_fn.compute(&right_labels_count);
64
65            let total_error = left_error.0 + right_error.0;
66
67            if total_error < best_error {
68                best_error = total_error;
69                tree.update_root()
70                    .map(|updater| updater.test(candidate).error(total_error));
71                tree.update_node(left_index)
72                    .map(|updater| updater.error(left_error.0).output(left_error.1));
73                tree.update_node(right_index)
74                    .map(|updater| updater.error(right_error.0).output(right_error.1));
75            }
76        }
77
78        if tree.root_error().is_infinite() {
79            return Err(FitError::EmptyTree);
80        }
81
82        Ok(tree)
83    }
84
85    fn find_optimal_depth_two_tree(
86        &self,
87        min_sup: usize,
88        cover: &mut Cover,
89        provided_candidates: Option<&[usize]>,
90    ) -> Result<Tree, FitError> {
91        let candidates = self.get_candidates(cover, min_sup, provided_candidates);
92        if candidates.is_empty() {
93            return Err(FitError::EmptyCandidates);
94        }
95        if candidates.len() < 2 {
96            return self.find_optimal_depth_one_tree(min_sup, cover, Some(&candidates));
97        }
98
99        let matrix = build_labels_count_distribution_matrix(cover, &candidates);
100
101        let mut best_tree = Tree::empty_tree(2);
102
103        let classes_distribution = cover.labels_count();
104        let total_support = cover.count();
105        let base_error = self.error_fn.compute(&classes_distribution);
106
107        let mut left_distribution = vec![0; cover.num_labels];
108
109        for (i, &first_attribute) in candidates.iter().enumerate() {
110            let right_distribution = &matrix[i][i];
111            let right_support = right_distribution.iter().sum::<usize>();
112
113            deduce_sibling_error_with_buffer(
114                &classes_distribution,
115                right_distribution,
116                &mut left_distribution,
117            );
118            let left_support = total_support - right_support;
119
120            if left_support < min_sup || right_support < min_sup {
121                continue;
122            }
123
124            let mut candidate_tree = Tree::empty_tree(2);
125
126            let (left_index, right_index) =
127                candidate_tree.update_root().map_or((0, 0), |updater| {
128                    updater
129                        .test(first_attribute)
130                        .error(base_error.0)
131                        .get_children()
132                });
133
134            let left_error = self.error_fn.compute(&left_distribution);
135
136            // Too few instances on the left to split it again.
137            if left_support < 2 * min_sup {
138                candidate_tree
139                    .update_node(left_index)
140                    .map(|updater| updater.error(left_error.0).output(left_error.1).leaf());
141
142                if best_tree.root_error() < left_error.0 {
143                    continue;
144                }
145            } else {
146                let mut feature_error = candidate_tree.root_error();
147                if left_error.0 < feature_error {
148                    candidate_tree
149                        .update_node(left_index)
150                        .map(|updater| updater.error(left_error.0));
151                }
152
153                for (j, &second_attribute) in candidates.iter().enumerate() {
154                    if i == j {
155                        continue;
156                    }
157
158                    // Left child of `i`, split on `j`.
159                    let i_left_j_right_classes_support =
160                        deduce_sibling_error(&matrix[j][j], &matrix[i][j]);
161                    let j_right_support = matrix[j][j].iter().sum::<usize>();
162                    let i_right_j_right_support = matrix[i][j].iter().sum::<usize>();
163                    let i_left_j_right_support = j_right_support - i_right_j_right_support;
164                    let i_left_j_left_support = left_support - i_left_j_right_support;
165
166                    if i_left_j_right_support < min_sup || i_left_j_left_support < min_sup {
167                        continue;
168                    }
169
170                    let right_leaf_error = self.error_fn.compute(&i_left_j_right_classes_support);
171
172                    if right_leaf_error.0 >= feature_error {
173                        continue;
174                    }
175
176                    let i_left_j_left_classes_support =
177                        deduce_sibling_error(&left_distribution, &i_left_j_right_classes_support);
178                    let left_leaf_error = self.error_fn.compute(&i_left_j_left_classes_support);
179
180                    let branch_error = left_leaf_error.0 + right_leaf_error.0;
181                    if branch_error >= feature_error {
182                        continue;
183                    }
184                    feature_error = branch_error;
185
186                    let (left_leaf_index, right_leaf_index) = candidate_tree
187                        .update_node(left_index)
188                        .map_or((0, 0), |updater| {
189                            updater
190                                .test(second_attribute)
191                                .error(feature_error)
192                                .get_children()
193                        });
194
195                    candidate_tree
196                        .update_leaf_node(left_leaf_index, left_leaf_error)
197                        .update_leaf_node(right_leaf_index, right_leaf_error);
198
199                    if float_is_null(feature_error) {
200                        break;
201                    }
202                }
203            }
204
205            let right_error = self.error_fn.compute(right_distribution);
206            if right_support < 2 * min_sup {
207                candidate_tree
208                    .update_node(right_index)
209                    .map(|updater| updater.error(right_error.0).output(right_error.1).leaf());
210
211                let best_error = best_tree.root_error();
212                let current_left_error = candidate_tree.node_error(left_index);
213                if current_left_error > best_error
214                    || right_error.0 >= best_error - current_left_error
215                {
216                    // This candidate cannot beat the best tree found so far.
217                    continue;
218                }
219            } else {
220                let mut feature_error = best_tree.root_error();
221                let current_left_error = candidate_tree.node_error(left_index);
222
223                if current_left_error > feature_error
224                    || right_error.0 < feature_error - current_left_error
225                {
226                    candidate_tree
227                        .update_node(right_index)
228                        .map(|updater| updater.error(right_error.0).output(right_error.1));
229                }
230
231                let mut i_right_j_left_classes_support = vec![0; cover.num_labels];
232
233                for (j, &second_attribute) in candidates.iter().enumerate() {
234                    if i == j {
235                        continue;
236                    }
237                    // Right child of `i`, split on `j`.
238                    let i_right_j_right_classes_support = &matrix[i][j];
239                    deduce_sibling_error_with_buffer(
240                        &matrix[i][i],
241                        &matrix[i][j],
242                        &mut i_right_j_left_classes_support,
243                    );
244                    let i_right_j_right_support = matrix[i][j].iter().sum::<usize>();
245                    let i_right_j_left_support = right_support - i_right_j_right_support;
246
247                    if i_right_j_left_support < min_sup || i_right_j_right_support < min_sup {
248                        continue;
249                    }
250
251                    let left_leaf_error = self.error_fn.compute(&i_right_j_left_classes_support);
252
253                    if left_leaf_error.0 >= feature_error {
254                        continue;
255                    }
256
257                    let right_leaf_error = self.error_fn.compute(i_right_j_right_classes_support);
258
259                    let branch_error = left_leaf_error.0 + right_leaf_error.0;
260
261                    if branch_error >= feature_error {
262                        continue;
263                    }
264
265                    feature_error = branch_error;
266
267                    let (left_leaf_index, right_leaf_index) = candidate_tree
268                        .update_node(right_index)
269                        .map_or((0, 0), |updater| {
270                            updater
271                                .test(second_attribute)
272                                .error(feature_error)
273                                .get_children()
274                        });
275
276                    candidate_tree
277                        .update_leaf_node(left_leaf_index, left_leaf_error)
278                        .update_leaf_node(right_leaf_index, right_leaf_error);
279
280                    if float_is_null(feature_error) {
281                        break;
282                    }
283                }
284
285                let feature_error =
286                    candidate_tree.node_error(left_index) + candidate_tree.node_error(right_index);
287                candidate_tree
288                    .update_root()
289                    .map(|updater| updater.error(feature_error));
290
291                if best_tree.root_error() > feature_error {
292                    best_tree = candidate_tree
293                }
294                if float_is_null(feature_error) {
295                    break;
296                }
297            }
298        }
299
300        if best_tree.root_error().is_infinite() {
301            return Err(FitError::EmptyTree);
302        }
303        Ok(best_tree)
304    }
305
306    fn error(&self, distribution: &[usize]) -> (f64, f64) {
307        self.error_fn.compute(distribution)
308    }
309}
310
311impl<E> ErrorMinimizer<E>
312where
313    E: ErrorWrapper + ?Sized,
314{
315    /// A solver using `error_function` for the error of a leaf.
316    pub fn new(error_function: Box<E>) -> Self {
317        Self {
318            error_fn: error_function,
319        }
320    }
321}
322
323#[cfg(test)]
324mod tests {
325    use crate::algorithms::optimal::depth2::error_minimizer::ErrorMinimizer;
326    use crate::algorithms::optimal::depth2::OptimalDepth2Tree;
327    use crate::reader::data_reader::DataReader;
328    use std::path::Path;
329
330    #[test]
331    fn run_small_data() {
332        let reader = DataReader::default();
333        let path = Path::new("test_data/anneal.txt");
334        let cover_result = reader.read_file(path);
335
336        let mut cover = cover_result.expect("the test data is readable");
337
338        let error_minimizer = ErrorMinimizer::default();
339        let tree = error_minimizer.fit(1, 2, &mut cover, None);
340
341        if let Ok(t) = tree {
342            println!("Error {}", t.root_error());
343            println!("{}", t)
344        }
345    }
346}