-
-
Notifications
You must be signed in to change notification settings - Fork 109
feat(tree,ensemble): add sample weights to regressor fit #467
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -79,9 +79,11 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, 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<BaseForestRegressor<TX, TY, X, Y>, Failed> { | ||
| let (n_rows, num_attributes) = x.shape(); | ||
|
|
@@ -111,10 +113,20 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1 | |
|
|
||
| let mut samples: Vec<usize> = (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::<TX, TY, X, Y>::sample_with_replacement(n_rows, &mut rng); | ||
| samples = BaseForestRegressor::<TX, TY, X, Y>::sample_with_replacement( | ||
| n_rows, | ||
| &mut rng, | ||
| dist.as_ref(), | ||
| ); | ||
| } | ||
|
|
||
| // keep samples is flag is on | ||
|
|
@@ -129,7 +141,14 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, 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<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1 | |
| result / TY::from(n_trees).unwrap() | ||
| } | ||
|
|
||
| fn sample_with_replacement(nrows: usize, rng: &mut impl rand::Rng) -> Vec<usize> { | ||
| fn sample_with_replacement( | ||
| nrows: usize, | ||
| rng: &mut impl rand::Rng, | ||
| distribution: Option<&rand::distr::weighted::WeightedIndex<f64>>, | ||
| ) -> Vec<usize> { | ||
| 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; | ||
| } | ||
|
|
||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The two branches of for _ in 0..nrows {
let xi = if let Some(dist) = distribution {
rng.sample(dist)
} else {
rng.random_range(0..nrows)
};
samples[xi] += 1;
}Or, construct a closure/enum once before the loop to avoid the per-iteration branch if performance matters. |
||
| 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<f64> = DenseMatrix::rand(1000, 10); | ||
| let model_parameters = (0..=9).map(|x| x as f64).collect::<Vec<_>>(); | ||
| let y: Vec<f64> = 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::<f64>() - y.iter().sum::<f64>()).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<f64> = (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::<Vec<_>>(); | ||
|
|
||
| let expected = y | ||
| .iter() | ||
| .zip(normalized_weights.iter()) | ||
| .map(|(&yi, &w)| yi * w) | ||
| .sum::<f64>(); | ||
| let actual = y_hat | ||
| .iter() | ||
| .zip(normalized_weights.iter()) | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The |
||
| .map(|(&yi, &w)| yi * w) | ||
| .sum::<f64>(); | ||
|
|
||
| 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<f64> = | ||
| DenseMatrix::from_iterator((0..2 * n_rows).map(|k| k as f64), n_rows, 2, 0); | ||
| let y: Vec<f64> = (0..n_rows).map(|i| i as f64).collect(); | ||
| let sample_weights: Vec<f64> = (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:?})" | ||
| ); | ||
| } | ||
| } | ||
| } | ||
| } | ||
|
|
||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Nit: the comment in the test says |
||
| #[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<f64> = DenseMatrix::from_iterator((0..20).map(|i| i as f64), 20, 1, 0); | ||
| let y: Vec<f64> = (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<f64> = (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}" | ||
| ); | ||
| } | ||
| } | ||
| } | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The public signature of
BaseForestRegressor::fitchanged (newsample_weightsparameter). The doc comment was removed in this hunk. Please restore it, and document the newsample_weightsparameter.