From eb8c01c0716227b3be7d57b82a5c34adc57b5c77 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Tue, 22 Sep 2026 13:47:00 +0200 Subject: [PATCH 1/2] Fix issue 462. Changed the tree growth loop in base_tree_regressor.rs and decision_tree_classifier.rs. Added two tests. Removed the unused `depth` method in BaseTreeRegressor. --- src/tree/base_tree_regressor.rs | 43 ++++++++++++++++++++++------ src/tree/decision_tree_classifier.rs | 40 ++++++++++++++++++++++---- 2 files changed, 69 insertions(+), 14 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index fb03e01c..04bca6c4 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -68,10 +68,6 @@ impl, Y: Array1> fn parameters(&self) -> &BaseTreeRegressorParameters { self.parameters.as_ref().unwrap() } - /// Get estimate of intercept, return value - fn depth(&self) -> u16 { - self.depth - } } #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] @@ -244,11 +240,11 @@ impl, Y: Array1> visitor_queue.push_back(visitor); } - while base_tree.depth() < base_tree.parameters().max_depth.unwrap_or(u16::MAX) { - match visitor_queue.pop_front() { - Some(node) => base_tree.split(node, mtry, &mut visitor_queue, &mut rng), - None => break, - }; + let max_depth = base_tree.parameters().max_depth.unwrap_or(u16::MAX); + while let Some(node) = visitor_queue.pop_front() { + if node.level < max_depth { + base_tree.split(node, mtry, &mut visitor_queue, &mut rng); + } } Ok(base_tree) @@ -553,6 +549,7 @@ mod tests { use super::*; use crate::linalg::basic::arrays::Array; use crate::linalg::basic::matrix::DenseMatrix; + use crate::metrics::mean_absolute_error; #[test] fn test_fit_on_empty_data_returns_error() { @@ -599,4 +596,32 @@ mod tests { assert!(result.is_err()); assert_eq!(result.err().unwrap().error(), FailedError::ParametersError); } + + #[test] + fn full_depth() { + let x = DenseMatrix::from_2d_vec(&vec![ + vec![1.0_f64], + vec![2.0], + vec![3.0], + vec![4.0], + vec![5.0], + vec![6.0], + ]) + .unwrap(); + let y = vec![1.0f64, 2.0, 6.0, 7.0, 11., 12.]; + + let parameters = BaseTreeRegressorParameters { + max_depth: Some(3), + min_samples_leaf: 1, + min_samples_split: 2, + seed: None, + splitter: Splitter::Best, + }; + + let tree = BaseTreeRegressor::fit(&x, &y, parameters).expect("Fit should work"); + let y_expected = vec![1.0, 2.0, 6.5, 6.5, 11.50, 11.50]; + let y_hat = tree.predict(&x).expect("Predict should work"); + assert_eq!(tree.nodes().len(), 7); + assert!(mean_absolute_error(&y_expected, &y_hat) < 1e-9); + } } diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index 595063a8..a046061d 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -624,11 +624,11 @@ impl, Y: Array1> visitor_queue.push_back(visitor); } - while tree.depth() < tree.parameters().max_depth.unwrap_or(u16::MAX) { - match visitor_queue.pop_front() { - Some(node) => tree.split(node, mtry, &mut visitor_queue, &mut rng), - None => break, - }; + let max_depth = tree.parameters().max_depth.unwrap_or(u16::MAX); + while let Some(node) = visitor_queue.pop_front() { + if node.level < max_depth { + tree.split(node, mtry, &mut visitor_queue, &mut rng); + } } Ok(tree) @@ -967,6 +967,7 @@ mod tests { use super::*; use crate::linalg::basic::arrays::Array; use crate::linalg::basic::matrix::DenseMatrix; + use crate::metrics::accuracy; #[test] fn search_parameters() { @@ -1053,6 +1054,35 @@ mod tests { } } + #[test] + fn full_depth() { + let x = DenseMatrix::from_2d_vec(&vec![ + vec![1.0_f64], + vec![2.0], + vec![3.0], + vec![4.0], + vec![5.0], + vec![6.0], + ]) + .unwrap(); + let y = vec![0, 1, 2, 2, 3, 4]; + + let parameters = DecisionTreeClassifierParameters { + max_depth: Some(3), + min_samples_leaf: 1, + min_samples_split: 1, // Fix later + seed: None, + criterion: SplitCriterion::Gini, + }; + + let tree = DecisionTreeClassifier::fit(&x, &y, parameters).expect("Fit should work"); + let y_hat = tree.predict(&x).expect("Predict should work"); + assert_eq!(tree.nodes().len(), 7); + + // Tree should have 5 out of 6 examples correct + assert!((accuracy(&y, &y_hat) - 5.0 / 6.0).abs() < 1e-9); + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test From 010c702b013cab96a3c8433710f4b7e21a8296ea Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Wed, 23 Sep 2026 10:11:12 +0200 Subject: [PATCH 2/2] Fix issue 463. min_samples_split --- src/tree/decision_tree_classifier.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index a046061d..de516f4f 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -707,7 +707,7 @@ impl, Y: Array1> return false; } - if n <= self.parameters().min_samples_split { + if n < self.parameters().min_samples_split { return false; } @@ -1070,7 +1070,7 @@ mod tests { let parameters = DecisionTreeClassifierParameters { max_depth: Some(3), min_samples_leaf: 1, - min_samples_split: 1, // Fix later + min_samples_split: 2, seed: None, criterion: SplitCriterion::Gini, };