From d6371f830c909f366d91d8741fd21e3a78109987 Mon Sep 17 00:00:00 2001 From: Montana Low Date: Wed, 14 Sep 2022 18:03:41 -0700 Subject: [PATCH 1/6] grid search draft --- src/model_selection/mod.rs | 113 ++++++++++++++++++++++++++++++++++++- 1 file changed, 112 insertions(+), 1 deletion(-) diff --git a/src/model_selection/mod.rs b/src/model_selection/mod.rs index 68f06350..57b87018 100644 --- a/src/model_selection/mod.rs +++ b/src/model_selection/mod.rs @@ -103,7 +103,7 @@ //! but instead of test error it calculates predictions for all samples in the test set. use crate::api::Predictor; -use crate::error::Failed; +use crate::error::{Failed, FailedError}; use crate::linalg::BaseVector; use crate::linalg::Matrix; use crate::math::num::RealNumber; @@ -276,6 +276,66 @@ where Ok(y_hat) } +/// grid search results. +#[derive(Clone, Debug)] +pub struct GridSearchResult { + /// Vector with test scores on each cv split + pub cross_validation_result: CrossValidationResult, + /// Vector with training scores on each cv split + pub parameters: I, +} + +/// Search for the best estimator by testing all possible combinations with cross-validation using given metric. +/// * `fit_estimator` - a `fit` function of an estimator +/// * `x` - features, matrix of size _NxM_ where _N_ is number of samples and _M_ is number of attributes. +/// * `y` - target values, should be of size _N_ +/// * `parameter_search` - an iterator for parameters that will be tested. +/// * `cv` - the cross-validation splitting strategy, should be an instance of [`BaseKFold`](./trait.BaseKFold.html) +/// * `score` - a metric to use for evaluation, see [metrics](../metrics/index.html) +pub fn grid_search( + fit_estimator: F, + x: &M, + y: &M::RowVector, + parameter_search: I, + cv: K, + score: S, +) -> Result, Failed> +where + T: RealNumber, + M: Matrix, + I: Iterator, + I::Item: Clone, + E: Predictor, + K: BaseKFold, + F: Fn(&M, &M::RowVector, I::Item) -> Result, + S: Fn(&M::RowVector, &M::RowVector) -> T, +{ + let mut best_result: Option> = None; + let mut best_parameters = None; + + for parameters in parameter_search { + let result = cross_validate(&fit_estimator, x, y, ¶meters, &cv, &score)?; + if best_result.is_none() + || result.mean_test_score() > best_result.as_ref().unwrap().mean_test_score() + { + best_parameters = Some(parameters); + best_result = Some(result); + } + } + + if let (Some(bp), Some(br)) = (best_parameters, best_result) { + Ok(GridSearchResult { + cross_validation_result: br, + parameters: bp, + }) + } else { + Err(Failed::because( + FailedError::FindFailed, + "there were no parameter sets found", + )) + } +} + #[cfg(test)] mod tests { @@ -307,6 +367,57 @@ mod tests { assert_eq!(x_test.shape().0, y_test.len()); } + #[test] + fn test_grid_search() { + let x = DenseMatrix::from_2d_array(&[ + &[5.1, 3.5, 1.4, 0.2], + &[4.9, 3.0, 1.4, 0.2], + &[4.7, 3.2, 1.3, 0.2], + &[4.6, 3.1, 1.5, 0.2], + &[5.0, 3.6, 1.4, 0.2], + &[5.4, 3.9, 1.7, 0.4], + &[4.6, 3.4, 1.4, 0.3], + &[5.0, 3.4, 1.5, 0.2], + &[4.4, 2.9, 1.4, 0.2], + &[4.9, 3.1, 1.5, 0.1], + &[7.0, 3.2, 4.7, 1.4], + &[6.4, 3.2, 4.5, 1.5], + &[6.9, 3.1, 4.9, 1.5], + &[5.5, 2.3, 4.0, 1.3], + &[6.5, 2.8, 4.6, 1.5], + &[5.7, 2.8, 4.5, 1.3], + &[6.3, 3.3, 4.7, 1.6], + &[4.9, 2.4, 3.3, 1.0], + &[6.6, 2.9, 4.6, 1.3], + &[5.2, 2.7, 3.9, 1.4], + ]); + let y = vec![ + 0., 0., 0., 0., 0., 0., 0., 0., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., + ]; + + let cv = KFold { + n_splits: 5, + ..KFold::default() + }; + + let parameters = LogisticRegressionSearchParameters { + alpha: vec![0., 1.], + ..Default::default() + }; + + let results = grid_search( + LogisticRegression::fit, + &x, + &y, + parameters.into_iter(), + cv, + &accuracy, + ) + .unwrap(); + + assert!([0., 1.].contains(&results.parameters.alpha)); + } + #[derive(Clone)] struct NoParameters {} From 48d3e2c9438cb97368db4140f8ed695211681527 Mon Sep 17 00:00:00 2001 From: Montana Low Date: Thu, 15 Sep 2022 08:26:39 -0700 Subject: [PATCH 2/6] hyperparam search for linear estimators --- src/model_selection/mod.rs | 113 +------------------------------------ 1 file changed, 1 insertion(+), 112 deletions(-) diff --git a/src/model_selection/mod.rs b/src/model_selection/mod.rs index 57b87018..68f06350 100644 --- a/src/model_selection/mod.rs +++ b/src/model_selection/mod.rs @@ -103,7 +103,7 @@ //! but instead of test error it calculates predictions for all samples in the test set. use crate::api::Predictor; -use crate::error::{Failed, FailedError}; +use crate::error::Failed; use crate::linalg::BaseVector; use crate::linalg::Matrix; use crate::math::num::RealNumber; @@ -276,66 +276,6 @@ where Ok(y_hat) } -/// grid search results. -#[derive(Clone, Debug)] -pub struct GridSearchResult { - /// Vector with test scores on each cv split - pub cross_validation_result: CrossValidationResult, - /// Vector with training scores on each cv split - pub parameters: I, -} - -/// Search for the best estimator by testing all possible combinations with cross-validation using given metric. -/// * `fit_estimator` - a `fit` function of an estimator -/// * `x` - features, matrix of size _NxM_ where _N_ is number of samples and _M_ is number of attributes. -/// * `y` - target values, should be of size _N_ -/// * `parameter_search` - an iterator for parameters that will be tested. -/// * `cv` - the cross-validation splitting strategy, should be an instance of [`BaseKFold`](./trait.BaseKFold.html) -/// * `score` - a metric to use for evaluation, see [metrics](../metrics/index.html) -pub fn grid_search( - fit_estimator: F, - x: &M, - y: &M::RowVector, - parameter_search: I, - cv: K, - score: S, -) -> Result, Failed> -where - T: RealNumber, - M: Matrix, - I: Iterator, - I::Item: Clone, - E: Predictor, - K: BaseKFold, - F: Fn(&M, &M::RowVector, I::Item) -> Result, - S: Fn(&M::RowVector, &M::RowVector) -> T, -{ - let mut best_result: Option> = None; - let mut best_parameters = None; - - for parameters in parameter_search { - let result = cross_validate(&fit_estimator, x, y, ¶meters, &cv, &score)?; - if best_result.is_none() - || result.mean_test_score() > best_result.as_ref().unwrap().mean_test_score() - { - best_parameters = Some(parameters); - best_result = Some(result); - } - } - - if let (Some(bp), Some(br)) = (best_parameters, best_result) { - Ok(GridSearchResult { - cross_validation_result: br, - parameters: bp, - }) - } else { - Err(Failed::because( - FailedError::FindFailed, - "there were no parameter sets found", - )) - } -} - #[cfg(test)] mod tests { @@ -367,57 +307,6 @@ mod tests { assert_eq!(x_test.shape().0, y_test.len()); } - #[test] - fn test_grid_search() { - let x = DenseMatrix::from_2d_array(&[ - &[5.1, 3.5, 1.4, 0.2], - &[4.9, 3.0, 1.4, 0.2], - &[4.7, 3.2, 1.3, 0.2], - &[4.6, 3.1, 1.5, 0.2], - &[5.0, 3.6, 1.4, 0.2], - &[5.4, 3.9, 1.7, 0.4], - &[4.6, 3.4, 1.4, 0.3], - &[5.0, 3.4, 1.5, 0.2], - &[4.4, 2.9, 1.4, 0.2], - &[4.9, 3.1, 1.5, 0.1], - &[7.0, 3.2, 4.7, 1.4], - &[6.4, 3.2, 4.5, 1.5], - &[6.9, 3.1, 4.9, 1.5], - &[5.5, 2.3, 4.0, 1.3], - &[6.5, 2.8, 4.6, 1.5], - &[5.7, 2.8, 4.5, 1.3], - &[6.3, 3.3, 4.7, 1.6], - &[4.9, 2.4, 3.3, 1.0], - &[6.6, 2.9, 4.6, 1.3], - &[5.2, 2.7, 3.9, 1.4], - ]); - let y = vec![ - 0., 0., 0., 0., 0., 0., 0., 0., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., - ]; - - let cv = KFold { - n_splits: 5, - ..KFold::default() - }; - - let parameters = LogisticRegressionSearchParameters { - alpha: vec![0., 1.], - ..Default::default() - }; - - let results = grid_search( - LogisticRegression::fit, - &x, - &y, - parameters.into_iter(), - cv, - &accuracy, - ) - .unwrap(); - - assert!([0., 1.].contains(&results.parameters.alpha)); - } - #[derive(Clone)] struct NoParameters {} From 357a92f19d82a66a130d2a8e662bc439c328905a Mon Sep 17 00:00:00 2001 From: Montana Low Date: Thu, 15 Sep 2022 08:56:53 -0700 Subject: [PATCH 3/6] grid search for ensembles --- src/ensemble/random_forest_classifier.rs | 244 ++++++++++++++++++++++- src/ensemble/random_forest_regressor.rs | 208 +++++++++++++++++++ src/linear/lasso.rs | 31 ++- src/tree/decision_tree_classifier.rs | 2 +- 4 files changed, 466 insertions(+), 19 deletions(-) diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index 247b5025..53c48451 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -193,6 +193,225 @@ impl> Predictor for RandomForestCla } } +/// RandomForestClassifier grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct RandomForestClassifierSearchParameters { + /// Split criteria to use when building a tree. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub criterion: Vec, + /// Tree max depth. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub max_depth: Vec>, + /// The minimum number of samples required to be at a leaf node. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub min_samples_leaf: Vec, + /// The minimum number of samples required to split an internal node. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub min_samples_split: Vec, + /// The number of trees in the forest. + pub n_trees: Vec, + /// Number of random sample of predictors to use as split candidates. + pub m: Vec>, + /// Whether to keep samples used for tree generation. This is required for OOB prediction. + pub keep_samples: Vec, + /// Seed used for bootstrap sampling and feature selection for each tree. + pub seed: Vec, +} + +/// RandomForestClassifier grid search iterator +pub struct RandomForestClassifierSearchParametersIterator { + random_forest_classifier_search_parameters: RandomForestClassifierSearchParameters, + current_criterion: usize, + current_max_depth: usize, + current_min_samples_leaf: usize, + current_min_samples_split: usize, + current_n_trees: usize, + current_m: usize, + current_keep_samples: usize, + current_seed: usize, +} + +impl IntoIterator for RandomForestClassifierSearchParameters { + type Item = RandomForestClassifierParameters; + type IntoIter = RandomForestClassifierSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + RandomForestClassifierSearchParametersIterator { + random_forest_classifier_search_parameters: self, + current_criterion: 0, + current_max_depth: 0, + current_min_samples_leaf: 0, + current_min_samples_split: 0, + current_n_trees: 0, + current_m: 0, + current_keep_samples: 0, + current_seed: 0, + } + } +} + +impl Iterator for RandomForestClassifierSearchParametersIterator { + type Item = RandomForestClassifierParameters; + + fn next(&mut self) -> Option { + if self.current_criterion + == self + .random_forest_classifier_search_parameters + .criterion + .len() + && self.current_max_depth + == self + .random_forest_classifier_search_parameters + .max_depth + .len() + && self.current_min_samples_leaf + == self + .random_forest_classifier_search_parameters + .min_samples_leaf + .len() + && self.current_min_samples_split + == self + .random_forest_classifier_search_parameters + .min_samples_split + .len() + && self.current_n_trees + == self + .random_forest_classifier_search_parameters + .n_trees + .len() + && self.current_m == self.random_forest_classifier_search_parameters.m.len() + && self.current_keep_samples + == self + .random_forest_classifier_search_parameters + .keep_samples + .len() + && self.current_seed == self.random_forest_classifier_search_parameters.seed.len() + { + return None; + } + + let next = RandomForestClassifierParameters { + criterion: self.random_forest_classifier_search_parameters.criterion + [self.current_criterion], + max_depth: self.random_forest_classifier_search_parameters.max_depth + [self.current_max_depth], + min_samples_leaf: self + .random_forest_classifier_search_parameters + .min_samples_leaf[self.current_min_samples_leaf], + min_samples_split: self + .random_forest_classifier_search_parameters + .min_samples_split[self.current_min_samples_split], + n_trees: self.random_forest_classifier_search_parameters.n_trees[self.current_n_trees], + m: self.random_forest_classifier_search_parameters.m[self.current_m], + keep_samples: self.random_forest_classifier_search_parameters.keep_samples + [self.current_keep_samples], + seed: self.random_forest_classifier_search_parameters.seed[self.current_seed], + }; + + if self.current_criterion + 1 + < self + .random_forest_classifier_search_parameters + .criterion + .len() + { + self.current_criterion += 1; + } else if self.current_max_depth + 1 + < self + .random_forest_classifier_search_parameters + .max_depth + .len() + { + self.current_criterion = 0; + self.current_max_depth += 1; + } else if self.current_min_samples_leaf + 1 + < self + .random_forest_classifier_search_parameters + .min_samples_leaf + .len() + { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf += 1; + } else if self.current_min_samples_split + 1 + < self + .random_forest_classifier_search_parameters + .min_samples_split + .len() + { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split += 1; + } else if self.current_n_trees + 1 + < self + .random_forest_classifier_search_parameters + .n_trees + .len() + { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees += 1; + } else if self.current_m + 1 < self.random_forest_classifier_search_parameters.m.len() { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees = 0; + self.current_m += 1; + } else if self.current_keep_samples + 1 + < self + .random_forest_classifier_search_parameters + .keep_samples + .len() + { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees = 0; + self.current_m = 0; + self.current_keep_samples += 1; + } else if self.current_seed + 1 < self.random_forest_classifier_search_parameters.seed.len() + { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees = 0; + self.current_m = 0; + self.current_keep_samples = 0; + self.current_seed += 1; + } else { + self.current_criterion += 1; + self.current_max_depth += 1; + self.current_min_samples_leaf += 1; + self.current_min_samples_split += 1; + self.current_n_trees += 1; + self.current_m += 1; + self.current_keep_samples += 1; + self.current_seed += 1; + } + + Some(next) + } +} + +impl Default for RandomForestClassifierSearchParameters { + fn default() -> Self { + let default_params = RandomForestClassifierParameters::default(); + + RandomForestClassifierSearchParameters { + criterion: vec![default_params.criterion], + max_depth: vec![default_params.max_depth], + min_samples_leaf: vec![default_params.min_samples_leaf], + min_samples_split: vec![default_params.min_samples_split], + n_trees: vec![default_params.n_trees], + m: vec![default_params.m], + keep_samples: vec![default_params.keep_samples], + seed: vec![default_params.seed], + } + } +} + impl RandomForestClassifier { /// Build a forest of trees from the training set. /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. @@ -238,7 +457,7 @@ impl RandomForestClassifier { } let params = DecisionTreeClassifierParameters { - criterion: parameters.criterion.clone(), + criterion: parameters.criterion, max_depth: parameters.max_depth, min_samples_leaf: parameters.min_samples_leaf, min_samples_split: parameters.min_samples_split, @@ -346,6 +565,29 @@ mod tests { use crate::linalg::naive::dense_matrix::DenseMatrix; use crate::metrics::*; + #[test] + fn search_parameters() { + let parameters = RandomForestClassifierSearchParameters { + n_trees: vec![10, 100], + m: vec![None, Some(1)], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 10); + assert_eq!(next.m, None); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 100); + assert_eq!(next.m, None); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 10); + assert_eq!(next.m, Some(1)); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 100); + assert_eq!(next.m, Some(1)); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn fit_predict_iris() { diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 08a7dcc7..ec781375 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -176,6 +176,191 @@ impl> Predictor for RandomForestReg } } +/// RandomForestRegressor grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct RandomForestRegressorSearchParameters { + /// Tree max depth. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub max_depth: Vec>, + /// The minimum number of samples required to be at a leaf node. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub min_samples_leaf: Vec, + /// The minimum number of samples required to split an internal node. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub min_samples_split: Vec, + /// The number of trees in the forest. + pub n_trees: Vec, + /// Number of random sample of predictors to use as split candidates. + pub m: Vec>, + /// Whether to keep samples used for tree generation. This is required for OOB prediction. + pub keep_samples: Vec, + /// Seed used for bootstrap sampling and feature selection for each tree. + pub seed: Vec, +} + +/// RandomForestRegressor grid search iterator +pub struct RandomForestRegressorSearchParametersIterator { + random_forest_regressor_search_parameters: RandomForestRegressorSearchParameters, + current_max_depth: usize, + current_min_samples_leaf: usize, + current_min_samples_split: usize, + current_n_trees: usize, + current_m: usize, + current_keep_samples: usize, + current_seed: usize, +} + +impl IntoIterator for RandomForestRegressorSearchParameters { + type Item = RandomForestRegressorParameters; + type IntoIter = RandomForestRegressorSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + RandomForestRegressorSearchParametersIterator { + random_forest_regressor_search_parameters: self, + current_max_depth: 0, + current_min_samples_leaf: 0, + current_min_samples_split: 0, + current_n_trees: 0, + current_m: 0, + current_keep_samples: 0, + current_seed: 0, + } + } +} + +impl Iterator for RandomForestRegressorSearchParametersIterator { + type Item = RandomForestRegressorParameters; + + fn next(&mut self) -> Option { + if self.current_max_depth + == self + .random_forest_regressor_search_parameters + .max_depth + .len() + && self.current_min_samples_leaf + == self + .random_forest_regressor_search_parameters + .min_samples_leaf + .len() + && self.current_min_samples_split + == self + .random_forest_regressor_search_parameters + .min_samples_split + .len() + && self.current_n_trees == self.random_forest_regressor_search_parameters.n_trees.len() + && self.current_m == self.random_forest_regressor_search_parameters.m.len() + && self.current_keep_samples + == self + .random_forest_regressor_search_parameters + .keep_samples + .len() + && self.current_seed == self.random_forest_regressor_search_parameters.seed.len() + { + return None; + } + + let next = RandomForestRegressorParameters { + max_depth: self.random_forest_regressor_search_parameters.max_depth + [self.current_max_depth], + min_samples_leaf: self + .random_forest_regressor_search_parameters + .min_samples_leaf[self.current_min_samples_leaf], + min_samples_split: self + .random_forest_regressor_search_parameters + .min_samples_split[self.current_min_samples_split], + n_trees: self.random_forest_regressor_search_parameters.n_trees[self.current_n_trees], + m: self.random_forest_regressor_search_parameters.m[self.current_m], + keep_samples: self.random_forest_regressor_search_parameters.keep_samples + [self.current_keep_samples], + seed: self.random_forest_regressor_search_parameters.seed[self.current_seed], + }; + + if self.current_max_depth + 1 + < self + .random_forest_regressor_search_parameters + .max_depth + .len() + { + self.current_max_depth += 1; + } else if self.current_min_samples_leaf + 1 + < self + .random_forest_regressor_search_parameters + .min_samples_leaf + .len() + { + self.current_max_depth = 0; + self.current_min_samples_leaf += 1; + } else if self.current_min_samples_split + 1 + < self + .random_forest_regressor_search_parameters + .min_samples_split + .len() + { + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split += 1; + } else if self.current_n_trees + 1 + < self.random_forest_regressor_search_parameters.n_trees.len() + { + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees += 1; + } else if self.current_m + 1 < self.random_forest_regressor_search_parameters.m.len() { + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees = 0; + self.current_m += 1; + } else if self.current_keep_samples + 1 + < self + .random_forest_regressor_search_parameters + .keep_samples + .len() + { + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees = 0; + self.current_m = 0; + self.current_keep_samples += 1; + } else if self.current_seed + 1 < self.random_forest_regressor_search_parameters.seed.len() + { + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split = 0; + self.current_n_trees = 0; + self.current_m = 0; + self.current_keep_samples = 0; + self.current_seed += 1; + } else { + self.current_max_depth += 1; + self.current_min_samples_leaf += 1; + self.current_min_samples_split += 1; + self.current_n_trees += 1; + self.current_m += 1; + self.current_keep_samples += 1; + self.current_seed += 1; + } + + Some(next) + } +} + +impl Default for RandomForestRegressorSearchParameters { + fn default() -> Self { + let default_params = RandomForestRegressorParameters::default(); + + RandomForestRegressorSearchParameters { + max_depth: vec![default_params.max_depth], + min_samples_leaf: vec![default_params.min_samples_leaf], + min_samples_split: vec![default_params.min_samples_split], + n_trees: vec![default_params.n_trees], + m: vec![default_params.m], + keep_samples: vec![default_params.keep_samples], + seed: vec![default_params.seed], + } + } +} + impl RandomForestRegressor { /// Build a forest of trees from the training set. /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. @@ -302,6 +487,29 @@ mod tests { use crate::linalg::naive::dense_matrix::DenseMatrix; use crate::metrics::mean_absolute_error; + #[test] + fn search_parameters() { + let parameters = RandomForestRegressorSearchParameters { + n_trees: vec![10, 100], + m: vec![None, Some(1)], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 10); + assert_eq!(next.m, None); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 100); + assert_eq!(next.m, None); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 10); + assert_eq!(next.m, Some(1)); + let next = iter.next().unwrap(); + assert_eq!(next.n_trees, 100); + assert_eq!(next.m, Some(1)); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn fit_longley() { diff --git a/src/linear/lasso.rs b/src/linear/lasso.rs index 7e80a8bb..aae7e500 100644 --- a/src/linear/lasso.rs +++ b/src/linear/lasso.rs @@ -129,7 +129,7 @@ pub struct LassoSearchParameters { /// Lasso grid search iterator pub struct LassoSearchParametersIterator { - lasso_regression_search_parameters: LassoSearchParameters, + lasso_search_parameters: LassoSearchParameters, current_alpha: usize, current_normalize: usize, current_tol: usize, @@ -142,7 +142,7 @@ impl IntoIterator for LassoSearchParameters { fn into_iter(self) -> Self::IntoIter { LassoSearchParametersIterator { - lasso_regression_search_parameters: self, + lasso_search_parameters: self, current_alpha: 0, current_normalize: 0, current_tol: 0, @@ -155,34 +155,31 @@ impl Iterator for LassoSearchParametersIterator { type Item = LassoParameters; fn next(&mut self) -> Option { - if self.current_alpha == self.lasso_regression_search_parameters.alpha.len() - && self.current_normalize == self.lasso_regression_search_parameters.normalize.len() - && self.current_tol == self.lasso_regression_search_parameters.tol.len() - && self.current_max_iter == self.lasso_regression_search_parameters.max_iter.len() + if self.current_alpha == self.lasso_search_parameters.alpha.len() + && self.current_normalize == self.lasso_search_parameters.normalize.len() + && self.current_tol == self.lasso_search_parameters.tol.len() + && self.current_max_iter == self.lasso_search_parameters.max_iter.len() { return None; } let next = LassoParameters { - alpha: self.lasso_regression_search_parameters.alpha[self.current_alpha], - normalize: self.lasso_regression_search_parameters.normalize[self.current_normalize], - tol: self.lasso_regression_search_parameters.tol[self.current_tol], - max_iter: self.lasso_regression_search_parameters.max_iter[self.current_max_iter], + alpha: self.lasso_search_parameters.alpha[self.current_alpha], + normalize: self.lasso_search_parameters.normalize[self.current_normalize], + tol: self.lasso_search_parameters.tol[self.current_tol], + max_iter: self.lasso_search_parameters.max_iter[self.current_max_iter], }; - if self.current_alpha + 1 < self.lasso_regression_search_parameters.alpha.len() { + if self.current_alpha + 1 < self.lasso_search_parameters.alpha.len() { self.current_alpha += 1; - } else if self.current_normalize + 1 - < self.lasso_regression_search_parameters.normalize.len() - { + } else if self.current_normalize + 1 < self.lasso_search_parameters.normalize.len() { self.current_alpha = 0; self.current_normalize += 1; - } else if self.current_tol + 1 < self.lasso_regression_search_parameters.tol.len() { + } else if self.current_tol + 1 < self.lasso_search_parameters.tol.len() { self.current_alpha = 0; self.current_normalize = 0; self.current_tol += 1; - } else if self.current_max_iter + 1 < self.lasso_regression_search_parameters.max_iter.len() - { + } else if self.current_max_iter + 1 < self.lasso_search_parameters.max_iter.len() { self.current_alpha = 0; self.current_normalize = 0; self.current_tol = 0; diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index 35889e4e..d58125d4 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -105,7 +105,7 @@ pub struct DecisionTreeClassifier { /// The function to measure the quality of a split. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Copy)] pub enum SplitCriterion { /// [Gini index](../decision_tree_classifier/index.html) Gini, From aec315f889735024e6486b1a32d921dfae8bda0e Mon Sep 17 00:00:00 2001 From: Montana Low Date: Thu, 15 Sep 2022 10:55:32 -0700 Subject: [PATCH 4/6] support grid search for more algos --- src/naive_bayes/bernoulli.rs | 96 ++++++++++++++++ src/naive_bayes/categorical.rs | 68 ++++++++++++ src/naive_bayes/gaussian.rs | 76 ++++++++++++- src/naive_bayes/multinomial.rs | 84 ++++++++++++++ src/svm/mod.rs | 10 +- src/svm/svc.rs | 125 ++++++++++++++++++++- src/svm/svr.rs | 121 ++++++++++++++++++++ src/tree/decision_tree_classifier.rs | 160 +++++++++++++++++++++++++++ src/tree/decision_tree_regressor.rs | 137 +++++++++++++++++++++++ 9 files changed, 869 insertions(+), 8 deletions(-) diff --git a/src/naive_bayes/bernoulli.rs b/src/naive_bayes/bernoulli.rs index 95c4d369..29c6c84d 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -150,6 +150,88 @@ impl Default for BernoulliNBParameters { } } +/// BernoulliNB grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct BernoulliNBSearchParameters { + /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). + pub alpha: Vec, + /// Prior probabilities of the classes. If specified the priors are not adjusted according to the data + pub priors: Vec>>, + /// Threshold for binarizing (mapping to booleans) of sample features. If None, input is presumed to already consist of binary vectors. + pub binarize: Vec>, +} + +/// BernoulliNB grid search iterator +pub struct BernoulliNBSearchParametersIterator { + bernoulli_nb_search_parameters: BernoulliNBSearchParameters, + current_alpha: usize, + current_priors: usize, + current_binarize: usize, +} + +impl IntoIterator for BernoulliNBSearchParameters { + type Item = BernoulliNBParameters; + type IntoIter = BernoulliNBSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + BernoulliNBSearchParametersIterator { + bernoulli_nb_search_parameters: self, + current_alpha: 0, + current_priors: 0, + current_binarize: 0, + } + } +} + +impl Iterator for BernoulliNBSearchParametersIterator { + type Item = BernoulliNBParameters; + + fn next(&mut self) -> Option { + if self.current_alpha == self.bernoulli_nb_search_parameters.alpha.len() + && self.current_priors == self.bernoulli_nb_search_parameters.priors.len() + && self.current_binarize == self.bernoulli_nb_search_parameters.binarize.len() + { + return None; + } + + let next = BernoulliNBParameters { + alpha: self.bernoulli_nb_search_parameters.alpha[self.current_alpha], + priors: self.bernoulli_nb_search_parameters.priors[self.current_priors].clone(), + binarize: self.bernoulli_nb_search_parameters.binarize[self.current_binarize], + }; + + if self.current_alpha + 1 < self.bernoulli_nb_search_parameters.alpha.len() { + self.current_alpha += 1; + } else if self.current_priors + 1 < self.bernoulli_nb_search_parameters.priors.len() { + self.current_alpha = 0; + self.current_priors += 1; + } else if self.current_binarize + 1 < self.bernoulli_nb_search_parameters.binarize.len() { + self.current_alpha = 0; + self.current_priors = 0; + self.current_binarize += 1; + } else { + self.current_alpha += 1; + self.current_priors += 1; + self.current_binarize += 1; + } + + Some(next) + } +} + +impl Default for BernoulliNBSearchParameters { + fn default() -> Self { + let default_params = BernoulliNBParameters::default(); + + BernoulliNBSearchParameters { + alpha: vec![default_params.alpha], + priors: vec![default_params.priors], + binarize: vec![default_params.binarize], + } + } +} + impl BernoulliNBDistribution { /// Fits the distribution to a NxM matrix where N is number of samples and M is number of features. /// * `x` - training data. @@ -347,6 +429,20 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::DenseMatrix; + #[test] + fn search_parameters() { + let parameters = BernoulliNBSearchParameters { + alpha: vec![1., 2.], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.alpha, 1.); + let next = iter.next().unwrap(); + assert_eq!(next.alpha, 2.); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn run_bernoulli_naive_bayes() { diff --git a/src/naive_bayes/categorical.rs b/src/naive_bayes/categorical.rs index 87067028..78556889 100644 --- a/src/naive_bayes/categorical.rs +++ b/src/naive_bayes/categorical.rs @@ -261,6 +261,60 @@ impl Default for CategoricalNBParameters { } } +/// CategoricalNB grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct CategoricalNBSearchParameters { + /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). + pub alpha: Vec, +} + +/// CategoricalNB grid search iterator +pub struct CategoricalNBSearchParametersIterator { + categorical_nb_search_parameters: CategoricalNBSearchParameters, + current_alpha: usize, +} + +impl IntoIterator for CategoricalNBSearchParameters { + type Item = CategoricalNBParameters; + type IntoIter = CategoricalNBSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + CategoricalNBSearchParametersIterator { + categorical_nb_search_parameters: self, + current_alpha: 0, + } + } +} + +impl Iterator for CategoricalNBSearchParametersIterator { + type Item = CategoricalNBParameters; + + fn next(&mut self) -> Option { + if self.current_alpha == self.categorical_nb_search_parameters.alpha.len() { + return None; + } + + let next = CategoricalNBParameters { + alpha: self.categorical_nb_search_parameters.alpha[self.current_alpha], + }; + + self.current_alpha += 1; + + Some(next) + } +} + +impl Default for CategoricalNBSearchParameters { + fn default() -> Self { + let default_params = CategoricalNBParameters::default(); + + CategoricalNBSearchParameters { + alpha: vec![default_params.alpha], + } + } +} + /// CategoricalNB implements the categorical naive Bayes algorithm for categorically distributed data. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, PartialEq)] @@ -351,6 +405,20 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::DenseMatrix; + #[test] + fn search_parameters() { + let parameters = CategoricalNBSearchParameters { + alpha: vec![1., 2.], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.alpha, 1.); + let next = iter.next().unwrap(); + assert_eq!(next.alpha, 2.); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn run_categorical_naive_bayes() { diff --git a/src/naive_bayes/gaussian.rs b/src/naive_bayes/gaussian.rs index bd239190..24bbdd33 100644 --- a/src/naive_bayes/gaussian.rs +++ b/src/naive_bayes/gaussian.rs @@ -76,7 +76,7 @@ impl> NBDistribution for GaussianNBDistributio /// `GaussianNB` parameters. Use `Default::default()` for default values. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Default, Clone)] +#[derive(Debug, Clone)] pub struct GaussianNBParameters { /// Prior probabilities of the classes. If specified the priors are not adjusted according to the data pub priors: Option>, @@ -90,6 +90,66 @@ impl GaussianNBParameters { } } +impl Default for GaussianNBParameters { + fn default() -> Self { + Self { priors: None } + } +} + +/// GaussianNB grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct GaussianNBSearchParameters { + /// Prior probabilities of the classes. If specified the priors are not adjusted according to the data + pub priors: Vec>>, +} + +/// GaussianNB grid search iterator +pub struct GaussianNBSearchParametersIterator { + gaussian_nb_search_parameters: GaussianNBSearchParameters, + current_priors: usize, +} + +impl IntoIterator for GaussianNBSearchParameters { + type Item = GaussianNBParameters; + type IntoIter = GaussianNBSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + GaussianNBSearchParametersIterator { + gaussian_nb_search_parameters: self, + current_priors: 0, + } + } +} + +impl Iterator for GaussianNBSearchParametersIterator { + type Item = GaussianNBParameters; + + fn next(&mut self) -> Option { + if self.current_priors == self.gaussian_nb_search_parameters.priors.len() { + return None; + } + + let next = GaussianNBParameters { + priors: self.gaussian_nb_search_parameters.priors[self.current_priors].clone(), + }; + + self.current_priors += 1; + + Some(next) + } +} + +impl Default for GaussianNBSearchParameters { + fn default() -> Self { + let default_params = GaussianNBParameters::default(); + + GaussianNBSearchParameters { + priors: vec![default_params.priors], + } + } +} + impl GaussianNBDistribution { /// Fits the distribution to a NxM matrix where N is number of samples and M is number of features. /// * `x` - training data. @@ -260,6 +320,20 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::DenseMatrix; + #[test] + fn search_parameters() { + let parameters = GaussianNBSearchParameters { + priors: vec![Some(vec![1.]), Some(vec![2.])], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.priors, Some(vec![1.])); + let next = iter.next().unwrap(); + assert_eq!(next.priors, Some(vec![2.])); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn run_gaussian_naive_bayes() { diff --git a/src/naive_bayes/multinomial.rs b/src/naive_bayes/multinomial.rs index f42b99e1..6e846c1a 100644 --- a/src/naive_bayes/multinomial.rs +++ b/src/naive_bayes/multinomial.rs @@ -114,6 +114,76 @@ impl Default for MultinomialNBParameters { } } +/// MultinomialNB grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct MultinomialNBSearchParameters { + /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). + pub alpha: Vec, + /// Prior probabilities of the classes. If specified the priors are not adjusted according to the data + pub priors: Vec>>, +} + +/// MultinomialNB grid search iterator +pub struct MultinomialNBSearchParametersIterator { + multinomial_nb_search_parameters: MultinomialNBSearchParameters, + current_alpha: usize, + current_priors: usize, +} + +impl IntoIterator for MultinomialNBSearchParameters { + type Item = MultinomialNBParameters; + type IntoIter = MultinomialNBSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + MultinomialNBSearchParametersIterator { + multinomial_nb_search_parameters: self, + current_alpha: 0, + current_priors: 0, + } + } +} + +impl Iterator for MultinomialNBSearchParametersIterator { + type Item = MultinomialNBParameters; + + fn next(&mut self) -> Option { + if self.current_alpha == self.multinomial_nb_search_parameters.alpha.len() + && self.current_priors == self.multinomial_nb_search_parameters.priors.len() + { + return None; + } + + let next = MultinomialNBParameters { + alpha: self.multinomial_nb_search_parameters.alpha[self.current_alpha], + priors: self.multinomial_nb_search_parameters.priors[self.current_priors].clone(), + }; + + if self.current_alpha + 1 < self.multinomial_nb_search_parameters.alpha.len() { + self.current_alpha += 1; + } else if self.current_priors + 1 < self.multinomial_nb_search_parameters.priors.len() { + self.current_alpha = 0; + self.current_priors += 1; + } else { + self.current_alpha += 1; + self.current_priors += 1; + } + + Some(next) + } +} + +impl Default for MultinomialNBSearchParameters { + fn default() -> Self { + let default_params = MultinomialNBParameters::default(); + + MultinomialNBSearchParameters { + alpha: vec![default_params.alpha], + priors: vec![default_params.priors], + } + } +} + impl MultinomialNBDistribution { /// Fits the distribution to a NxM matrix where N is number of samples and M is number of features. /// * `x` - training data. @@ -297,6 +367,20 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::DenseMatrix; + #[test] + fn search_parameters() { + let parameters = MultinomialNBSearchParameters { + alpha: vec![1., 2.], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.alpha, 1.); + let next = iter.next().unwrap(); + assert_eq!(next.alpha, 2.); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn run_multinomial_naive_bayes() { diff --git a/src/svm/mod.rs b/src/svm/mod.rs index 55df5840..4c71b3f2 100644 --- a/src/svm/mod.rs +++ b/src/svm/mod.rs @@ -33,7 +33,7 @@ use crate::linalg::BaseVector; use crate::math::num::RealNumber; /// Defines a kernel function -pub trait Kernel> { +pub trait Kernel>: Clone { /// Apply kernel function to x_i and x_j fn apply(&self, x_i: &V, x_j: &V) -> T; } @@ -95,12 +95,12 @@ impl Kernels { /// Linear Kernel #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, PartialEq, Eq)] pub struct LinearKernel {} /// Radial basis function (Gaussian) kernel #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, PartialEq, Eq)] pub struct RBFKernel { /// kernel coefficient pub gamma: T, @@ -108,7 +108,7 @@ pub struct RBFKernel { /// Polynomial kernel #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, PartialEq, Eq)] pub struct PolynomialKernel { /// degree of the polynomial pub degree: T, @@ -120,7 +120,7 @@ pub struct PolynomialKernel { /// Sigmoid (hyperbolic tangent) kernel #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, PartialEq, Eq)] pub struct SigmoidKernel { /// kernel coefficient pub gamma: T, diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 87fb7431..7cfc1db0 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -94,7 +94,7 @@ pub struct SVCParameters, K: Kernel pub epoch: usize, /// Regularization parameter. pub c: T, - /// Tolerance for stopping criterion. + /// Tolerance for stopping epoch. pub tol: T, /// The kernel function. pub kernel: K, @@ -102,6 +102,109 @@ pub struct SVCParameters, K: Kernel m: PhantomData, } +/// SVC grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct SVCSearchParameters, K: Kernel> { + /// Number of epochs. + pub epoch: Vec, + /// Regularization parameter. + pub c: Vec, + /// Tolerance for stopping epoch. + pub tol: Vec, + /// The kernel function. + pub kernel: Vec, + /// Unused parameter. + m: PhantomData, +} + +/// SVC grid search iterator +pub struct SVCSearchParametersIterator, K: Kernel> { + svc_search_parameters: SVCSearchParameters, + current_epoch: usize, + current_c: usize, + current_tol: usize, + current_kernel: usize, +} + +impl, K: Kernel> IntoIterator + for SVCSearchParameters +{ + type Item = SVCParameters; + type IntoIter = SVCSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + SVCSearchParametersIterator { + svc_search_parameters: self, + current_epoch: 0, + current_c: 0, + current_tol: 0, + current_kernel: 0, + } + } +} + +impl, K: Kernel> Iterator + for SVCSearchParametersIterator +{ + type Item = SVCParameters; + + fn next(&mut self) -> Option { + if self.current_epoch == self.svc_search_parameters.epoch.len() + && self.current_c == self.svc_search_parameters.c.len() + && self.current_tol == self.svc_search_parameters.tol.len() + && self.current_kernel == self.svc_search_parameters.kernel.len() + { + return None; + } + + let next = SVCParameters:: { + epoch: self.svc_search_parameters.epoch[self.current_epoch], + c: self.svc_search_parameters.c[self.current_c], + tol: self.svc_search_parameters.tol[self.current_tol], + kernel: self.svc_search_parameters.kernel[self.current_kernel].clone(), + m: PhantomData, + }; + + if self.current_epoch + 1 < self.svc_search_parameters.epoch.len() { + self.current_epoch += 1; + } else if self.current_c + 1 < self.svc_search_parameters.c.len() { + self.current_epoch = 0; + self.current_c += 1; + } else if self.current_tol + 1 < self.svc_search_parameters.tol.len() { + self.current_epoch = 0; + self.current_c = 0; + self.current_tol += 1; + } else if self.current_kernel + 1 < self.svc_search_parameters.kernel.len() { + self.current_epoch = 0; + self.current_c = 0; + self.current_tol = 0; + self.current_kernel += 1; + } else { + self.current_epoch += 1; + self.current_c += 1; + self.current_tol += 1; + self.current_kernel += 1; + } + + Some(next) + } +} + +impl> Default for SVCSearchParameters { + fn default() -> Self { + let default_params: SVCParameters = SVCParameters::default(); + + SVCSearchParameters { + epoch: vec![default_params.epoch], + c: vec![default_params.c], + tol: vec![default_params.tol], + kernel: vec![default_params.kernel], + m: PhantomData, + } + } +} + #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug)] #[cfg_attr( @@ -163,7 +266,7 @@ impl, K: Kernel> SVCParameters Self { self.tol = tol; self @@ -737,6 +840,24 @@ mod tests { #[cfg(feature = "serde")] use crate::svm::*; + #[test] + fn search_parameters() { + let parameters: SVCSearchParameters, LinearKernel> = + SVCSearchParameters { + epoch: vec![10, 100], + kernel: vec![LinearKernel {}], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.epoch, 10); + assert_eq!(next.kernel, LinearKernel {}); + let next = iter.next().unwrap(); + assert_eq!(next.epoch, 100); + assert_eq!(next.kernel, LinearKernel {}); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn svc_fit_predict() { diff --git a/src/svm/svr.rs b/src/svm/svr.rs index 18c73d11..25326d4c 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -94,6 +94,109 @@ pub struct SVRParameters, K: Kernel m: PhantomData, } +/// SVR grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct SVRSearchParameters, K: Kernel> { + /// Epsilon in the epsilon-SVR model. + pub eps: Vec, + /// Regularization parameter. + pub c: Vec, + /// Tolerance for stopping eps. + pub tol: Vec, + /// The kernel function. + pub kernel: Vec, + /// Unused parameter. + m: PhantomData, +} + +/// SVR grid search iterator +pub struct SVRSearchParametersIterator, K: Kernel> { + svr_search_parameters: SVRSearchParameters, + current_eps: usize, + current_c: usize, + current_tol: usize, + current_kernel: usize, +} + +impl, K: Kernel> IntoIterator + for SVRSearchParameters +{ + type Item = SVRParameters; + type IntoIter = SVRSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + SVRSearchParametersIterator { + svr_search_parameters: self, + current_eps: 0, + current_c: 0, + current_tol: 0, + current_kernel: 0, + } + } +} + +impl, K: Kernel> Iterator + for SVRSearchParametersIterator +{ + type Item = SVRParameters; + + fn next(&mut self) -> Option { + if self.current_eps == self.svr_search_parameters.eps.len() + && self.current_c == self.svr_search_parameters.c.len() + && self.current_tol == self.svr_search_parameters.tol.len() + && self.current_kernel == self.svr_search_parameters.kernel.len() + { + return None; + } + + let next = SVRParameters:: { + eps: self.svr_search_parameters.eps[self.current_eps], + c: self.svr_search_parameters.c[self.current_c], + tol: self.svr_search_parameters.tol[self.current_tol], + kernel: self.svr_search_parameters.kernel[self.current_kernel].clone(), + m: PhantomData, + }; + + if self.current_eps + 1 < self.svr_search_parameters.eps.len() { + self.current_eps += 1; + } else if self.current_c + 1 < self.svr_search_parameters.c.len() { + self.current_eps = 0; + self.current_c += 1; + } else if self.current_tol + 1 < self.svr_search_parameters.tol.len() { + self.current_eps = 0; + self.current_c = 0; + self.current_tol += 1; + } else if self.current_kernel + 1 < self.svr_search_parameters.kernel.len() { + self.current_eps = 0; + self.current_c = 0; + self.current_tol = 0; + self.current_kernel += 1; + } else { + self.current_eps += 1; + self.current_c += 1; + self.current_tol += 1; + self.current_kernel += 1; + } + + Some(next) + } +} + +impl> Default for SVRSearchParameters { + fn default() -> Self { + let default_params: SVRParameters = SVRParameters::default(); + + SVRSearchParameters { + eps: vec![default_params.eps], + c: vec![default_params.c], + tol: vec![default_params.tol], + kernel: vec![default_params.kernel], + m: PhantomData, + } + } +} + #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug)] #[cfg_attr( @@ -536,6 +639,24 @@ mod tests { #[cfg(feature = "serde")] use crate::svm::*; + #[test] + fn search_parameters() { + let parameters: SVRSearchParameters, LinearKernel> = + SVRSearchParameters { + eps: vec![0., 1.], + kernel: vec![LinearKernel {}], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.eps, 0.); + assert_eq!(next.kernel, LinearKernel {}); + let next = iter.next().unwrap(); + assert_eq!(next.eps, 1.); + assert_eq!(next.kernel, LinearKernel {}); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn svr_fit_predict() { diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index d58125d4..278ad7f8 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -201,6 +201,143 @@ impl Default for DecisionTreeClassifierParameters { } } +/// DecisionTreeClassifier grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct DecisionTreeClassifierSearchParameters { + /// Split criteria to use when building a tree. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub criterion: Vec, + /// Tree max depth. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub max_depth: Vec>, + /// The minimum number of samples required to be at a leaf node. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub min_samples_leaf: Vec, + /// The minimum number of samples required to split an internal node. See [Decision Tree Classifier](../../tree/decision_tree_classifier/index.html) + pub min_samples_split: Vec, +} + +/// DecisionTreeClassifier grid search iterator +pub struct DecisionTreeClassifierSearchParametersIterator { + decision_tree_classifier_search_parameters: DecisionTreeClassifierSearchParameters, + current_criterion: usize, + current_max_depth: usize, + current_min_samples_leaf: usize, + current_min_samples_split: usize, +} + +impl IntoIterator for DecisionTreeClassifierSearchParameters { + type Item = DecisionTreeClassifierParameters; + type IntoIter = DecisionTreeClassifierSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + DecisionTreeClassifierSearchParametersIterator { + decision_tree_classifier_search_parameters: self, + current_criterion: 0, + current_max_depth: 0, + current_min_samples_leaf: 0, + current_min_samples_split: 0, + } + } +} + +impl Iterator for DecisionTreeClassifierSearchParametersIterator { + type Item = DecisionTreeClassifierParameters; + + fn next(&mut self) -> Option { + if self.current_criterion + == self + .decision_tree_classifier_search_parameters + .criterion + .len() + && self.current_max_depth + == self + .decision_tree_classifier_search_parameters + .max_depth + .len() + && self.current_min_samples_leaf + == self + .decision_tree_classifier_search_parameters + .min_samples_leaf + .len() + && self.current_min_samples_split + == self + .decision_tree_classifier_search_parameters + .min_samples_split + .len() + { + return None; + } + + let next = DecisionTreeClassifierParameters { + criterion: self.decision_tree_classifier_search_parameters.criterion + [self.current_criterion], + max_depth: self.decision_tree_classifier_search_parameters.max_depth + [self.current_max_depth], + min_samples_leaf: self + .decision_tree_classifier_search_parameters + .min_samples_leaf[self.current_min_samples_leaf], + min_samples_split: self + .decision_tree_classifier_search_parameters + .min_samples_split[self.current_min_samples_split], + }; + + if self.current_criterion + 1 + < self + .decision_tree_classifier_search_parameters + .criterion + .len() + { + self.current_criterion += 1; + } else if self.current_max_depth + 1 + < self + .decision_tree_classifier_search_parameters + .max_depth + .len() + { + self.current_criterion = 0; + self.current_max_depth += 1; + } else if self.current_min_samples_leaf + 1 + < self + .decision_tree_classifier_search_parameters + .min_samples_leaf + .len() + { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf += 1; + } else if self.current_min_samples_split + 1 + < self + .decision_tree_classifier_search_parameters + .min_samples_split + .len() + { + self.current_criterion = 0; + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split += 1; + } else { + self.current_criterion += 1; + self.current_max_depth += 1; + self.current_min_samples_leaf += 1; + self.current_min_samples_split += 1; + } + + Some(next) + } +} + +impl Default for DecisionTreeClassifierSearchParameters { + fn default() -> Self { + let default_params = DecisionTreeClassifierParameters::default(); + + DecisionTreeClassifierSearchParameters { + criterion: vec![default_params.criterion], + max_depth: vec![default_params.max_depth], + min_samples_leaf: vec![default_params.min_samples_leaf], + min_samples_split: vec![default_params.min_samples_split], + } + } +} + impl Node { fn new(index: usize, output: usize) -> Self { Node { @@ -651,6 +788,29 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::DenseMatrix; + #[test] + fn search_parameters() { + let parameters = DecisionTreeClassifierSearchParameters { + max_depth: vec![Some(10), Some(100)], + min_samples_split: vec![1, 2], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(10)); + assert_eq!(next.min_samples_split, 1); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(100)); + assert_eq!(next.min_samples_split, 1); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(10)); + assert_eq!(next.min_samples_split, 2); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(100)); + assert_eq!(next.min_samples_split, 2); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn gini_impurity() { diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index 25f5e7e5..f48de33f 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -134,6 +134,120 @@ impl Default for DecisionTreeRegressorParameters { } } +/// DecisionTreeRegressor grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct DecisionTreeRegressorSearchParameters { + /// Tree max depth. See [Decision Tree Regressor](../../tree/decision_tree_regressor/index.html) + pub max_depth: Vec>, + /// The minimum number of samples required to be at a leaf node. See [Decision Tree Regressor](../../tree/decision_tree_regressor/index.html) + pub min_samples_leaf: Vec, + /// The minimum number of samples required to split an internal node. See [Decision Tree Regressor](../../tree/decision_tree_regressor/index.html) + pub min_samples_split: Vec, +} + +/// DecisionTreeRegressor grid search iterator +pub struct DecisionTreeRegressorSearchParametersIterator { + decision_tree_regressor_search_parameters: DecisionTreeRegressorSearchParameters, + current_max_depth: usize, + current_min_samples_leaf: usize, + current_min_samples_split: usize, +} + +impl IntoIterator for DecisionTreeRegressorSearchParameters { + type Item = DecisionTreeRegressorParameters; + type IntoIter = DecisionTreeRegressorSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + DecisionTreeRegressorSearchParametersIterator { + decision_tree_regressor_search_parameters: self, + current_max_depth: 0, + current_min_samples_leaf: 0, + current_min_samples_split: 0, + } + } +} + +impl Iterator for DecisionTreeRegressorSearchParametersIterator { + type Item = DecisionTreeRegressorParameters; + + fn next(&mut self) -> Option { + if self.current_max_depth + == self + .decision_tree_regressor_search_parameters + .max_depth + .len() + && self.current_min_samples_leaf + == self + .decision_tree_regressor_search_parameters + .min_samples_leaf + .len() + && self.current_min_samples_split + == self + .decision_tree_regressor_search_parameters + .min_samples_split + .len() + { + return None; + } + + let next = DecisionTreeRegressorParameters { + max_depth: self.decision_tree_regressor_search_parameters.max_depth + [self.current_max_depth], + min_samples_leaf: self + .decision_tree_regressor_search_parameters + .min_samples_leaf[self.current_min_samples_leaf], + min_samples_split: self + .decision_tree_regressor_search_parameters + .min_samples_split[self.current_min_samples_split], + }; + + if self.current_max_depth + 1 + < self + .decision_tree_regressor_search_parameters + .max_depth + .len() + { + self.current_max_depth += 1; + } else if self.current_min_samples_leaf + 1 + < self + .decision_tree_regressor_search_parameters + .min_samples_leaf + .len() + { + self.current_max_depth = 0; + self.current_min_samples_leaf += 1; + } else if self.current_min_samples_split + 1 + < self + .decision_tree_regressor_search_parameters + .min_samples_split + .len() + { + self.current_max_depth = 0; + self.current_min_samples_leaf = 0; + self.current_min_samples_split += 1; + } else { + self.current_max_depth += 1; + self.current_min_samples_leaf += 1; + self.current_min_samples_split += 1; + } + + Some(next) + } +} + +impl Default for DecisionTreeRegressorSearchParameters { + fn default() -> Self { + let default_params = DecisionTreeRegressorParameters::default(); + + DecisionTreeRegressorSearchParameters { + max_depth: vec![default_params.max_depth], + min_samples_leaf: vec![default_params.min_samples_leaf], + min_samples_split: vec![default_params.min_samples_split], + } + } +} + impl Node { fn new(index: usize, output: T) -> Self { Node { @@ -517,6 +631,29 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::DenseMatrix; + #[test] + fn search_parameters() { + let parameters = DecisionTreeRegressorSearchParameters { + max_depth: vec![Some(10), Some(100)], + min_samples_split: vec![1, 2], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(10)); + assert_eq!(next.min_samples_split, 1); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(100)); + assert_eq!(next.min_samples_split, 1); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(10)); + assert_eq!(next.min_samples_split, 2); + let next = iter.next().unwrap(); + assert_eq!(next.max_depth, Some(100)); + assert_eq!(next.min_samples_split, 2); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn fit_longley() { From 11cf75540bc2fed33c4cdfe8d99bca3d0bec583c Mon Sep 17 00:00:00 2001 From: Montana Low Date: Thu, 15 Sep 2022 11:27:23 -0700 Subject: [PATCH 5/6] grid search for unsupervised algos --- src/cluster/dbscan.rs | 120 ++++++++++++++++++++++++++++ src/cluster/kmeans.rs | 93 +++++++++++++++++++++ src/decomposition/pca.rs | 98 +++++++++++++++++++++++ src/decomposition/svd.rs | 68 ++++++++++++++++ src/model_selection/hyper_tuning.rs | 2 +- src/model_selection/mod.rs | 1 - 6 files changed, 380 insertions(+), 2 deletions(-) diff --git a/src/cluster/dbscan.rs b/src/cluster/dbscan.rs index 7f2baef0..621d0173 100644 --- a/src/cluster/dbscan.rs +++ b/src/cluster/dbscan.rs @@ -109,6 +109,103 @@ impl, T>> DBSCANParameters { } } +/// DBSCAN grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct DBSCANSearchParameters, T>> { + /// a function that defines a distance between each pair of point in training data. + /// This function should extend [`Distance`](../../math/distance/trait.Distance.html) trait. + /// See [`Distances`](../../math/distance/struct.Distances.html) for a list of available functions. + pub distance: Vec, + /// The number of samples (or total weight) in a neighborhood for a point to be considered as a core point. + pub min_samples: Vec, + /// The maximum distance between two samples for one to be considered as in the neighborhood of the other. + pub eps: Vec, + /// KNN algorithm to use. + pub algorithm: Vec, +} + +/// DBSCAN grid search iterator +pub struct DBSCANSearchParametersIterator, T>> { + dbscan_search_parameters: DBSCANSearchParameters, + current_distance: usize, + current_min_samples: usize, + current_eps: usize, + current_algorithm: usize, +} + +impl, T>> IntoIterator for DBSCANSearchParameters { + type Item = DBSCANParameters; + type IntoIter = DBSCANSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + DBSCANSearchParametersIterator { + dbscan_search_parameters: self, + current_distance: 0, + current_min_samples: 0, + current_eps: 0, + current_algorithm: 0, + } + } +} + +impl, T>> Iterator for DBSCANSearchParametersIterator { + type Item = DBSCANParameters; + + fn next(&mut self) -> Option { + if self.current_distance == self.dbscan_search_parameters.distance.len() + && self.current_min_samples == self.dbscan_search_parameters.min_samples.len() + && self.current_eps == self.dbscan_search_parameters.eps.len() + && self.current_algorithm == self.dbscan_search_parameters.algorithm.len() + { + return None; + } + + let next = DBSCANParameters { + distance: self.dbscan_search_parameters.distance[self.current_distance].clone(), + min_samples: self.dbscan_search_parameters.min_samples[self.current_min_samples], + eps: self.dbscan_search_parameters.eps[self.current_eps], + algorithm: self.dbscan_search_parameters.algorithm[self.current_algorithm].clone(), + }; + + if self.current_distance + 1 < self.dbscan_search_parameters.distance.len() { + self.current_distance += 1; + } else if self.current_min_samples + 1 < self.dbscan_search_parameters.min_samples.len() { + self.current_distance = 0; + self.current_min_samples += 1; + } else if self.current_eps + 1 < self.dbscan_search_parameters.eps.len() { + self.current_distance = 0; + self.current_min_samples = 0; + self.current_eps += 1; + } else if self.current_algorithm + 1 < self.dbscan_search_parameters.algorithm.len() { + self.current_distance = 0; + self.current_min_samples = 0; + self.current_eps = 0; + self.current_algorithm += 1; + } else { + self.current_distance += 1; + self.current_min_samples += 1; + self.current_eps += 1; + self.current_algorithm += 1; + } + + Some(next) + } +} + +impl Default for DBSCANSearchParameters { + fn default() -> Self { + let default_params = DBSCANParameters::default(); + + DBSCANSearchParameters { + distance: vec![default_params.distance], + min_samples: vec![default_params.min_samples], + eps: vec![default_params.eps], + algorithm: vec![default_params.algorithm], + } + } +} + impl, T>> PartialEq for DBSCAN { fn eq(&self, other: &Self) -> bool { self.cluster_labels.len() == other.cluster_labels.len() @@ -268,6 +365,29 @@ mod tests { #[cfg(feature = "serde")] use crate::math::distance::euclidian::Euclidian; + #[test] + fn search_parameters() { + let parameters = DBSCANSearchParameters { + min_samples: vec![10, 100], + eps: vec![1., 2.], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.min_samples, 10); + assert_eq!(next.eps, 1.); + let next = iter.next().unwrap(); + assert_eq!(next.min_samples, 100); + assert_eq!(next.eps, 1.); + let next = iter.next().unwrap(); + assert_eq!(next.min_samples, 10); + assert_eq!(next.eps, 2.); + let next = iter.next().unwrap(); + assert_eq!(next.min_samples, 100); + assert_eq!(next.eps, 2.); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn fit_predict_dbscan() { diff --git a/src/cluster/kmeans.rs b/src/cluster/kmeans.rs index 05af6809..8ecbb2e9 100644 --- a/src/cluster/kmeans.rs +++ b/src/cluster/kmeans.rs @@ -132,6 +132,76 @@ impl Default for KMeansParameters { } } +/// KMeans grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct KMeansSearchParameters { + /// Number of clusters. + pub k: Vec, + /// Maximum number of iterations of the k-means algorithm for a single run. + pub max_iter: Vec, +} + +/// KMeans grid search iterator +pub struct KMeansSearchParametersIterator { + kmeans_search_parameters: KMeansSearchParameters, + current_k: usize, + current_max_iter: usize, +} + +impl IntoIterator for KMeansSearchParameters { + type Item = KMeansParameters; + type IntoIter = KMeansSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + KMeansSearchParametersIterator { + kmeans_search_parameters: self, + current_k: 0, + current_max_iter: 0, + } + } +} + +impl Iterator for KMeansSearchParametersIterator { + type Item = KMeansParameters; + + fn next(&mut self) -> Option { + if self.current_k == self.kmeans_search_parameters.k.len() + && self.current_max_iter == self.kmeans_search_parameters.max_iter.len() + { + return None; + } + + let next = KMeansParameters { + k: self.kmeans_search_parameters.k[self.current_k], + max_iter: self.kmeans_search_parameters.max_iter[self.current_max_iter], + }; + + if self.current_k + 1 < self.kmeans_search_parameters.k.len() { + self.current_k += 1; + } else if self.current_max_iter + 1 < self.kmeans_search_parameters.max_iter.len() { + self.current_k = 0; + self.current_max_iter += 1; + } else { + self.current_k += 1; + self.current_max_iter += 1; + } + + Some(next) + } +} + +impl Default for KMeansSearchParameters { + fn default() -> Self { + let default_params = KMeansParameters::default(); + + KMeansSearchParameters { + k: vec![default_params.k], + max_iter: vec![default_params.max_iter], + } + } +} + impl> UnsupervisedEstimator for KMeans { fn fit(x: &M, parameters: KMeansParameters) -> Result { KMeans::fit(x, parameters) @@ -313,6 +383,29 @@ mod tests { ); } + #[test] + fn search_parameters() { + let parameters = KMeansSearchParameters { + k: vec![2, 4], + max_iter: vec![10, 100], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.k, 2); + assert_eq!(next.max_iter, 10); + let next = iter.next().unwrap(); + assert_eq!(next.k, 4); + assert_eq!(next.max_iter, 10); + let next = iter.next().unwrap(); + assert_eq!(next.k, 2); + assert_eq!(next.max_iter, 100); + let next = iter.next().unwrap(); + assert_eq!(next.k, 4); + assert_eq!(next.max_iter, 100); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn fit_predict_iris() { diff --git a/src/decomposition/pca.rs b/src/decomposition/pca.rs index 9aebae20..296926a4 100644 --- a/src/decomposition/pca.rs +++ b/src/decomposition/pca.rs @@ -116,6 +116,81 @@ impl Default for PCAParameters { } } +/// PCA grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct PCASearchParameters { + /// Number of components to keep. + pub n_components: Vec, + /// By default, covariance matrix is used to compute principal components. + /// Enable this flag if you want to use correlation matrix instead. + pub use_correlation_matrix: Vec, +} + +/// PCA grid search iterator +pub struct PCASearchParametersIterator { + pca_search_parameters: PCASearchParameters, + current_k: usize, + current_use_correlation_matrix: usize, +} + +impl IntoIterator for PCASearchParameters { + type Item = PCAParameters; + type IntoIter = PCASearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + PCASearchParametersIterator { + pca_search_parameters: self, + current_k: 0, + current_use_correlation_matrix: 0, + } + } +} + +impl Iterator for PCASearchParametersIterator { + type Item = PCAParameters; + + fn next(&mut self) -> Option { + if self.current_k == self.pca_search_parameters.n_components.len() + && self.current_use_correlation_matrix + == self.pca_search_parameters.use_correlation_matrix.len() + { + return None; + } + + let next = PCAParameters { + n_components: self.pca_search_parameters.n_components[self.current_k], + use_correlation_matrix: self.pca_search_parameters.use_correlation_matrix + [self.current_use_correlation_matrix], + }; + + if self.current_k + 1 < self.pca_search_parameters.n_components.len() { + self.current_k += 1; + } else if self.current_use_correlation_matrix + 1 + < self.pca_search_parameters.use_correlation_matrix.len() + { + self.current_k = 0; + self.current_use_correlation_matrix += 1; + } else { + self.current_k += 1; + self.current_use_correlation_matrix += 1; + } + + Some(next) + } +} + +impl Default for PCASearchParameters { + fn default() -> Self { + let default_params = PCAParameters::default(); + + PCASearchParameters { + n_components: vec![default_params.n_components], + use_correlation_matrix: vec![default_params.use_correlation_matrix], + } + } +} + impl> UnsupervisedEstimator for PCA { fn fit(x: &M, parameters: PCAParameters) -> Result { PCA::fit(x, parameters) @@ -271,6 +346,29 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::*; + #[test] + fn search_parameters() { + let parameters = PCASearchParameters { + n_components: vec![2, 4], + use_correlation_matrix: vec![true, false], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.n_components, 2); + assert_eq!(next.use_correlation_matrix, true); + let next = iter.next().unwrap(); + assert_eq!(next.n_components, 4); + assert_eq!(next.use_correlation_matrix, true); + let next = iter.next().unwrap(); + assert_eq!(next.n_components, 2); + assert_eq!(next.use_correlation_matrix, false); + let next = iter.next().unwrap(); + assert_eq!(next.n_components, 4); + assert_eq!(next.use_correlation_matrix, false); + assert!(iter.next().is_none()); + } + fn us_arrests_data() -> DenseMatrix { DenseMatrix::from_2d_array(&[ &[13.2, 236.0, 58.0, 21.2], diff --git a/src/decomposition/svd.rs b/src/decomposition/svd.rs index 38077603..3001fd9e 100644 --- a/src/decomposition/svd.rs +++ b/src/decomposition/svd.rs @@ -90,6 +90,60 @@ impl SVDParameters { } } +/// SVD grid search parameters +#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[derive(Debug, Clone)] +pub struct SVDSearchParameters { + /// Maximum number of iterations of the k-means algorithm for a single run. + pub n_components: Vec, +} + +/// SVD grid search iterator +pub struct SVDSearchParametersIterator { + svd_search_parameters: SVDSearchParameters, + current_n_components: usize, +} + +impl IntoIterator for SVDSearchParameters { + type Item = SVDParameters; + type IntoIter = SVDSearchParametersIterator; + + fn into_iter(self) -> Self::IntoIter { + SVDSearchParametersIterator { + svd_search_parameters: self, + current_n_components: 0, + } + } +} + +impl Iterator for SVDSearchParametersIterator { + type Item = SVDParameters; + + fn next(&mut self) -> Option { + if self.current_n_components == self.svd_search_parameters.n_components.len() { + return None; + } + + let next = SVDParameters { + n_components: self.svd_search_parameters.n_components[self.current_n_components], + }; + + self.current_n_components += 1; + + Some(next) + } +} + +impl Default for SVDSearchParameters { + fn default() -> Self { + let default_params = SVDParameters::default(); + + SVDSearchParameters { + n_components: vec![default_params.n_components], + } + } +} + impl> UnsupervisedEstimator for SVD { fn fit(x: &M, parameters: SVDParameters) -> Result { SVD::fit(x, parameters) @@ -153,6 +207,20 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::*; + #[test] + fn search_parameters() { + let parameters = SVDSearchParameters { + n_components: vec![10, 100], + ..Default::default() + }; + let mut iter = parameters.into_iter(); + let next = iter.next().unwrap(); + assert_eq!(next.n_components, 10); + let next = iter.next().unwrap(); + assert_eq!(next.n_components, 100); + assert!(iter.next().is_none()); + } + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] #[test] fn svd_decompose() { diff --git a/src/model_selection/hyper_tuning.rs b/src/model_selection/hyper_tuning.rs index 3093fbdd..cb69da18 100644 --- a/src/model_selection/hyper_tuning.rs +++ b/src/model_selection/hyper_tuning.rs @@ -114,4 +114,4 @@ mod tests { assert!([0., 1.].contains(&results.parameters.alpha)); } -} \ No newline at end of file +} diff --git a/src/model_selection/mod.rs b/src/model_selection/mod.rs index 68f06350..6f737d6a 100644 --- a/src/model_selection/mod.rs +++ b/src/model_selection/mod.rs @@ -281,7 +281,6 @@ mod tests { use super::*; use crate::linalg::naive::dense_matrix::*; - use crate::metrics::{accuracy, mean_absolute_error}; use crate::model_selection::kfold::KFold; use crate::neighbors::knn_regressor::KNNRegressor; From 08d74565ff985ffa555802882b0316de236bd240 Mon Sep 17 00:00:00 2001 From: Montana Low Date: Thu, 15 Sep 2022 11:34:29 -0700 Subject: [PATCH 6/6] minor cleanup --- src/ensemble/random_forest_classifier.rs | 5 +++-- src/svm/svc.rs | 4 ++-- src/tree/decision_tree_classifier.rs | 5 +++-- 3 files changed, 8 insertions(+), 6 deletions(-) diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index 53c48451..a4d6e75d 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -289,7 +289,8 @@ impl Iterator for RandomForestClassifierSearchParametersIterator { let next = RandomForestClassifierParameters { criterion: self.random_forest_classifier_search_parameters.criterion - [self.current_criterion], + [self.current_criterion] + .clone(), max_depth: self.random_forest_classifier_search_parameters.max_depth [self.current_max_depth], min_samples_leaf: self @@ -457,7 +458,7 @@ impl RandomForestClassifier { } let params = DecisionTreeClassifierParameters { - criterion: parameters.criterion, + criterion: parameters.criterion.clone(), max_depth: parameters.max_depth, min_samples_leaf: parameters.min_samples_leaf, min_samples_split: parameters.min_samples_split, diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 7cfc1db0..46b0b68c 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -94,7 +94,7 @@ pub struct SVCParameters, K: Kernel pub epoch: usize, /// Regularization parameter. pub c: T, - /// Tolerance for stopping epoch. + /// Tolerance for stopping criterion. pub tol: T, /// The kernel function. pub kernel: K, @@ -266,7 +266,7 @@ impl, K: Kernel> SVCParameters Self { self.tol = tol; self diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index 278ad7f8..a1699afa 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -105,7 +105,7 @@ pub struct DecisionTreeClassifier { /// The function to measure the quality of a split. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone)] pub enum SplitCriterion { /// [Gini index](../decision_tree_classifier/index.html) Gini, @@ -269,7 +269,8 @@ impl Iterator for DecisionTreeClassifierSearchParametersIterator { let next = DecisionTreeClassifierParameters { criterion: self.decision_tree_classifier_search_parameters.criterion - [self.current_criterion], + [self.current_criterion] + .clone(), max_depth: self.decision_tree_classifier_search_parameters.max_depth [self.current_max_depth], min_samples_leaf: self