dtrees_rs/algorithms/optimal/depth2/
error_minimizer.rs1use 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
11pub 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 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 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 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 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 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}