diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index 51eb2c41..b80a2f76 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -79,9 +79,11 @@ impl, Y: Array1 /// Build a forest of trees from the training set. /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. /// * `y` - the target class values + /// * `sample_weights`: optional sample_weights to use during fitting pub fn fit( x: &X, y: &Y, + sample_weights: Option<&[f64]>, parameters: BaseForestRegressorParameters, ) -> Result, Failed> { let (n_rows, num_attributes) = x.shape(); @@ -111,10 +113,20 @@ impl, Y: Array1 let mut samples: Vec = (0..n_rows).map(|_| 1).collect(); + let dist = sample_weights + .map(|weights| { + rand::distr::weighted::WeightedIndex::new(weights) + .map_err(|e| Failed::fit(&e.to_string())) + }) + .transpose()?; + for _ in 0..parameters.n_trees { if parameters.bootstrap { - samples = - BaseForestRegressor::::sample_with_replacement(n_rows, &mut rng); + samples = BaseForestRegressor::::sample_with_replacement( + n_rows, + &mut rng, + dist.as_ref(), + ); } // keep samples is flag is on @@ -129,7 +141,14 @@ impl, Y: Array1 seed: Some(parameters.seed), splitter: parameters.splitter.clone(), }; - let tree = BaseTreeRegressor::fit_weak_learner(x, y, samples.clone(), mtry, params)?; + let tree = BaseTreeRegressor::fit_weak_learner( + x, + y, + sample_weights, + samples.clone(), + mtry, + params, + )?; trees.push(tree); } @@ -216,18 +235,27 @@ impl, Y: Array1 result / TY::from(n_trees).unwrap() } - fn sample_with_replacement(nrows: usize, rng: &mut impl rand::Rng) -> Vec { + fn sample_with_replacement( + nrows: usize, + rng: &mut impl rand::Rng, + distribution: Option<&rand::distr::weighted::WeightedIndex>, + ) -> Vec { let mut samples = vec![0; nrows]; for _ in 0..nrows { - let xi = rng.random_range(0..nrows); + let xi = match distribution { + Some(dist) => rng.sample(dist), + None => rng.random_range(0..nrows), + }; samples[xi] += 1; } + samples } } #[cfg(test)] mod tests { + use super::*; use crate::linalg::basic::arrays::Array; use crate::linalg::basic::matrix::DenseMatrix; @@ -247,7 +275,7 @@ mod tests { bootstrap: true, splitter: crate::tree::base_tree_regressor::Splitter::Best, }; - let regressor = BaseForestRegressor::fit(&x, &y, params).unwrap(); + let regressor = BaseForestRegressor::fit(&x, &y, None, params).unwrap(); assert_eq!(regressor.trees.unwrap().len(), 5); assert!(regressor.samples.is_some()); } @@ -263,6 +291,7 @@ mod tests { let result = BaseForestRegressor::fit( &empty, &y, + None, BaseForestRegressorParameters { max_depth: None, min_samples_leaf: 1, @@ -290,6 +319,7 @@ mod tests { let result = BaseForestRegressor::fit( &no_features, &y, + None, BaseForestRegressorParameters { max_depth: None, min_samples_leaf: 1, @@ -305,4 +335,176 @@ mod tests { assert!(result.is_err()); assert_eq!(result.err().unwrap().error(), FailedError::ParametersError); } + + #[test] + fn balance_property() { + // Test the balance property on a random dataset + // Fit random forest with bootstrapping = False. + // The (weighted) average of predictions on the training data should equal the + // weighted average of the actual targets of the training data + let x: DenseMatrix = DenseMatrix::rand(1000, 10); + let model_parameters = (0..=9).map(|x| x as f64).collect::>(); + let y: Vec = model_parameters.xa(true, &x); + + let forest_parameters = BaseForestRegressorParameters { + max_depth: None, + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 5, + m: None, + keep_samples: true, + seed: 42, + bootstrap: false, + splitter: crate::tree::base_tree_regressor::Splitter::Best, + }; + + let forest = BaseForestRegressor::fit(&x, &y, None, forest_parameters.clone()) + .expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + assert!((y_hat.iter().sum::() - y.iter().sum::()).abs() < 1e-9); + + // Seeded RNG: the test gives the same result on each run + let mut rng = get_rng_impl(Some(42)); + + // Positive weights in [0.5, 2.0) + let sample_weights: Vec = (0..1000).map(|_| rng.random_range(0.5..2.0)).collect(); + let forest = + BaseForestRegressor::fit(&x, &y, Some(&sample_weights), forest_parameters.clone()) + .expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + let s: f64 = sample_weights.iter().sum(); + let normalized_weights = sample_weights.iter().map(|w| *w / s).collect::>(); + + let expected = y + .iter() + .zip(normalized_weights.iter()) + .map(|(&yi, &w)| yi * w) + .sum::(); + let actual = y_hat + .iter() + .zip(normalized_weights.iter()) + .map(|(&yi, &w)| yi * w) + .sum::(); + + assert!((expected - actual).abs() < 1e-9); + } + + #[test] + fn test_each_tree_gets_different_bootstrap_sample() { + // actual data uses is irrelevant + let n_rows = 100; + let x: DenseMatrix = + DenseMatrix::from_iterator((0..2 * n_rows).map(|k| k as f64), n_rows, 2, 0); + let y: Vec = (0..n_rows).map(|i| i as f64).collect(); + let sample_weights: Vec = (0..n_rows).map(|i| 1.0 + (i % 4) as f64).collect(); + + for weights in [None, Some(sample_weights.as_slice())] { + let params = BaseForestRegressorParameters { + max_depth: Some(1), + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 10, + m: None, + keep_samples: true, // keep samples used for each tree, so we can check that they are different + seed: 42, + bootstrap: true, // Use bootstrapping + splitter: crate::tree::base_tree_regressor::Splitter::Best, + }; + let regressor = BaseForestRegressor::fit(&x, &y, weights, params).unwrap(); + let samples = regressor.samples.unwrap(); + + for (t, in_bag) in samples.iter().enumerate() { + assert!( + in_bag.iter().any(|b| !b), + "tree {t} has no out-of-bag rows (weights: {weights:?})" + ); + for (u, other) in samples.iter().enumerate().skip(t + 1) { + assert_ne!( + in_bag, other, + "trees {t} and {u} have the same bootstrap sample (weights: {weights:?})" + ); + } + } + } + } + + #[test] + fn fit_with_weights_predicts_approx_weighted_mean() { + // 20 rows, 1 feature. Stumps (max_depth = 0): each tree predicts the + // weighted mean of y over its bootstrap sample (which is also weighted) + let x: DenseMatrix = DenseMatrix::from_iterator((0..20).map(|i| i as f64), 20, 1, 0); + let y: Vec = (0..20).map(|i| if i < 10 { 0.0 } else { 10.0 }).collect(); + // Rows with y = 10 have weight 9, rows with y = 0 have weight 1. + let sample_weights: Vec = (0..20).map(|i| if i < 10 { 1.0 } else { 9.0 }).collect(); + + let parameters = BaseForestRegressorParameters { + max_depth: Some(1), + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 10, + m: None, + keep_samples: true, // keep samples used for each tree, so we can check that they are different + seed: 42, + bootstrap: false, // No bootstrapping + splitter: crate::tree::base_tree_regressor::Splitter::Best, + }; + + let forest = BaseForestRegressor::fit(&x, &y, Some(&sample_weights), parameters.clone()) + .expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + // weighted mean is (10.0 * 0 + 90*10) / 100 = 9 + for p in y_hat.iter() { + assert!( + (p - 9.0f64).abs() < 1e-9, + "expected value very close to 9, got {p}" + ); + } + + // Without weights, the predicted value should be close to 5 + let forest = BaseForestRegressor::fit(&x, &y, None, parameters).expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + for p in y_hat.iter() { + assert!( + (p - 5.0).abs() < 1e-9, + "expected value very close to 5, got {p}" + ); + } + + // Use bootstrapping + let parameters = BaseForestRegressorParameters { + max_depth: Some(1), + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 500, // Use more trees than before to smooth out randomness + 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 + splitter: crate::tree::base_tree_regressor::Splitter::Best, + }; + + let forest = BaseForestRegressor::fit(&x, &y, Some(&sample_weights), parameters.clone()) + .expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + // weighted mean is (10.0 * 0 + 90*10) / 100 = 9 + for p in y_hat.iter() { + assert!(p > &9.0f64, "expected value well above 9, got {p}"); + } + + // Without weights, the predicted value should be reasonably close to 5 + let forest = BaseForestRegressor::fit(&x, &y, None, parameters).expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + // The bootstrapping makes that we will not be very close to 5 + for p in y_hat.iter() { + assert!( + (p - 5.0).abs() < 0.1, + "expected value reasonably close to 5, got {p}" + ); + } + } } diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index ab0555f7..ae996ccc 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -64,7 +64,7 @@ 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; +use crate::tree::base_tree_regressor::{Splitter, validate_sample_weights}; #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)] @@ -192,6 +192,29 @@ impl, Y: Array1 x: &X, y: &Y, parameters: ExtraTreesRegressorParameters, + ) -> Result, Failed> { + Self::fit_inner(x, y, None, parameters) + } + + /// Build a forest of trees from the training set. + /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. + /// * `y` - the target class values + /// * `sample_weights`: sample_weights to use during fitting + pub fn fit_with_weights( + x: &X, + y: &Y, + sample_weights: &[f64], + parameters: ExtraTreesRegressorParameters, + ) -> Result, Failed> { + validate_sample_weights(sample_weights, x.shape().0)?; + Self::fit_inner(x, y, Some(sample_weights), parameters) + } + + fn fit_inner( + x: &X, + y: &Y, + sample_weights: Option<&[f64]>, + parameters: ExtraTreesRegressorParameters, ) -> Result, Failed> { let regressor_params = BaseForestRegressorParameters { max_depth: parameters.max_depth, @@ -204,7 +227,7 @@ impl, Y: Array1 bootstrap: false, splitter: Splitter::Random, }; - let forest_regressor = BaseForestRegressor::fit(x, y, regressor_params)?; + let forest_regressor = BaseForestRegressor::fit(x, y, sample_weights, regressor_params)?; Ok(ExtraTreesRegressor { forest_regressor: Some(forest_regressor), @@ -262,6 +285,41 @@ mod tests { assert!(mse < 1.0); } + #[test] + fn fit_with_weights_validates_weights() { + let x: DenseMatrix = DenseMatrix::from_iterator((0..6).map(|i| i as f64), 3, 2, 0); + let y = vec![1.0_f64, 2.0, 3.0]; + let parameters = ExtraTreesRegressorParameters::default() + .with_n_trees(5) + .with_seed(42); + + // Valid weights: a zero weight is permitted when the sum is positive + for weights in [vec![1.0, 2.0, 3.0], vec![0.0, 0.0, 0.5]] { + assert!( + ExtraTreesRegressor::fit_with_weights(&x, &y, &weights, parameters.clone()).is_ok(), + "weights: {weights:?}" + ); + } + + let wrong_length = "Number of sample weights must equal number of rows in x"; + let not_finite_or_negative = "Sample weights must be finite and non-negative"; + let zero_sum = "Sum of sample weights must be positive"; + let cases: Vec<(Vec, &str)> = vec![ + (vec![], wrong_length), + (vec![1.0, 2.0], wrong_length), + (vec![1.0, 2.0, 3.0, 4.0], wrong_length), + (vec![1.0, -1.0, 3.0], not_finite_or_negative), + (vec![1.0, f64::NAN, 3.0], not_finite_or_negative), + (vec![1.0, f64::INFINITY, 3.0], not_finite_or_negative), + (vec![0.0, 0.0, 0.0], zero_sum), + ]; + for (weights, msg) in cases { + let result = + ExtraTreesRegressor::fit_with_weights(&x, &y, &weights, parameters.clone()); + assert_eq!(result.err(), Some(Failed::fit(msg)), "weights: {weights:?}"); + } + } + #[test] fn test_fit_predict_higher_dims() { // Dataset with 10 features, but y is only dependent on the 3rd feature (index 2). @@ -316,4 +374,43 @@ mod tests { assert_eq!(y_hat1, y_hat2); } + + #[test] + fn fit_with_weights_predicts_approx_weighted_mean() { + // 20 rows, 1 feature. Stumps (max_depth = 0): each tree predicts the + // weighted mean of y over its bootstrap sample (which is also weighted) + let x: DenseMatrix = DenseMatrix::from_iterator((0..20).map(|i| i as f64), 20, 1, 0); + let y: Vec = (0..20).map(|i| if i < 10 { 0.0 } else { 10.0 }).collect(); + // Rows with y = 10 have weight 9, rows with y = 0 have weight 1. + let sample_weights: Vec = (0..20).map(|i| if i < 10 { 1.0 } else { 9.0 }).collect(); + + let parameters = ExtraTreesRegressorParameters::default() + .with_max_depth(0) + .with_n_trees(50) + .with_seed(42); + + let forest = + ExtraTreesRegressor::fit_with_weights(&x, &y, &sample_weights, parameters.clone()) + .expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + // weighted mean is (10.0 * 0 + 90*10) / 100 = 9 + for p in y_hat.iter() { + assert!( + (p - 9.0f64).abs() < 1e-9, + "expected value very close to 9, got {p}" + ); + } + + // Without weights, the predicted value should be close to 5 + let forest = ExtraTreesRegressor::fit(&x, &y, parameters).expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + for p in y_hat.iter() { + assert!( + (p - 5.0).abs() < 1e-9, + "expected value very close to 5, got {p}" + ); + } + } } diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index f03b2940..9a7dc103 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -55,7 +55,7 @@ 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; +use crate::tree::base_tree_regressor::{Splitter, validate_sample_weights}; #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)] @@ -385,6 +385,29 @@ impl, Y: Array1 x: &X, y: &Y, parameters: RandomForestRegressorParameters, + ) -> Result, Failed> { + Self::fit_inner(x, y, None, parameters) + } + + /// Build a forest of trees from the training set. + /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. + /// * `y` - the target class values + /// * `sample_weights`: sample_weights to use during fitting + pub fn fit_with_weights( + x: &X, + y: &Y, + sample_weights: &[f64], + parameters: RandomForestRegressorParameters, + ) -> Result, Failed> { + validate_sample_weights(sample_weights, x.shape().0)?; + Self::fit_inner(x, y, Some(sample_weights), parameters) + } + + fn fit_inner( + x: &X, + y: &Y, + sample_weights: Option<&[f64]>, + parameters: RandomForestRegressorParameters, ) -> Result, Failed> { let regressor_params = BaseForestRegressorParameters { max_depth: parameters.max_depth, @@ -397,7 +420,7 @@ impl, Y: Array1 bootstrap: true, splitter: Splitter::Best, }; - let forest_regressor = BaseForestRegressor::fit(x, y, regressor_params)?; + let forest_regressor = BaseForestRegressor::fit(x, y, sample_weights, regressor_params)?; Ok(RandomForestRegressor { forest_regressor: Some(forest_regressor), @@ -576,6 +599,110 @@ mod tests { assert!(mean_absolute_error(&y, &y_hat) < mean_absolute_error(&y, &y_hat_oob)); } + #[test] + fn fit_with_weights_validates_weights() { + let x: DenseMatrix = DenseMatrix::from_iterator((0..6).map(|i| i as f64), 3, 2, 0); + let y = vec![1.0_f64, 2.0, 3.0]; + let parameters = RandomForestRegressorParameters::default() + .with_n_trees(5) + .with_seed(42); + + // Valid weights: a zero weight is permitted when the sum is positive + for weights in [vec![1.0, 2.0, 3.0], vec![0.0, 0.0, 0.5]] { + assert!( + RandomForestRegressor::fit_with_weights(&x, &y, &weights, parameters.clone()) + .is_ok(), + "weights: {weights:?}" + ); + } + + let wrong_length = "Number of sample weights must equal number of rows in x"; + let not_finite_or_negative = "Sample weights must be finite and non-negative"; + let zero_sum = "Sum of sample weights must be positive"; + let cases: Vec<(Vec, &str)> = vec![ + (vec![], wrong_length), + (vec![1.0, 2.0], wrong_length), + (vec![1.0, 2.0, 3.0, 4.0], wrong_length), + (vec![1.0, -1.0, 3.0], not_finite_or_negative), + (vec![1.0, f64::NAN, 3.0], not_finite_or_negative), + (vec![1.0, f64::INFINITY, 3.0], not_finite_or_negative), + (vec![0.0, 0.0, 0.0], zero_sum), + ]; + for (weights, msg) in cases { + let result = + RandomForestRegressor::fit_with_weights(&x, &y, &weights, parameters.clone()); + assert_eq!(result.err(), Some(Failed::fit(msg)), "weights: {weights:?}"); + } + } + + #[test] + fn fit_with_weights_predicts_approx_weighted_mean() { + // 20 rows, 1 feature. Stumps (max_depth = 0): each tree predicts the + // weighted mean of y over its bootstrap sample (which is also weighted) + let x: DenseMatrix = DenseMatrix::from_iterator((0..20).map(|i| i as f64), 20, 1, 0); + let y: Vec = (0..20).map(|i| if i < 10 { 0.0 } else { 10.0 }).collect(); + // Rows with y = 10 have weight 9, rows with y = 0 have weight 1. + let sample_weights: Vec = (0..20).map(|i| if i < 10 { 1.0 } else { 9.0 }).collect(); + + let parameters = RandomForestRegressorParameters::default() + .with_max_depth(0) + .with_n_trees(500) + .with_seed(42); + + let forest = + RandomForestRegressor::fit_with_weights(&x, &y, &sample_weights, parameters.clone()) + .expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + // Due to bootstrapping, weighted value should be well above 9. + for p in y_hat.iter() { + assert!(p > &9.0f64, "expected value greater than 9, got {p}"); + } + + // Without weights, the predicted value should be close to 5 + let forest = RandomForestRegressor::fit(&x, &y, parameters).expect("Fit should work"); + let y_hat = forest.predict(&x).expect("Predict should work"); + + // Due to bootstrapping predicted value is not exactly 5 + for p in y_hat.iter() { + assert!((p - 5.0).abs() < 0.1, "expected value around 5, got {p}"); + } + } + + #[test] + fn fit_with_same_seed_is_deterministic() { + // 30 rows, 3 features, deterministic non-linear data + let x: DenseMatrix = DenseMatrix::from_iterator( + (0..90).map(|k| ((k % 17) as f64) / (10.0 + (k % 17) as f64)), + 30, + 3, + 0, + ); + let model_parameters = (1..=3).map(|x| x as f64).collect::>(); + let y: Vec = model_parameters.xa(true, &x); + let sample_weights: Vec = (0..30).map(|i| 1.0 + (i % 4) as f64).collect(); + + let parameters = RandomForestRegressorParameters::default() + .with_n_trees(20) + .with_m(2) // only use 2 attributes + .with_seed(42); + + // Without weights + let forest_1 = RandomForestRegressor::fit(&x, &y, parameters.clone()).unwrap(); + let forest_2 = RandomForestRegressor::fit(&x, &y, parameters.clone()).unwrap(); + assert_eq!(&forest_1, &forest_2); + assert_eq!(forest_1.predict(&x).unwrap(), forest_2.predict(&x).unwrap()); + + // With weights + let forest_1 = + RandomForestRegressor::fit_with_weights(&x, &y, &sample_weights, parameters.clone()) + .unwrap(); + let forest_2 = + RandomForestRegressor::fit_with_weights(&x, &y, &sample_weights, parameters).unwrap(); + assert_eq!(forest_1, forest_2); + assert_eq!(forest_1.predict(&x).unwrap(), forest_2.predict(&x).unwrap()); + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index ceed7e89..5f330c6a 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -131,6 +131,7 @@ struct NodeVisitor<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Ar y: &'a Y, node: usize, samples: Vec, + sample_weights: Option<&'a [f64]>, order: &'a [Vec], true_child_output: f64, false_child_output: f64, @@ -145,6 +146,7 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> fn new( node_id: usize, samples: Vec, + sample_weights: Option<&'a [f64]>, order: &'a [Vec], x: &'a X, y: &'a Y, @@ -155,6 +157,7 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> y, node: node_id, samples, + sample_weights, order, true_child_output: 0f64, false_child_output: 0f64, @@ -163,17 +166,46 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> _phantom_ty: PhantomData, } } + + /// Weighted count of sample `i`. The weight is 1.0 if no weights are given. + fn mass_of(&self, i: usize) -> f64 { + mass_of(i, &self.samples, self.sample_weights) + } +} + +/// Weighted count of sample `i`. The weight is 1.0 if no weights are given. +fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { + match sample_weights { + Some(weights) => samples[i] as f64 * weights[i], + None => samples[i] as 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 { - /// Build a decision base_tree regressor from the training data. - /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. - /// * `y` - the target values - pub fn fit( + pub(crate) fn fit_inner( x: &X, y: &Y, + sample_weights: Option<&[f64]>, parameters: BaseTreeRegressorParameters, ) -> Result, Failed> { let (x_nrows, num_attributes) = x.shape(); @@ -188,12 +220,20 @@ impl, Y: Array1> } let samples = vec![1; x_nrows]; - BaseTreeRegressor::fit_weak_learner(x, y, samples, num_attributes, parameters) + BaseTreeRegressor::fit_weak_learner( + x, + y, + sample_weights, + samples, + num_attributes, + parameters, + ) } pub(crate) fn fit_weak_learner( x: &X, y: &Y, + sample_weights: Option<&[f64]>, samples: Vec, mtry: usize, parameters: BaseTreeRegressorParameters, @@ -206,14 +246,16 @@ impl, Y: Array1> let mut nodes: Vec = Vec::new(); let mut rng = get_rng_impl(parameters.seed); - let mut n = 0; let mut sum = 0f64; - for (i, sample_i) in samples.iter().enumerate().take(y_ncols) { - n += *sample_i; - sum += *sample_i as f64 * y_m.get(i).to_f64().unwrap(); + let mut mass = 0f64; + + for i in 0..y_ncols { + let mass_i = mass_of(i, &samples, sample_weights); + mass += mass_i; + sum += mass_i * y_m.get(i).to_f64().unwrap(); } - let root = Node::new(sum / (n as f64)); + let root = Node::new(sum / mass); nodes.push(root); let mut order: Vec> = Vec::new(); @@ -232,7 +274,8 @@ impl, Y: Array1> _phantom_y: PhantomData, }; - let mut visitor = NodeVisitor::::new(0, samples, &order, x, &y_m, 1); + let mut visitor = + NodeVisitor::::new(0, samples, sample_weights, &order, x, &y_m, 1); let mut visitor_queue: LinkedList> = LinkedList::new(); @@ -296,7 +339,11 @@ impl, Y: Array1> return false; } - let sum = self.nodes()[visitor.node].output * n as f64; + 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::>(); @@ -305,17 +352,17 @@ impl, Y: Array1> } let parent_gain = - n as f64 * self.nodes()[visitor.node].output * self.nodes()[visitor.node].output; + mass * self.nodes()[visitor.node].output * self.nodes()[visitor.node].output; let splitter = self.parameters().splitter.clone(); for variable in variables.iter().take(mtry) { match splitter { Splitter::Random => { - self.find_random_split(visitor, n, sum, parent_gain, *variable, rng); + self.find_random_split(visitor, n, mass, sum, parent_gain, *variable, rng); } Splitter::Best => { - self.find_best_split(visitor, n, sum, parent_gain, *variable); + self.find_best_split(visitor, n, mass, sum, parent_gain, *variable); } } } @@ -327,6 +374,7 @@ impl, Y: Array1> &mut self, visitor: &mut NodeVisitor<'_, TX, TY, X, Y>, n: usize, + mass: f64, sum: f64, parent_gain: f64, j: usize, @@ -360,12 +408,14 @@ impl, Y: Array1> let split_value = rng.random_range(min_val.to_f64().unwrap()..max_val.to_f64().unwrap()); let mut true_sum = 0f64; + let mut true_mass = 0f64; let mut true_count = 0; for &i in &visitor.order[j] { if visitor.samples[i] > 0 { if visitor.x.get((i, j)).to_f64().unwrap() <= split_value { - true_sum += visitor.samples[i] as f64 * visitor.y.get(i).to_f64().unwrap(); + true_sum += visitor.mass_of(i) * visitor.y.get(i).to_f64().unwrap(); true_count += visitor.samples[i]; + true_mass += visitor.mass_of(i); } else { break; } @@ -380,18 +430,18 @@ impl, Y: Array1> return; } - let true_mean = if true_count > 0 { - true_sum / true_count as f64 + let true_mean = if true_mass > 0f64 { + true_sum / true_mass } else { 0.0 }; - let false_mean = if false_count > 0 { - (sum - true_sum) / false_count as f64 + let false_mass = mass - true_mass; + let false_mean = if false_mass > 0f64 { + (sum - true_sum) / false_mass } else { 0.0 }; - let gain = (true_count as f64 * true_mean * true_mean - + false_count as f64 * false_mean * false_mean) + let gain = (true_mass * true_mean * true_mean + false_mass * false_mean * false_mean) - parent_gain; if self.nodes[visitor.node].split_score.is_none() @@ -409,12 +459,14 @@ impl, Y: Array1> &mut self, visitor: &mut NodeVisitor<'_, TX, TY, X, Y>, n: usize, + mass: f64, sum: f64, parent_gain: f64, j: usize, ) { let mut true_sum = 0f64; let mut true_count = 0; + let mut true_mass = 0f64; let mut prevx = Option::None; for i in visitor.order[j].iter() { @@ -424,7 +476,8 @@ impl, Y: Array1> if prevx.is_none() || x_ij == prevx.unwrap() { prevx = Some(x_ij); true_count += visitor.samples[*i]; - true_sum += visitor.samples[*i] as f64 * visitor.y.get(*i).to_f64().unwrap(); + true_mass += visitor.mass_of(*i); + true_sum += visitor.mass_of(*i) * visitor.y.get(*i).to_f64().unwrap(); continue; } @@ -435,15 +488,25 @@ impl, Y: Array1> { prevx = Some(x_ij); true_count += visitor.samples[*i]; - true_sum += visitor.samples[*i] as f64 * visitor.y.get(*i).to_f64().unwrap(); + true_mass += visitor.mass_of(*i); + true_sum += visitor.mass_of(*i) * visitor.y.get(*i).to_f64().unwrap(); continue; } - let true_mean = true_sum / true_count as f64; - let false_mean = (sum - true_sum) / false_count as f64; + let true_mean = if true_mass > 0.0 { + true_sum / true_mass + } else { + 0.0 + }; + let false_mass = mass - true_mass; + let false_mean = if false_mass > 0.0 { + (sum - true_sum) / false_mass + } else { + 0.0 + }; - let gain = (true_count as f64 * true_mean * true_mean - + false_count as f64 * false_mean * false_mean) + let gain = (true_mass * true_mean * true_mean + + false_mass * false_mean * false_mean) - parent_gain; if self.nodes()[visitor.node].split_score.is_none() @@ -459,8 +522,9 @@ impl, Y: Array1> } prevx = Some(x_ij); - true_sum += visitor.samples[*i] as f64 * visitor.y.get(*i).to_f64().unwrap(); + true_sum += visitor.mass_of(*i) * visitor.y.get(*i).to_f64().unwrap(); true_count += visitor.samples[*i]; + true_mass += visitor.mass_of(*i); } } } @@ -517,6 +581,7 @@ impl, Y: Array1> let mut true_visitor = NodeVisitor::::new( true_child_idx, true_samples, + visitor.sample_weights, visitor.order, visitor.x, visitor.y, @@ -530,6 +595,7 @@ impl, Y: Array1> let mut false_visitor = NodeVisitor::::new( false_child_idx, visitor.samples, + visitor.sample_weights, visitor.order, visitor.x, visitor.y, @@ -559,9 +625,10 @@ mod tests { assert_eq!(empty.shape(), (0, 2)); let y: Vec = vec![]; - let result = BaseTreeRegressor::fit( + let result = BaseTreeRegressor::fit_inner( &empty, &y, + None, BaseTreeRegressorParameters { max_depth: None, min_samples_leaf: 1, @@ -582,9 +649,10 @@ mod tests { assert_eq!(no_features.shape(), (2, 0)); let y = vec![1.0_f64, 2.0]; - let result = BaseTreeRegressor::fit( + let result = BaseTreeRegressor::fit_inner( &no_features, &y, + None, BaseTreeRegressorParameters { max_depth: None, min_samples_leaf: 1, @@ -597,6 +665,118 @@ mod tests { assert_eq!(result.err().unwrap().error(), FailedError::ParametersError); } + #[test] + fn root_prediction_is_weighted_mean() { + // Create a tree with no splits. Assert that the prediction is the weighted mean of the targets. + + // X-values are arbitrary. 3 examples + let x = DenseMatrix::from_2d_vec(&vec![vec![1.0_f64, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]]) + .unwrap(); + let y = vec![1.0, 2.0, 3.0]; + let sample_weights = vec![5.0, 6.0, 7.0]; + + let result = BaseTreeRegressor::fit_inner( + &x, + &y, + Some(&sample_weights), + BaseTreeRegressorParameters { + max_depth: Some(0), + min_samples_leaf: 1, + min_samples_split: 2, + seed: None, + splitter: Splitter::Best, + }, + ); + assert!(result.is_ok()); + let tree = result.unwrap(); + let expected = y + .iter() + .zip(sample_weights.iter()) + .map(|(yi, wi)| yi * wi) + .sum::() + / sample_weights.iter().sum::(); + + assert!((tree.predict_for_row(&x, 0) - expected).abs() < 1e-9); + } + + #[test] + fn uniform_weights_match_unweighted() { + // Test that using no weights is equivalent to using uniform weights + let x_rand: DenseMatrix = DenseMatrix::::rand(17, 5); + let y_rand: Vec = (0..17).collect::>().map(|y| *y as f64); + let parameters = BaseTreeRegressorParameters { + max_depth: Some(5), + min_samples_leaf: 1, + min_samples_split: 1, + seed: Some(42), + splitter: Splitter::Best, + }; + + let tree_no_weights = + BaseTreeRegressor::fit_inner(&x_rand, &y_rand, None, parameters.clone()) + .expect("Fit should work"); + let uniform_weights = vec![1.0f64; 17]; + let tree_with_weights = + BaseTreeRegressor::fit_inner(&x_rand, &y_rand, Some(&uniform_weights), parameters) + .expect("Fit should work"); + + let y_pred_no_weights = tree_no_weights + .predict(&x_rand) + .expect("Predict should work"); + let y_pred_with_weights = tree_with_weights + .predict(&x_rand) + .expect("Predict should work"); + assert!(mean_absolute_error(&y_pred_no_weights, &y_pred_with_weights) < 1e-9); + } + + #[test] + fn integer_weights_equivalent_to_repeating_sample() { + // Test that setting weight to "2" (or "3") is equivalent to having the same sample twice (or trice) + let x = DenseMatrix::from_2d_vec(&vec![vec![1.0_f64, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]]) + .unwrap(); + let y = vec![4.0, 5.0, 6.0]; + let sample_weights = vec![1.0, 2.0, 3.0]; + + let x_repeated = DenseMatrix::from_2d_vec(&vec![ + vec![1.0_f64, 2.0], + vec![3.0, 4.0], + vec![3.0, 4.0], + vec![5.0, 6.0], + vec![5.0, 6.0], + vec![5.0, 6.0], + ]) + .unwrap(); + let y_repeated = vec![4.0, 5.0, 5.0, 6.0, 6.0, 6.0]; + + let weighted_parameters = BaseTreeRegressorParameters { + max_depth: Some(1), // tree should not be able to fully separate all the examples + min_samples_leaf: 1, + min_samples_split: 1, + seed: Some(42), + splitter: Splitter::Best, + }; + + let repeated_parameters = BaseTreeRegressorParameters { + max_depth: Some(1), + min_samples_leaf: 1, + min_samples_split: 1, + seed: Some(42), + splitter: Splitter::Best, + }; + + let tree_weighted = + BaseTreeRegressor::fit_inner(&x, &y, Some(&sample_weights), weighted_parameters) + .expect("Fit should work"); + let tree_repeated = + BaseTreeRegressor::fit_inner(&x_repeated, &y_repeated, None, repeated_parameters) + .expect("Fit should work"); + // Predict on the same data + let y_pred_weighted = tree_weighted.predict(&x).expect("Predict should work"); + let y_pred_repeated = tree_repeated.predict(&x).expect("Predict should work"); + + assert!(mean_absolute_error(&y_pred_weighted, &y_pred_repeated) < 1e-9); + } + #[test] fn full_depth() { let x = DenseMatrix::from_2d_vec(&vec![ @@ -618,7 +798,7 @@ mod tests { splitter: Splitter::Best, }; - let tree = BaseTreeRegressor::fit(&x, &y, parameters).expect("Fit should work"); + let tree = BaseTreeRegressor::fit_inner(&x, &y, None, 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); @@ -626,6 +806,40 @@ mod tests { assert!(mean_absolute_error(&y_expected, &y_hat) < 1e-9); } + #[test] + fn full_depth_with_weights() { + 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 sample_weights = [1.0f64, 1.0, 1.0, 1.0, 10.0, 10.0]; + let tree = BaseTreeRegressor::fit_inner(&x, &y, Some(&sample_weights), parameters) + .expect("Fit should work"); + let x_test: DenseMatrix = + DenseMatrix::from_iterator((1..=12).map(|i| i as f64), 12, 1, 0); + let y_expected = vec![ + 1.50, 1.50, 6.50, 6.50, 11.0, 12.0, 12.0, 12.0, 12.0, 12.0, 12.0, 12.0, + ]; + let y_hat = tree.predict(&x_test).expect("Predict should work"); + assert_eq!(tree.nodes().len(), 7); + assert_eq!(tree.depth, 3); + assert!(mean_absolute_error(&y_expected, &y_hat) < 1e-9); + } + #[test] fn min_samples_split_boundary() { // A node that holds exactly `min_samples_split` samples must still split. @@ -640,7 +854,7 @@ mod tests { splitter: Splitter::Best, }; - let tree = BaseTreeRegressor::fit(&x, &y, parameters).expect("Fit should work"); + let tree = BaseTreeRegressor::fit_inner(&x, &y, None, parameters).expect("Fit should work"); assert_eq!(tree.nodes().len(), 3); assert_eq!(tree.depth, 2); @@ -653,8 +867,121 @@ mod tests { splitter: Splitter::Best, }; - let tree = BaseTreeRegressor::fit(&x, &y, parameters).expect("Fit should work"); + let tree = BaseTreeRegressor::fit_inner(&x, &y, None, parameters).expect("Fit should work"); assert_eq!(tree.nodes().len(), 1); assert_eq!(tree.depth, 0); } + + #[test] + fn zero_weight_on_true_side_gives_finite_predictions() { + // The first candidate split has only zero-weight samples on its true side, + // so true_mass is 0 in find_best_split. This must not produce NaN. + let x = DenseMatrix::from_2d_vec(&vec![vec![1.0_f64], vec![2.0], vec![3.0]]).unwrap(); + let y = vec![1.0f64, 2.0, 3.0]; + let sample_weights = [0.0f64, 1.0, 1.0]; + + let parameters = BaseTreeRegressorParameters { + max_depth: Some(2), + min_samples_leaf: 1, + min_samples_split: 2, + seed: None, + splitter: Splitter::Best, + }; + + let tree = BaseTreeRegressor::fit_inner(&x, &y, Some(&sample_weights), parameters) + .expect("Fit should work"); + + assert!(tree.nodes().iter().all(|node| node.output.is_finite())); + assert!( + tree.nodes() + .iter() + .all(|node| node.split_score.is_none_or(f64::is_finite)) + ); + + let y_hat = tree.predict(&x).expect("Predict should work"); + assert!(y_hat.iter().all(|v| v.is_finite())); + + // The best split separates x=3 from the other rows. + let y_expected = vec![2.0, 2.0, 3.0]; + assert!(mean_absolute_error(&y_expected, &y_hat) < 1e-9); + } + + #[test] + fn zero_weight_on_false_side_gives_finite_predictions() { + // The two first rows share x=1, so the only candidate split is between x=1 and x=2. + // The false side of that split holds only a zero-weight sample, + // so false_mass is 0 in find_best_split. This must not produce NaN. + let x = DenseMatrix::from_2d_vec(&vec![vec![1.0_f64], vec![1.0], vec![2.0]]).unwrap(); + let y = vec![1.0f64, 2.0, 3.0]; + let sample_weights = [1.0f64, 1.0, 0.0]; + + let parameters = BaseTreeRegressorParameters { + max_depth: Some(2), + min_samples_leaf: 1, + min_samples_split: 2, + seed: None, + splitter: Splitter::Best, + }; + + let tree = BaseTreeRegressor::fit_inner(&x, &y, Some(&sample_weights), parameters) + .expect("Fit should work"); + + assert!(tree.nodes().iter().all(|node| node.output.is_finite())); + assert!( + tree.nodes() + .iter() + .all(|node| node.split_score.is_none_or(f64::is_finite)) + ); + + let y_hat = tree.predict(&x).expect("Predict should work"); + assert!(y_hat.iter().all(|v| v.is_finite())); + } + + #[test] + fn weights_on_tied_feature_values() { + // Rows with the same feature value but different weights. The split must keep each + // tie group together, and each leaf must give the weighted mean of its group. + let x = DenseMatrix::from_2d_vec(&vec![vec![1.0_f64], vec![1.0], vec![2.0], vec![2.0]]) + .unwrap(); + let y = vec![0.0f64, 10.0, 20.0, 30.0]; + let sample_weights = [3.0f64, 1.0, 1.0, 3.0]; + + let parameters = BaseTreeRegressorParameters { + max_depth: None, + min_samples_leaf: 1, + min_samples_split: 2, + seed: None, + splitter: Splitter::Best, + }; + + let tree = BaseTreeRegressor::fit_inner(&x, &y, Some(&sample_weights), parameters.clone()) + .expect("Fit should work"); + + assert_eq!(tree.nodes().len(), 3); + assert_eq!(tree.depth, 2); + assert!((tree.nodes()[0].split_value.unwrap() - 1.5).abs() < 1e-9); // Split should be at 1.5 + + let y_hat = tree.predict(&x).expect("Predict should work"); + let y_expected = vec![2.5, 2.5, 27.5, 27.5]; // Expected values are weighted means + assert!(mean_absolute_error(&y_expected, &y_hat) < 1e-9); + + // Integer weights must give the same tree as repeated rows. + let x_repeated = DenseMatrix::from_2d_vec(&vec![ + vec![1.0_f64], + vec![1.0], + vec![1.0], + vec![1.0], + vec![2.0], + vec![2.0], + vec![2.0], + vec![2.0], + ]) + .unwrap(); + let y_repeated = vec![0.0f64, 0.0, 0.0, 10.0, 20.0, 30.0, 30.0, 30.0]; + let tree_repeated = + BaseTreeRegressor::fit_inner(&x_repeated, &y_repeated, None, parameters) + .expect("Fit should work"); + let y_hat_repeated = tree_repeated.predict(&x).expect("Predict should work"); + assert!(mean_absolute_error(&y_hat, &y_hat_repeated) < 1e-9); + } } diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index 2474155f..92c2480b 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -69,6 +69,7 @@ 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; #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)] @@ -300,6 +301,29 @@ impl, Y: Array1> x: &X, y: &Y, parameters: DecisionTreeRegressorParameters, + ) -> Result, Failed> { + Self::fit_inner(x, y, None, parameters) + } + + /// Build a decision tree regressor from the training data. + /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. + /// * `y` - the target values + /// * `sample_weights` - weights to use during fitting + pub fn fit_with_weights( + x: &X, + y: &Y, + sample_weights: &[f64], + parameters: DecisionTreeRegressorParameters, + ) -> Result, Failed> { + validate_sample_weights(sample_weights, x.shape().0)?; + Self::fit_inner(x, y, Some(sample_weights), parameters) + } + + fn fit_inner( + x: &X, + y: &Y, + sample_weights: Option<&[f64]>, + parameters: DecisionTreeRegressorParameters, ) -> Result, Failed> { let tree_parameters = BaseTreeRegressorParameters { max_depth: parameters.max_depth, @@ -308,7 +332,7 @@ impl, Y: Array1> seed: parameters.seed, splitter: Splitter::Best, }; - let tree = BaseTreeRegressor::fit(x, y, tree_parameters)?; + let tree = BaseTreeRegressor::fit_inner(x, y, sample_weights, tree_parameters)?; Ok(Self { tree_regressor: Some(tree), }) @@ -349,6 +373,40 @@ mod tests { assert!(iter.next().is_none()); } + #[test] + fn fit_with_weights_validates_weights() { + let x: DenseMatrix = DenseMatrix::from_iterator((0..6).map(|i| i as f64), 3, 2, 0); + let y = vec![1.0_f64, 2.0, 3.0]; + let parameters = DecisionTreeRegressorParameters::default(); + + // Valid weights: a zero weight is permitted when the sum is positive + for weights in [vec![1.0, 2.0, 3.0], vec![0.0, 0.0, 0.5]] { + assert!( + DecisionTreeRegressor::fit_with_weights(&x, &y, &weights, parameters.clone()) + .is_ok(), + "weights: {weights:?}" + ); + } + + let wrong_length = "Number of sample weights must equal number of rows in x"; + let not_finite_or_negative = "Sample weights must be finite and non-negative"; + let zero_sum = "Sum of sample weights must be positive"; + let cases: Vec<(Vec, &str)> = vec![ + (vec![], wrong_length), + (vec![1.0, 2.0], wrong_length), + (vec![1.0, 2.0, 3.0, 4.0], wrong_length), + (vec![1.0, -1.0, 3.0], not_finite_or_negative), + (vec![1.0, f64::NAN, 3.0], not_finite_or_negative), + (vec![1.0, f64::INFINITY, 3.0], not_finite_or_negative), + (vec![0.0, 0.0, 0.0], zero_sum), + ]; + for (weights, msg) in cases { + let result = + DecisionTreeRegressor::fit_with_weights(&x, &y, &weights, parameters.clone()); + assert_eq!(result.err(), Some(Failed::fit(msg)), "weights: {weights:?}"); + } + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test @@ -430,6 +488,110 @@ mod tests { } } + #[test] + fn fit_with_weights_matches_sklearn() { + // Reference: sklearn DecisionTreeRegressor(max_depth=3, random_state=0) + // fitted with the same sample_weight. It gives this tree: + // + // |--- x1 <= 2.50 + // | |--- x2 <= 0.50 + // | | |--- value: [2.00] + // | |--- x2 > 0.50 + // | | |--- x1 <= 1.50 + // | | | |--- value: [3.00] + // | | |--- x1 > 1.50 + // | | | |--- value: [2.50] + // |--- x1 > 2.50 + // | |--- x1 <= 6.50 + // | | |--- x2 <= 0.50 + // | | | |--- value: [6.00] + // | | |--- x2 > 0.50 + // | | | |--- value: [7.50] + // | |--- x1 > 6.50 + // | | |--- x2 <= 0.50 + // | | | |--- value: [9.00] + // | | |--- x2 > 0.50 + // | | | |--- value: [9.80] + let x = DenseMatrix::from_2d_array(&[ + &[1., 0.], + &[1., 1.], + &[2., 1.], + &[3., 0.], + &[3., 0.], + &[3., 1.], + &[5., 0.], + &[5., 1.], + &[6., 1.], + &[7., 0.], + &[8., 1.], + &[8., 1.], + ]) + .unwrap(); + let y: Vec = vec![2.0, 3.0, 2.5, 6.0, 5.0, 7.0, 6.5, 8.0, 7.5, 9.0, 10.0, 9.5]; + let sample_weights = [2.0, 1.0, 3.0, 1.0, 2.0, 1.0, 4.0, 1.0, 2.0, 1.0, 3.0, 2.0]; + + // smartcore counts the root as level 1, so sklearn max_depth=3 is max_depth=4 here. + let parameters = DecisionTreeRegressorParameters::default().with_max_depth(4); + let tree = DecisionTreeRegressor::fit_with_weights(&x, &y, &sample_weights, parameters) + .expect("Fit should work"); + + // Each training row gets the weighted mean of its leaf. + let y_hat = tree.predict(&x).unwrap(); + let y_expected = [2.0, 3.0, 2.5, 6.0, 6.0, 7.5, 6.0, 7.5, 7.5, 9.0, 9.8, 9.8]; + for i in 0..y_expected.len() { + assert!( + (y_hat[i] - y_expected[i]).abs() < 1e-9, + "row {i}: got {}, expected {}", + y_hat[i], + y_expected[i] + ); + } + + // Probe points on each side of each threshold check the split features and the + // midpoint thresholds. At the node x1 > 6.5, the splits x2 <= 0.5 and x1 <= 7.5 give + // the same partition and the same gain. sklearn picks x2 because of its random + // feature order. Thus we probe only points where both splits give the same value. + let probes = DenseMatrix::from_2d_array(&[ + // root: x1 <= 2.5 + &[2.4, 0.], + &[2.6, 0.], + &[2.4, 1.], + &[2.6, 1.], + // left: x2 <= 0.5 + &[1.0, 0.4], + &[1.0, 0.6], + // left, x2 > 0.5: x1 <= 1.5 + &[1.4, 1.], + &[1.6, 1.], + // right: x1 <= 6.5 + &[6.4, 0.], + &[6.6, 0.], + &[6.4, 1.], + // right, x1 <= 6.5: x2 <= 0.5 + &[4.0, 0.4], + &[4.0, 0.6], + // right, x1 > 6.5: tied split + &[7.0, 0.4], + &[8.0, 0.6], + // out of the training range + &[0., 0.], + &[10., 1.], + ]) + .unwrap(); + let probes_hat = tree.predict(&probes).unwrap(); + let probes_expected = [ + 2.0, 6.0, 2.5, 7.5, 2.0, 3.0, 3.0, 2.5, 6.0, 9.0, 7.5, 6.0, 7.5, 9.0, 9.8, 2.0, 9.8, + ]; + for i in 0..probes_expected.len() { + assert!( + (probes_hat[i] - probes_expected[i]).abs() < 1e-9, + "probe {i}: got {}, expected {}", + probes_hat[i], + probes_expected[i] + ); + } + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test