diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index 223a8b90..198189a6 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -94,12 +94,12 @@ impl, Y: Array1 .unwrap_or((num_attributes as f64).sqrt().floor() as usize); let mut rng = get_rng_impl(Some(parameters.seed)); - let mut trees: Vec> = Vec::new(); + let n_trees = parameters.n_trees; + let mut trees: Vec> = Vec::with_capacity(n_trees); let mut maybe_all_samples: Option>> = Option::None; if parameters.keep_samples { - // TODO: use with_capacity here - maybe_all_samples = Some(Vec::new()); + maybe_all_samples = Some(Vec::with_capacity(n_trees)); } let mut samples: Vec = (0..n_rows).map(|_| 1).collect(); @@ -218,3 +218,29 @@ impl, Y: Array1 samples } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::linalg::basic::matrix::DenseMatrix; + + #[test] + fn test_base_forest_regressor_keep_samples() { + let x = DenseMatrix::from_2d_array(&[&[1.0, 2.0], &[3.0, 4.0], &[5.0, 6.0]]).unwrap(); + let y = vec![1.0, 2.0, 3.0]; + let params = BaseForestRegressorParameters { + max_depth: None, + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 5, + m: None, + keep_samples: true, + seed: 42, + bootstrap: true, + splitter: crate::tree::base_tree_regressor::Splitter::Best, + }; + let regressor = BaseForestRegressor::fit(&x, &y, params).unwrap(); + assert_eq!(regressor.trees.unwrap().len(), 5); + assert!(regressor.samples.is_some()); + } +} diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index 0f86a4df..8553472a 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -475,13 +475,12 @@ impl, Y: Array1> = Vec::new(); + let n_trees = parameters.n_trees as usize; + let mut trees: Vec> = Vec::with_capacity(n_trees); let mut maybe_all_samples: Option>> = Option::None; if parameters.keep_samples { - // TODO: use with_capacity here - maybe_all_samples = Some(Vec::new()); + maybe_all_samples = Some(Vec::with_capacity(n_trees)); } for _ in 0..parameters.n_trees {