From 126a44bd832dc128c33549481258bcb60dc50ae6 Mon Sep 17 00:00:00 2001 From: "Lorenzo (Mec-iS)" Date: Sun, 27 Sep 2026 14:36:26 +0100 Subject: [PATCH] refactor(tree,ensemble): address review follow-ups from #467 --- src/ensemble/base_forest_regressor.rs | 2 +- src/ensemble/extra_trees_regressor.rs | 20 ++++++++++++++- src/ensemble/random_forest_regressor.rs | 20 ++++++++++++++- src/tree/base_tree_regressor.rs | 33 ++++++------------------- src/tree/decision_tree_regressor.rs | 19 +++++++++++++- 5 files changed, 65 insertions(+), 29 deletions(-) diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index b80a2f76..768f64c4 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -482,7 +482,7 @@ mod tests { m: None, keep_samples: true, // keep samples used for each tree, so we can check that they are different seed: 42, - bootstrap: true, // No bootstrapping + bootstrap: true, // Use bootstrapping splitter: crate::tree::base_tree_regressor::Splitter::Best, }; diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index ae996ccc..5ae33ebe 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -64,7 +64,25 @@ use crate::error::Failed; use crate::linalg::basic::arrays::{Array1, Array2}; use crate::numbers::basenum::Number; use crate::numbers::floatnum::FloatNumber; -use crate::tree::base_tree_regressor::{Splitter, validate_sample_weights}; +use crate::tree::base_tree_regressor::Splitter; + +/// Validates the sample weights +fn validate_sample_weights(sample_weights: &[f64], n_rows: usize) -> Result<(), Failed> { + if sample_weights.len() != n_rows { + return Err(Failed::fit( + "Number of sample weights must equal number of rows in x", + )); + } + if sample_weights.iter().any(|v| !v.is_finite() || *v < 0.0) { + return Err(Failed::fit( + "Sample weights must be finite and non-negative", + )); + } + if sample_weights.iter().sum::() <= 0.0 { + return Err(Failed::fit("Sum of sample weights must be positive")); + } + Ok(()) +} #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)] diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 9a7dc103..d04fa760 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -55,7 +55,25 @@ use crate::error::Failed; use crate::linalg::basic::arrays::{Array1, Array2}; use crate::numbers::basenum::Number; use crate::numbers::floatnum::FloatNumber; -use crate::tree::base_tree_regressor::{Splitter, validate_sample_weights}; +use crate::tree::base_tree_regressor::Splitter; + +/// Validates the sample weights +fn validate_sample_weights(sample_weights: &[f64], n_rows: usize) -> Result<(), Failed> { + if sample_weights.len() != n_rows { + return Err(Failed::fit( + "Number of sample weights must equal number of rows in x", + )); + } + if sample_weights.iter().any(|v| !v.is_finite() || *v < 0.0) { + return Err(Failed::fit( + "Sample weights must be finite and non-negative", + )); + } + if sample_weights.iter().sum::() <= 0.0 { + return Err(Failed::fit("Sum of sample weights must be positive")); + } + Ok(()) +} #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)] diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 5f330c6a..b833e345 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -181,24 +181,6 @@ fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { } } -/// Validates the sample weights -pub(crate) fn validate_sample_weights(sample_weights: &[f64], n_rows: usize) -> Result<(), Failed> { - if sample_weights.len() != n_rows { - return Err(Failed::fit( - "Number of sample weights must equal number of rows in x", - )); - } - if sample_weights.iter().any(|v| !v.is_finite() || *v < 0.0) { - return Err(Failed::fit( - "Sample weights must be finite and non-negative", - )); - } - if sample_weights.iter().sum::() <= 0.0 { - return Err(Failed::fit("Sum of sample weights must be positive")); - } - Ok(()) -} - impl, Y: Array1> BaseTreeRegressor { @@ -279,7 +261,7 @@ impl, Y: Array1> let mut visitor_queue: LinkedList> = LinkedList::new(); - if base_tree.find_best_cutoff(&mut visitor, mtry, &mut rng) { + if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng) { visitor_queue.push_back(visitor); } @@ -329,6 +311,7 @@ impl, Y: Array1> &mut self, visitor: &mut NodeVisitor<'_, TX, TY, X, Y>, mtry: usize, + mass: f64, rng: &mut impl rand::Rng, ) -> bool { let (_, n_attr) = visitor.x.shape(); @@ -339,10 +322,6 @@ impl, Y: Array1> return false; } - let mass = match visitor.sample_weights { - Some(_) => (0..visitor.samples.len()).map(|i| visitor.mass_of(i)).sum(), - None => n as f64, - }; let sum = self.nodes()[visitor.node].output * mass; let mut variables = (0..n_attr).collect::>(); @@ -539,6 +518,8 @@ impl, Y: Array1> let (n, _) = visitor.x.shape(); let mut tc = 0; let mut fc = 0; + let mut true_mass = 0f64; + let mut false_mass = 0f64; let mut true_samples: Vec = vec![0; n]; for (i, true_sample) in true_samples.iter_mut().enumerate().take(n) { @@ -552,9 +533,11 @@ impl, Y: Array1> { *true_sample = visitor.samples[i]; tc += *true_sample; + true_mass += visitor.mass_of(i); visitor.samples[i] = 0; } else { fc += visitor.samples[i]; + false_mass += visitor.mass_of(i); } } } @@ -588,7 +571,7 @@ impl, Y: Array1> visitor.level + 1, ); - if self.find_best_cutoff(&mut true_visitor, mtry, rng) { + if self.find_best_cutoff(&mut true_visitor, mtry, true_mass, rng) { visitor_queue.push_back(true_visitor); } @@ -602,7 +585,7 @@ impl, Y: Array1> visitor.level + 1, ); - if self.find_best_cutoff(&mut false_visitor, mtry, rng) { + if self.find_best_cutoff(&mut false_visitor, mtry, false_mass, rng) { visitor_queue.push_back(false_visitor); } diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index 92c2480b..d1c91ec8 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -69,7 +69,24 @@ use crate::api::{Predictor, SupervisedEstimator}; use crate::error::Failed; use crate::linalg::basic::arrays::{Array1, Array2}; use crate::numbers::basenum::Number; -use crate::tree::base_tree_regressor::validate_sample_weights; + +/// Validates the sample weights +fn validate_sample_weights(sample_weights: &[f64], n_rows: usize) -> Result<(), Failed> { + if sample_weights.len() != n_rows { + return Err(Failed::fit( + "Number of sample weights must equal number of rows in x", + )); + } + if sample_weights.iter().any(|v| !v.is_finite() || *v < 0.0) { + return Err(Failed::fit( + "Sample weights must be finite and non-negative", + )); + } + if sample_weights.iter().sum::() <= 0.0 { + return Err(Failed::fit("Sum of sample weights must be positive")); + } + Ok(()) +} #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)]