From 291b30ce48a074a84946ca7a5b351a50755006a1 Mon Sep 17 00:00:00 2001 From: Hubert Majewski <36614010+HubertMajewski@users.noreply.github.com> Date: Thu, 9 Jul 2026 03:58:52 +0000 Subject: [PATCH 1/4] Add getter for parameters in all model structs. --- src/cluster/agglomerative.rs | 11 +++++++++++ src/cluster/dbscan.rs | 13 ++++++++++++- src/cluster/kmeans.rs | 11 +++++++++++ src/decomposition/pca.rs | 11 +++++++++++ src/decomposition/svd.rs | 11 +++++++++++ src/ensemble/extra_trees_regressor.rs | 12 ++++++++++++ src/ensemble/random_forest_classifier.rs | 12 ++++++++++++ src/ensemble/random_forest_regressor.rs | 12 ++++++++++++ src/linear/elastic_net.rs | 12 ++++++++++++ src/linear/lasso.rs | 12 ++++++++++++ src/linear/linear_regression.rs | 12 ++++++++++++ src/linear/logistic_regression.rs | 13 +++++++++++++ src/linear/ridge_regression.rs | 12 ++++++++++++ src/naive_bayes/bernoulli.rs | 20 ++++++++++++++++---- src/naive_bayes/categorical.rs | 15 +++++++++++++-- src/naive_bayes/gaussian.rs | 17 ++++++++++++++--- src/naive_bayes/multinomial.rs | 17 ++++++++++++++--- src/neighbors/knn_classifier.rs | 18 +++++++++++++++--- src/neighbors/knn_regressor.rs | 18 +++++++++++++++--- src/svm/svc.rs | 11 ++++++++++- src/svm/svr.rs | 9 +++++++++ src/tree/base_tree_regressor.rs | 8 ++++++-- src/tree/decision_tree_classifier.rs | 10 +++++++--- src/tree/decision_tree_regressor.rs | 14 +++++++++++++- src/xgboost/xgb_regressor.rs | 9 +++++++++ 25 files changed, 294 insertions(+), 26 deletions(-) diff --git a/src/cluster/agglomerative.rs b/src/cluster/agglomerative.rs index 373f6f95..b9936b0e 100644 --- a/src/cluster/agglomerative.rs +++ b/src/cluster/agglomerative.rs @@ -76,6 +76,7 @@ impl Default for AgglomerativeClusteringParameters { pub struct AgglomerativeClustering, Y: Array1> { /// The cluster label assigned to each sample. pub labels: Vec, + parameters: Option, _phantom_tx: PhantomData, _phantom_ty: PhantomData, _phantom_x: PhantomData, @@ -176,12 +177,22 @@ impl, Y: Array1> AgglomerativeClusteri } Ok(AgglomerativeClustering { labels, + parameters: Some(parameters), _phantom_tx: PhantomData, _phantom_ty: PhantomData, _phantom_x: PhantomData, _phantom_y: PhantomData, }) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &AgglomerativeClusteringParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } impl, Y: Array1> diff --git a/src/cluster/dbscan.rs b/src/cluster/dbscan.rs index 2e2aac10..6bf13005 100644 --- a/src/cluster/dbscan.rs +++ b/src/cluster/dbscan.rs @@ -64,6 +64,7 @@ pub struct DBSCAN, Y: Array1, D: Dista num_classes: usize, knn_algorithm: KNNAlgorithm, eps: f64, + parameters: Option>, _phantom_ty: PhantomData, _phantom_x: PhantomData, _phantom_y: PhantomData, @@ -295,7 +296,7 @@ impl, Y: Array1, D: Distance>> x.row_iter() .map(|row| row.iterator(0).cloned().collect()) .collect(), - parameters.distance, + parameters.distance.clone(), )?; let mut row = vec![TX::zero(); x.shape().1]; @@ -353,6 +354,7 @@ impl, Y: Array1, D: Distance>> num_classes: k as usize, knn_algorithm: algo, eps: parameters.eps, + parameters: Some(parameters), _phantom_ty: PhantomData, _phantom_x: PhantomData, _phantom_y: PhantomData, @@ -392,6 +394,15 @@ impl, Y: Array1, D: Distance>> Ok(result) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &DBSCANParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/cluster/kmeans.rs b/src/cluster/kmeans.rs index b81ffd7e..f5f683f6 100644 --- a/src/cluster/kmeans.rs +++ b/src/cluster/kmeans.rs @@ -76,6 +76,7 @@ pub struct KMeans, Y: Array1> { size: Vec, _distortion: f64, centroids: Vec>, + parameters: Option, _phantom_tx: PhantomData, _phantom_ty: PhantomData, _phantom_x: PhantomData, @@ -315,6 +316,7 @@ impl, Y: Array1> KMeans size, _distortion: distortion, centroids, + parameters: Some(parameters), _phantom_tx: PhantomData, _phantom_ty: PhantomData, _phantom_x: PhantomData, @@ -411,6 +413,15 @@ impl, Y: Array1> KMeans y } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &KMeansParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/decomposition/pca.rs b/src/decomposition/pca.rs index 11853648..0dd2b62e 100644 --- a/src/decomposition/pca.rs +++ b/src/decomposition/pca.rs @@ -65,6 +65,7 @@ pub struct PCA + SVDDecomposable + EVDDe eigenvectors: X, eigenvalues: Vec, projection: X, + parameters: Option, mu: Vec, pmu: Vec, } @@ -329,6 +330,7 @@ impl + SVDDecomposable + EVDDecomposable eigenvectors, eigenvalues, projection: projection.transpose(), + parameters: Some(parameters), mu, pmu, }) @@ -360,6 +362,15 @@ impl + SVDDecomposable + EVDDecomposable pub fn components(&self) -> &X { &self.projection } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &PCAParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/decomposition/svd.rs b/src/decomposition/svd.rs index 259bfbc0..524b0f17 100644 --- a/src/decomposition/svd.rs +++ b/src/decomposition/svd.rs @@ -62,6 +62,7 @@ use crate::numbers::realnum::RealNumber; #[derive(Debug)] pub struct SVD + SVDDecomposable + EVDDecomposable> { components: X, + parameters: Option, phantom: PhantomData, } @@ -190,6 +191,7 @@ impl + SVDDecomposable + EVDDecomposable Ok(SVD { components, + parameters: Some(parameters), phantom: PhantomData, }) } @@ -212,6 +214,15 @@ impl + SVDDecomposable + EVDDecomposable pub fn components(&self) -> &X { &self.components } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &SVDParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index 818ac6c7..bb5a9efa 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -104,6 +104,7 @@ pub struct ExtraTreesRegressor< Y: Array1, > { forest_regressor: Option>, + parameters: Option } impl ExtraTreesRegressorParameters { @@ -165,6 +166,7 @@ impl, Y: Array1 fn new() -> Self { Self { forest_regressor: Option::None, + parameters: Option::None } } @@ -207,6 +209,7 @@ impl, Y: Array1 Ok(ExtraTreesRegressor { forest_regressor: Some(forest_regressor), + parameters: Some(parameters) }) } @@ -222,6 +225,15 @@ impl, Y: Array1 let forest_regressor = self.forest_regressor.as_ref().unwrap(); forest_regressor.predict_oob(x) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &ExtraTreesRegressorParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index 0f86a4df..1a8478a7 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -107,6 +107,7 @@ pub struct RandomForestClassifier< trees: Option>>, classes: Option>, samples: Option>>, + parameters: Option } impl RandomForestClassifierParameters { @@ -200,6 +201,7 @@ impl, Y: trees: Option::None, classes: Option::None, samples: Option::None, + parameters: Option::None, } } fn fit(x: &X, y: &Y, parameters: RandomForestClassifierParameters) -> Result { @@ -506,6 +508,7 @@ impl, Y: Array1, Y: Array1 &RandomForestClassifierParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 0a8a888c..023e0757 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -95,6 +95,7 @@ pub struct RandomForestRegressor< Y: Array1, > { forest_regressor: Option>, + parameters: Option } impl RandomForestRegressorParameters { @@ -165,6 +166,7 @@ impl, Y: Array1 fn new() -> Self { Self { forest_regressor: Option::None, + parameters: Option::None, } } @@ -399,6 +401,7 @@ impl, Y: Array1 Ok(RandomForestRegressor { forest_regressor: Some(forest_regressor), + parameters: Some(parameters) }) } @@ -414,6 +417,15 @@ impl, Y: Array1 let forest_regressor = self.forest_regressor.as_ref().unwrap(); forest_regressor.predict_oob(x) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &RandomForestRegressorParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/linear/elastic_net.rs b/src/linear/elastic_net.rs index d5b1d4d5..dae81843 100644 --- a/src/linear/elastic_net.rs +++ b/src/linear/elastic_net.rs @@ -98,6 +98,7 @@ pub struct ElasticNetParameters { pub struct ElasticNet, Y: Array1> { coefficients: Option, intercept: Option, + parameters: Option, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -288,6 +289,7 @@ impl, Y: Array1> Self { coefficients: Option::None, intercept: Option::None, + parameters: Option::None, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -385,6 +387,7 @@ impl, Y: Array1> Ok(ElasticNet { intercept: Some(b), coefficients: Some(w), + parameters: Some(parameters), _phantom_ty: PhantomData, _phantom_y: PhantomData, }) @@ -459,6 +462,15 @@ impl, Y: Array1> (x2, y2, gamma) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &ElasticNetParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/linear/lasso.rs b/src/linear/lasso.rs index 59c60ddc..cbe922f7 100644 --- a/src/linear/lasso.rs +++ b/src/linear/lasso.rs @@ -64,6 +64,7 @@ pub struct LassoParameters { pub struct Lasso, Y: Array1> { coefficients: Option, intercept: Option, + parameters: Option, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -130,6 +131,7 @@ impl, Y: Array1> Self { coefficients: None, intercept: None, + parameters: None, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -352,6 +354,7 @@ impl, Y: Array1> Las Ok(Lasso { intercept: b, coefficients: Some(w), + parameters: Some(parameters), _phantom_ty: PhantomData, _phantom_y: PhantomData, }) @@ -402,6 +405,15 @@ impl, Y: Array1> Las scaled_x.scale_mut(&col_mean, &col_std, 0); Ok((scaled_x, col_mean, col_std)) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &LassoParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/linear/linear_regression.rs b/src/linear/linear_regression.rs index 43410bbb..f65c9434 100644 --- a/src/linear/linear_regression.rs +++ b/src/linear/linear_regression.rs @@ -113,6 +113,7 @@ pub struct LinearRegression< > { coefficients: Option, intercept: Option, + parameters: Option, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -209,6 +210,7 @@ impl< Self { coefficients: Option::None, intercept: Option::None, + parameters: Option::None, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -274,6 +276,7 @@ impl< Ok(LinearRegression { intercept: Some(*w.get((num_attributes, 0))), coefficients: Some(weights), + parameters: Some(parameters), _phantom_ty: PhantomData, _phantom_y: PhantomData, }) @@ -301,6 +304,15 @@ impl< pub fn intercept(&self) -> &TX { self.intercept.as_ref().unwrap() } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &LinearRegressionParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/linear/logistic_regression.rs b/src/linear/logistic_regression.rs index c28dc347..1d4a7177 100644 --- a/src/linear/logistic_regression.rs +++ b/src/linear/logistic_regression.rs @@ -178,6 +178,7 @@ pub struct LogisticRegression< classes: Option>, num_attributes: usize, num_classes: usize, + parameters: Option>, _phantom_tx: PhantomData, _phantom_y: PhantomData, } @@ -389,6 +390,7 @@ impl, Y: classes: Option::None, num_attributes: 0, num_classes: 0, + parameters: Option::None, _phantom_tx: PhantomData, _phantom_y: PhantomData, } @@ -465,6 +467,7 @@ impl, Y: classes: Some(classes), num_attributes, num_classes: k, + parameters: Some(parameters), _phantom_tx: PhantomData, _phantom_y: PhantomData, }) @@ -491,6 +494,7 @@ impl, Y: classes: Some(classes), num_attributes, num_classes: k, + parameters: Some(parameters), _phantom_tx: PhantomData, _phantom_y: PhantomData, }) @@ -560,6 +564,15 @@ impl, Y: optimizer.optimize(&f, &df, &x0, &ls) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &LogisticRegressionParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/linear/ridge_regression.rs b/src/linear/ridge_regression.rs index be2f3d41..e03f476a 100644 --- a/src/linear/ridge_regression.rs +++ b/src/linear/ridge_regression.rs @@ -192,6 +192,7 @@ pub struct RidgeRegression< > { coefficients: Option, intercept: Option, + parameters: Option>, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -253,6 +254,7 @@ impl< Self { coefficients: Option::None, intercept: Option::None, + parameters: Option::None, _phantom_ty: PhantomData, _phantom_y: PhantomData, } @@ -360,6 +362,7 @@ impl< Ok(RidgeRegression { intercept: Some(b), coefficients: Some(w), + parameters: Some(parameters), _phantom_ty: PhantomData, _phantom_y: PhantomData, }) @@ -409,6 +412,15 @@ impl< pub fn intercept(&self) -> &TX { self.intercept.as_ref().unwrap() } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &RidgeRegressionParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/naive_bayes/bernoulli.rs b/src/naive_bayes/bernoulli.rs index cdd5b83d..1e0bcd2c 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -127,7 +127,7 @@ impl NBDistribution /// `BernoulliNB` parameters. Use `Default::default()` for default values. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, PartialEq)] pub struct BernoulliNBParameters { #[cfg_attr(feature = "serde", serde(default))] /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). @@ -357,6 +357,7 @@ pub struct BernoulliNB< > { inner: Option>>, binarize: Option, + parameters: Option> } impl, Y: Array1> @@ -380,6 +381,7 @@ impl, Y: Arr Self { inner: Option::None, binarize: Option::None, + parameters: Option::None } } @@ -410,17 +412,18 @@ impl, Y: Arr BernoulliNBDistribution::fit( &Self::binarize(x, threshold), y, - parameters.alpha, - parameters.priors, + parameters.alpha.clone(), + parameters.priors.clone(), )? } else { - BernoulliNBDistribution::fit(x, y, parameters.alpha, parameters.priors)? + BernoulliNBDistribution::fit(x, y, parameters.alpha.clone(), parameters.priors.clone())? }; let inner = BaseNaiveBayes::fit(distribution)?; Ok(Self { inner: Some(inner), binarize: parameters.binarize, + parameters: Some(parameters) }) } @@ -485,6 +488,15 @@ impl, Y: Arr Self::binarize_mut(&mut new_x, threshold); new_x } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &BernoulliNBParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/naive_bayes/categorical.rs b/src/naive_bayes/categorical.rs index b60ee0d3..80c27fa7 100644 --- a/src/naive_bayes/categorical.rs +++ b/src/naive_bayes/categorical.rs @@ -257,7 +257,7 @@ impl CategoricalNBDistribution { /// `CategoricalNB` parameters. Use `Default::default()` for default values. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, PartialEq)] pub struct CategoricalNBParameters { #[cfg_attr(feature = "serde", serde(default))] /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). @@ -338,6 +338,7 @@ impl Default for CategoricalNBSearchParameters { #[derive(Debug, PartialEq)] pub struct CategoricalNB, Y: Array1> { inner: Option>>, + parameters: Option } impl, Y: Array1> @@ -346,6 +347,7 @@ impl, Y: Array1> fn new() -> Self { Self { inner: Option::None, + parameters: Option::None } } @@ -370,7 +372,7 @@ impl, Y: Array1> CategoricalNB { let alpha = parameters.alpha; let distribution = CategoricalNBDistribution::fit(x, y, alpha)?; let inner = BaseNaiveBayes::fit(distribution)?; - Ok(Self { inner: Some(inner) }) + Ok(Self { inner: Some(inner), parameters: Some(parameters) }) } /// Estimates the class labels for the provided data. @@ -415,6 +417,15 @@ impl, Y: Array1> CategoricalNB { pub fn feature_log_prob(&self) -> &Vec>> { &self.inner.as_ref().unwrap().distribution.coefficients } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &CategoricalNBParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/naive_bayes/gaussian.rs b/src/naive_bayes/gaussian.rs index dbf3fd81..3b93fa2a 100644 --- a/src/naive_bayes/gaussian.rs +++ b/src/naive_bayes/gaussian.rs @@ -92,7 +92,7 @@ impl NBDistribution /// `GaussianNB` parameters. Use `Default::default()` for default values. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Default, Clone)] +#[derive(Debug, Default, Clone, PartialEq)] pub struct GaussianNBParameters { #[cfg_attr(feature = "serde", serde(default))] /// Prior probabilities of the classes. If specified the priors are not adjusted according to the data @@ -266,6 +266,7 @@ pub struct GaussianNB< Y: Array1, > { inner: Option>>, + parameters: Option } impl< @@ -291,6 +292,7 @@ impl< fn new() -> Self { Self { inner: Option::None, + parameters: Option::None, } } @@ -320,9 +322,9 @@ impl, Y: Arr /// * `y` - vector with target values (classes) of length N. /// * `parameters` - additional parameters like class priors. pub fn fit(x: &X, y: &Y, parameters: GaussianNBParameters) -> Result { - let distribution = GaussianNBDistribution::fit(x, y, parameters.priors)?; + let distribution = GaussianNBDistribution::fit(x, y, parameters.priors.clone())?; let inner = BaseNaiveBayes::fit(distribution)?; - Ok(Self { inner: Some(inner) }) + Ok(Self { inner: Some(inner), parameters: Some(parameters) }) } /// Estimates the class labels for the provided data. @@ -362,6 +364,15 @@ impl, Y: Arr pub fn var(&self) -> &Vec> { &self.inner.as_ref().unwrap().distribution.var } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &GaussianNBParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/naive_bayes/multinomial.rs b/src/naive_bayes/multinomial.rs index ad873943..58d043e5 100644 --- a/src/naive_bayes/multinomial.rs +++ b/src/naive_bayes/multinomial.rs @@ -99,7 +99,7 @@ impl NBDistribution /// `MultinomialNB` parameters. Use `Default::default()` for default values. #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[derive(Debug, Clone, PartialEq)] pub struct MultinomialNBParameters { #[cfg_attr(feature = "serde", serde(default))] /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). @@ -302,6 +302,7 @@ pub struct MultinomialNB< Y: Array1, > { inner: Option>>, + parameters: Option } impl, Y: Array1> fmt::Display @@ -323,6 +324,7 @@ impl, Y: Array fn new() -> Self { Self { inner: Option::None, + parameters: Option::None } } @@ -350,9 +352,9 @@ impl, Y: Array /// binarizing threshold. pub fn fit(x: &X, y: &Y, parameters: MultinomialNBParameters) -> Result { let distribution = - MultinomialNBDistribution::fit(x, y, parameters.alpha, parameters.priors)?; + MultinomialNBDistribution::fit(x, y, parameters.alpha.clone(), parameters.priors.clone())?; let inner = BaseNaiveBayes::fit(distribution)?; - Ok(Self { inner: Some(inner) }) + Ok(Self { inner: Some(inner), parameters: Some(parameters) }) } /// Estimates the class labels for the provided data. @@ -391,6 +393,15 @@ impl, Y: Array pub fn feature_count(&self) -> &Vec> { &self.inner.as_ref().unwrap().distribution.feature_count } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &MultinomialNBParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/neighbors/knn_classifier.rs b/src/neighbors/knn_classifier.rs index 264ab0e2..2cfa781e 100644 --- a/src/neighbors/knn_classifier.rs +++ b/src/neighbors/knn_classifier.rs @@ -83,6 +83,7 @@ pub struct KNNClassifier< knn_algorithm: Option>, weight: Option, k: Option, + parameters: Option>, _phantom_tx: PhantomData, _phantom_x: PhantomData, _phantom_y: PhantomData, @@ -188,6 +189,7 @@ impl, Y: Array1, D: Distance, Y: Array1, D: Distance, Y: Array1, D: Distance &KNNClassifierParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/neighbors/knn_regressor.rs b/src/neighbors/knn_regressor.rs index b49743f8..a97c6684 100644 --- a/src/neighbors/knn_regressor.rs +++ b/src/neighbors/knn_regressor.rs @@ -80,6 +80,7 @@ pub struct KNNRegressor, Y: Array1, D: knn_algorithm: Option>, weight: Option, k: Option, + parameters: Option>, _phantom_tx: PhantomData, _phantom_ty: PhantomData, _phantom_x: PhantomData, @@ -179,6 +180,7 @@ impl, Y: Array1, D: Distance>> knn_algorithm: Option::None, weight: Option::None, k: Option::None, + parameters: Option::None, _phantom_tx: PhantomData, _phantom_ty: PhantomData, _phantom_x: PhantomData, @@ -231,13 +233,14 @@ impl, Y: Array1, D: Distance>> ))); } - let knn_algo = parameters.algorithm.fit(data, parameters.distance)?; + let knn_algo = parameters.algorithm.fit(data, parameters.distance.clone())?; Ok(KNNRegressor { y: Some(y.clone()), - k: Some(parameters.k), + k: Some(parameters.k.clone()), knn_algorithm: Some(knn_algo), - weight: Some(parameters.weight), + weight: Some(parameters.weight.clone()), + parameters: Some(parameters), _phantom_tx: PhantomData, _phantom_ty: PhantomData, _phantom_x: PhantomData, @@ -277,6 +280,15 @@ impl, Y: Array1, D: Distance>> Ok(result) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &KNNRegressorParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/svm/svc.rs b/src/svm/svc.rs index d72ecdac..692d3ef7 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -292,7 +292,7 @@ pub struct SVCParameters, /// Controls the pseudo random number generation for shuffling the data for probability estimates - seed: Option, + pub seed: Option, } #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] @@ -608,6 +608,15 @@ impl<'a, TX: Number + RealNumber, TY: Number + Ord, X: Array2 + 'a, Y: Array f } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &SVCParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } impl, Y: Array1> PartialEq diff --git a/src/svm/svr.rs b/src/svm/svr.rs index e912743b..679a77a3 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -280,6 +280,15 @@ impl<'a, T: Number + FloatNumber + PartialOrd, X: Array2, Y: Array1> SVR<' T::from(f).unwrap() } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &SVRParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } impl, Y: Array1> PartialEq diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index f84ae7e9..2230e8a1 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -63,9 +63,13 @@ impl, Y: Array1> fn nodes(&self) -> &Vec { self.nodes.as_ref() } - /// Get parameters, return a shared reference + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model fn parameters(&self) -> &BaseTreeRegressorParameters { - self.parameters.as_ref().unwrap() + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() } /// Get estimate of intercept, return value fn depth(&self) -> u16 { diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index 96007677..2a2e7859 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -131,9 +131,13 @@ impl, Y: Array1> fn nodes(&self) -> &Vec { self.nodes.as_ref() } - /// Get parameters, return a shared reference - fn parameters(&self) -> &DecisionTreeClassifierParameters { - self.parameters.as_ref().unwrap() + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &DecisionTreeClassifierParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() } /// get classes vector, return a shared reference fn classes(&self) -> &Vec { diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index 86b99343..ebfc1d8f 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -94,6 +94,7 @@ pub struct DecisionTreeRegressorParameters { pub struct DecisionTreeRegressor, Y: Array1> { tree_regressor: Option>, + parameters: Option } impl DecisionTreeRegressorParameters { @@ -271,7 +272,8 @@ impl, Y: Array1> { fn new() -> Self { Self { - tree_regressor: None, + tree_regressor: Option::None, + parameters: Option::None } } @@ -309,6 +311,7 @@ impl, Y: Array1> let tree = BaseTreeRegressor::fit(x, y, tree_parameters)?; Ok(Self { tree_regressor: Some(tree), + parameters: Some(parameters) }) } @@ -317,6 +320,15 @@ impl, Y: Array1> pub fn predict(&self, x: &X) -> Result { self.tree_regressor.as_ref().unwrap().predict(x) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &DecisionTreeRegressorParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } #[cfg(test)] diff --git a/src/xgboost/xgb_regressor.rs b/src/xgboost/xgb_regressor.rs index a3b9bf0a..92a3f4f0 100644 --- a/src/xgboost/xgb_regressor.rs +++ b/src/xgboost/xgb_regressor.rs @@ -609,6 +609,15 @@ impl, Y: Array1> XGRegres indices.truncate((population_size as f64 * subsample_ratio) as usize); indices } + + /// Getter for parameters used in the model + /// + /// # Returns + /// Parameters used to setup the model + pub fn parameters(&self) -> &XGRegressorParameters { + assert!(self.parameters.is_some()); + &self.parameters.as_ref().unwrap() + } } // Boilerplate implementation for the smartcore traits From 4894d3cf2e7099eb35efced95d53ac90471ccbc8 Mon Sep 17 00:00:00 2001 From: Hubert Majewski <36614010+HubertMajewski@users.noreply.github.com> Date: Tue, 21 Jul 2026 05:50:05 +0000 Subject: [PATCH 2/4] Convert parameters() to return an Option; add tests with parameters --- src/cluster/agglomerative.rs | 38 ++++++++++-- src/cluster/dbscan.rs | 53 +++++++++++++++-- src/cluster/kmeans.rs | 42 +++++++++++-- src/decomposition/pca.rs | 36 ++++++++++-- src/decomposition/svd.rs | 32 ++++++++-- src/ensemble/extra_trees_regressor.rs | 52 ++++++++++++++-- src/ensemble/random_forest_classifier.rs | 56 ++++++++++++++++-- src/ensemble/random_forest_regressor.rs | 52 ++++++++++++++-- src/linear/elastic_net.rs | 48 +++++++++++++-- src/linear/lasso.rs | 48 +++++++++++++-- src/linear/linear_regression.rs | 44 ++++++++++++-- src/linear/logistic_regression.rs | 47 +++++++++++++-- src/linear/ridge_regression.rs | 46 +++++++++++++-- src/naive_bayes/bernoulli.rs | 44 ++++++++++++-- src/naive_bayes/categorical.rs | 44 ++++++++++++-- src/naive_bayes/gaussian.rs | 44 ++++++++++++-- src/naive_bayes/multinomial.rs | 44 ++++++++++++-- src/neighbors/knn_classifier.rs | 60 +++++++++++++++++-- src/neighbors/knn_regressor.rs | 60 +++++++++++++++++-- src/svm/svc.rs | 44 ++++++++++++-- src/svm/svr.rs | 40 +++++++++++-- src/tree/base_tree_regressor.rs | 31 ++++++---- src/tree/decision_tree_classifier.rs | 75 ++++++++++++++++++++---- src/tree/decision_tree_regressor.rs | 47 +++++++++++++-- src/xgboost/xgb_regressor.rs | 56 ++++++++++++++++-- 25 files changed, 1067 insertions(+), 116 deletions(-) diff --git a/src/cluster/agglomerative.rs b/src/cluster/agglomerative.rs index b9936b0e..88f468e6 100644 --- a/src/cluster/agglomerative.rs +++ b/src/cluster/agglomerative.rs @@ -188,10 +188,9 @@ impl, Y: Array1> AgglomerativeClusteri /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &AgglomerativeClusteringParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&AgglomerativeClusteringParameters> { + self.parameters.as_ref() } } @@ -325,4 +324,35 @@ mod tests { assert!(result.is_err()); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 5.0, 5.0, 10.0, 10.0]; + let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); + let parameters = AgglomerativeClusteringParameters::default().with_n_clusters(1); + let expected_parameters = parameters.clone(); + let clustering = AgglomerativeClustering::, Vec>::fit( + &matrix, parameters, + ) + .unwrap(); + + let actual_parameters = clustering + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.n_clusters, expected_parameters.n_clusters); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let clustering = AgglomerativeClustering::, Vec> { + labels: Vec::new(), + parameters: None, + _phantom_tx: PhantomData, + _phantom_ty: PhantomData, + _phantom_x: PhantomData, + _phantom_y: PhantomData, + }; + + assert!(clustering.parameters().is_none()); + } } diff --git a/src/cluster/dbscan.rs b/src/cluster/dbscan.rs index 6bf13005..5a8c8fef 100644 --- a/src/cluster/dbscan.rs +++ b/src/cluster/dbscan.rs @@ -398,10 +398,9 @@ impl, Y: Array1, D: Distance>> /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &DBSCANParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&DBSCANParameters> { + self.parameters.as_ref() } } @@ -525,4 +524,50 @@ mod tests { println!("{labels:?}"); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.5, 0.5, 10.0, 10.0]; + let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); + let parameters: DBSCANParameters> = + DBSCANParameters::default().with_eps(1.0).with_min_samples(1); + let expected_parameters = parameters.clone(); + let clustering = + DBSCAN::, Vec, Euclidian>::fit( + &matrix, + parameters, + ) + .unwrap(); + + let actual_parameters = clustering + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!( + format!("{:?}", actual_parameters.distance), + format!("{:?}", expected_parameters.distance) + ); + assert_eq!(actual_parameters.min_samples, expected_parameters.min_samples); + assert_eq!(actual_parameters.eps, expected_parameters.eps); + assert_eq!( + std::mem::discriminant(&actual_parameters.algorithm), + std::mem::discriminant(&expected_parameters.algorithm) + ); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.5, 0.5, 10.0, 10.0]; + let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); + let parameters: DBSCANParameters> = + DBSCANParameters::default().with_eps(1.0).with_min_samples(1); + let mut clustering = + DBSCAN::, Vec, Euclidian>::fit( + &matrix, + parameters, + ) + .unwrap(); + clustering.parameters = None; + + assert!(clustering.parameters().is_none()); + } } diff --git a/src/cluster/kmeans.rs b/src/cluster/kmeans.rs index f5f683f6..6b6c7563 100644 --- a/src/cluster/kmeans.rs +++ b/src/cluster/kmeans.rs @@ -417,10 +417,9 @@ impl, Y: Array1> KMeans /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &KMeansParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&KMeansParameters> { + self.parameters.as_ref() } } @@ -554,4 +553,39 @@ mod tests { assert_eq!(kmeans, deserialized_kmeans); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 5.0, 5.0, 10.0, 10.0]; + let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); + let parameters = KMeansParameters::default().with_k(2); + let expected_parameters = parameters.clone(); + let clustering = KMeans::, Vec>::fit( + &matrix, + parameters, + ) + .unwrap(); + + let actual_parameters = clustering + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.k, expected_parameters.k); + assert_eq!(actual_parameters.max_iter, expected_parameters.max_iter); + assert_eq!(actual_parameters.seed, expected_parameters.seed); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 5.0, 5.0, 10.0, 10.0]; + let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); + let parameters = KMeansParameters::default().with_k(2); + let mut clustering = KMeans::, Vec>::fit( + &matrix, + parameters, + ) + .unwrap(); + clustering.parameters = None; + + assert!(clustering.parameters().is_none()); + } } diff --git a/src/decomposition/pca.rs b/src/decomposition/pca.rs index 0dd2b62e..8189a730 100644 --- a/src/decomposition/pca.rs +++ b/src/decomposition/pca.rs @@ -366,10 +366,9 @@ impl + SVDDecomposable + EVDDecomposable /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &PCAParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&PCAParameters> { + self.parameters.as_ref() } } @@ -758,4 +757,33 @@ mod tests { // assert_eq!(pca, deserialized_pca); // } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let parameters = PCAParameters::default().with_n_components(1); + let expected_parameters = parameters.clone(); + let pca = PCA::>::fit(&matrix, parameters).unwrap(); + + let actual_parameters = pca + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.n_components, expected_parameters.n_components); + assert_eq!( + actual_parameters.use_correlation_matrix, + expected_parameters.use_correlation_matrix + ); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let parameters = PCAParameters::default().with_n_components(1); + let mut pca = PCA::>::fit(&matrix, parameters).unwrap(); + pca.parameters = None; + + assert!(pca.parameters().is_none()); + } } diff --git a/src/decomposition/svd.rs b/src/decomposition/svd.rs index 524b0f17..96fd1948 100644 --- a/src/decomposition/svd.rs +++ b/src/decomposition/svd.rs @@ -218,10 +218,9 @@ impl + SVDDecomposable + EVDDecomposable /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &SVDParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&SVDParameters> { + self.parameters.as_ref() } } @@ -363,4 +362,29 @@ mod tests { // assert_eq!(svd, deserialized_svd); // } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 0.0, 1.0]; + let matrix = DenseMatrix::new(3, 3, data, false).unwrap(); + let parameters = SVDParameters::default().with_n_components(1); + let expected_parameters = parameters.clone(); + let svd = SVD::>::fit(&matrix, parameters).unwrap(); + + let actual_parameters = svd + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.n_components, expected_parameters.n_components); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 0.0, 1.0]; + let matrix = DenseMatrix::new(3, 3, data, false).unwrap(); + let parameters = SVDParameters::default().with_n_components(1); + let mut svd = SVD::>::fit(&matrix, parameters).unwrap(); + svd.parameters = None; + + assert!(svd.parameters().is_none()); + } } diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index bb5a9efa..9d94afd9 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -229,10 +229,9 @@ impl, Y: Array1 /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &ExtraTreesRegressorParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&ExtraTreesRegressorParameters> { + self.parameters.as_ref() } } @@ -327,4 +326,49 @@ mod tests { assert_eq!(y_hat1, y_hat2); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = ExtraTreesRegressorParameters::default(); + let expected_parameters = parameters.clone(); + let regressor = + ExtraTreesRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regressor + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); + assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); + assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!(actual_parameters.n_trees, expected_parameters.n_trees); + assert_eq!(actual_parameters.m, expected_parameters.m); + assert_eq!(actual_parameters.keep_samples, expected_parameters.keep_samples); + assert_eq!(actual_parameters.seed, expected_parameters.seed); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = ExtraTreesRegressorParameters::default(); + let mut regressor = + ExtraTreesRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regressor.parameters = None; + + assert!(regressor.parameters().is_none()); + } } diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index 1a8478a7..a457eaeb 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -620,10 +620,9 @@ impl, Y: Array1 &RandomForestClassifierParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&RandomForestClassifierParameters> { + self.parameters.as_ref() } } @@ -822,4 +821,53 @@ mod tests { assert_eq!(forest, deserialized_forest); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters = RandomForestClassifierParameters::default(); + let expected_parameters = parameters.clone(); + let classifier = + RandomForestClassifier::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!( + std::mem::discriminant(&actual_parameters.criterion), + std::mem::discriminant(&expected_parameters.criterion) + ); + assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); + assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); + assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!(actual_parameters.n_trees, expected_parameters.n_trees); + assert_eq!(actual_parameters.m, expected_parameters.m); + assert_eq!(actual_parameters.keep_samples, expected_parameters.keep_samples); + assert_eq!(actual_parameters.seed, expected_parameters.seed); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters = RandomForestClassifierParameters::default(); + let mut classifier = + RandomForestClassifier::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 023e0757..b7bcb508 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -421,10 +421,9 @@ impl, Y: Array1 /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &RandomForestRegressorParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&RandomForestRegressorParameters> { + self.parameters.as_ref() } } @@ -624,4 +623,49 @@ mod tests { assert_eq!(forest, deserialized_forest); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = RandomForestRegressorParameters::default(); + let expected_parameters = parameters.clone(); + let regressor = + RandomForestRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regressor + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); + assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); + assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!(actual_parameters.n_trees, expected_parameters.n_trees); + assert_eq!(actual_parameters.m, expected_parameters.m); + assert_eq!(actual_parameters.keep_samples, expected_parameters.keep_samples); + assert_eq!(actual_parameters.seed, expected_parameters.seed); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = RandomForestRegressorParameters::default(); + let mut regressor = + RandomForestRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regressor.parameters = None; + + assert!(regressor.parameters().is_none()); + } } diff --git a/src/linear/elastic_net.rs b/src/linear/elastic_net.rs index dae81843..0cafb9c1 100644 --- a/src/linear/elastic_net.rs +++ b/src/linear/elastic_net.rs @@ -466,10 +466,9 @@ impl, Y: Array1> /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &ElasticNetParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&ElasticNetParameters> { + self.parameters.as_ref() } } @@ -657,4 +656,45 @@ mod tests { // assert_eq!(lr, deserialized_lr); // } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = ElasticNetParameters::default(); + let expected_parameters = parameters.clone(); + let regression = ElasticNet::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regression + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.alpha, expected_parameters.alpha); + assert_eq!(actual_parameters.l1_ratio, expected_parameters.l1_ratio); + assert_eq!(actual_parameters.normalize, expected_parameters.normalize); + assert_eq!(actual_parameters.tol, expected_parameters.tol); + assert_eq!(actual_parameters.max_iter, expected_parameters.max_iter); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = ElasticNetParameters::default(); + let mut regression = ElasticNet::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regression.parameters = None; + + assert!(regression.parameters().is_none()); + } } diff --git a/src/linear/lasso.rs b/src/linear/lasso.rs index cbe922f7..81092cce 100644 --- a/src/linear/lasso.rs +++ b/src/linear/lasso.rs @@ -409,10 +409,9 @@ impl, Y: Array1> Las /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &LassoParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&LassoParameters> { + self.parameters.as_ref() } } @@ -586,4 +585,45 @@ mod tests { // assert_eq!(lr, deserialized_lr); // } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = LassoParameters::default(); + let expected_parameters = parameters.clone(); + let regression = Lasso::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regression + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.alpha, expected_parameters.alpha); + assert_eq!(actual_parameters.normalize, expected_parameters.normalize); + assert_eq!(actual_parameters.tol, expected_parameters.tol); + assert_eq!(actual_parameters.max_iter, expected_parameters.max_iter); + assert_eq!(actual_parameters.fit_intercept, expected_parameters.fit_intercept); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = LassoParameters::default(); + let mut regression = Lasso::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regression.parameters = None; + + assert!(regression.parameters().is_none()); + } } diff --git a/src/linear/linear_regression.rs b/src/linear/linear_regression.rs index f65c9434..ea97be3e 100644 --- a/src/linear/linear_regression.rs +++ b/src/linear/linear_regression.rs @@ -308,10 +308,9 @@ impl< /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &LinearRegressionParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&LinearRegressionParameters> { + self.parameters.as_ref() } } @@ -424,4 +423,41 @@ mod tests { // let parameters: LinearRegressionParameters = serde_json::from_str("{}").unwrap(); // assert_eq!(parameters.solver, default.solver); // } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = LinearRegressionParameters::default(); + let expected_parameters = parameters.clone(); + let regression = LinearRegression::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regression + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.solver, expected_parameters.solver); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = LinearRegressionParameters::default(); + let mut regression = LinearRegression::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regression.parameters = None; + + assert!(regression.parameters().is_none()); + } } diff --git a/src/linear/logistic_regression.rs b/src/linear/logistic_regression.rs index 1d4a7177..e2170768 100644 --- a/src/linear/logistic_regression.rs +++ b/src/linear/logistic_regression.rs @@ -568,10 +568,9 @@ impl, Y: /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &LogisticRegressionParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&LogisticRegressionParameters> { + self.parameters.as_ref() } } @@ -959,4 +958,44 @@ mod tests { assert_eq!(y_hat.shape(), 52181); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters: LogisticRegressionParameters = + LogisticRegressionParameters::default(); + let expected_parameters = parameters.clone(); + let regression = LogisticRegression::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regression + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.solver, expected_parameters.solver); + assert_eq!(actual_parameters.alpha, expected_parameters.alpha); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters: LogisticRegressionParameters = + LogisticRegressionParameters::default(); + let mut regression = LogisticRegression::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regression.parameters = None; + + assert!(regression.parameters().is_none()); + } } diff --git a/src/linear/ridge_regression.rs b/src/linear/ridge_regression.rs index e03f476a..10c28439 100644 --- a/src/linear/ridge_regression.rs +++ b/src/linear/ridge_regression.rs @@ -416,10 +416,9 @@ impl< /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &RidgeRegressionParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&RidgeRegressionParameters> { + self.parameters.as_ref() } } @@ -540,4 +539,43 @@ mod tests { // assert_eq!(lr, deserialized_lr); // } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters: RidgeRegressionParameters = RidgeRegressionParameters::default(); + let expected_parameters = parameters.clone(); + let regression = RidgeRegression::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regression + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.solver, expected_parameters.solver); + assert_eq!(actual_parameters.alpha, expected_parameters.alpha); + assert_eq!(actual_parameters.normalize, expected_parameters.normalize); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters: RidgeRegressionParameters = RidgeRegressionParameters::default(); + let mut regression = RidgeRegression::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regression.parameters = None; + + assert!(regression.parameters().is_none()); + } } diff --git a/src/naive_bayes/bernoulli.rs b/src/naive_bayes/bernoulli.rs index 1e0bcd2c..0aef35b0 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -492,10 +492,9 @@ impl, Y: Arr /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &BernoulliNBParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&BernoulliNBParameters> { + self.parameters.as_ref() } } @@ -667,4 +666,41 @@ mod tests { assert_eq!(bnb, deserialized_bnb); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0, 0, 0, 1, 1, 0, 1, 1]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters: BernoulliNBParameters = BernoulliNBParameters::default(); + let expected_parameters = parameters.clone(); + let classifier = BernoulliNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters, &expected_parameters); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0, 0, 0, 1, 1, 0, 1, 1]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters: BernoulliNBParameters = BernoulliNBParameters::default(); + let mut classifier = BernoulliNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/naive_bayes/categorical.rs b/src/naive_bayes/categorical.rs index 80c27fa7..020a0391 100644 --- a/src/naive_bayes/categorical.rs +++ b/src/naive_bayes/categorical.rs @@ -421,10 +421,9 @@ impl, Y: Array1> CategoricalNB { /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &CategoricalNBParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&CategoricalNBParameters> { + self.parameters.as_ref() } } @@ -596,4 +595,41 @@ mod tests { assert_eq!(cnb, deserialized_cnb); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0_u32, 0, 0, 1, 1, 0, 1, 1]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters = CategoricalNBParameters::default(); + let expected_parameters = parameters.clone(); + let classifier = CategoricalNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters, &expected_parameters); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0_u32, 0, 0, 1, 1, 0, 1, 1]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters = CategoricalNBParameters::default(); + let mut classifier = CategoricalNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/naive_bayes/gaussian.rs b/src/naive_bayes/gaussian.rs index 3b93fa2a..d41028fc 100644 --- a/src/naive_bayes/gaussian.rs +++ b/src/naive_bayes/gaussian.rs @@ -368,10 +368,9 @@ impl, Y: Arr /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &GaussianNBParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&GaussianNBParameters> { + self.parameters.as_ref() } } @@ -485,4 +484,41 @@ mod tests { assert_eq!(gnb, deserialized_gnb); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters = GaussianNBParameters::default(); + let expected_parameters = parameters.clone(); + let classifier = GaussianNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters, &expected_parameters); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters = GaussianNBParameters::default(); + let mut classifier = GaussianNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/naive_bayes/multinomial.rs b/src/naive_bayes/multinomial.rs index 58d043e5..bcd89ac3 100644 --- a/src/naive_bayes/multinomial.rs +++ b/src/naive_bayes/multinomial.rs @@ -397,10 +397,9 @@ impl, Y: Array /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &MultinomialNBParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&MultinomialNBParameters> { + self.parameters.as_ref() } } @@ -577,4 +576,41 @@ mod tests { assert_eq!(mnb, deserialized_mnb); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0_u32, 1, 1, 0, 1, 2, 2, 1]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters = MultinomialNBParameters::default(); + let expected_parameters = parameters.clone(); + let classifier = MultinomialNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters, &expected_parameters); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0_u32, 1, 1, 0, 1, 2, 2, 1]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0_u32, 0, 1, 1]; + let parameters = MultinomialNBParameters::default(); + let mut classifier = MultinomialNB::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/neighbors/knn_classifier.rs b/src/neighbors/knn_classifier.rs index 2cfa781e..cd6e04cc 100644 --- a/src/neighbors/knn_classifier.rs +++ b/src/neighbors/knn_classifier.rs @@ -351,10 +351,9 @@ impl, Y: Array1, D: Distance &KNNClassifierParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&KNNClassifierParameters> { + self.parameters.as_ref() } } @@ -737,4 +736,57 @@ mod tests { assert_eq!(knn, deserialized_knn); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters: KNNClassifierParameters> = + KNNClassifierParameters::default(); + let expected_parameters = parameters.clone(); + let classifier = + KNNClassifier::, Vec, Euclidian>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!( + format!("{:?}", actual_parameters.distance), + format!("{:?}", expected_parameters.distance) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.algorithm), + std::mem::discriminant(&expected_parameters.algorithm) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.weight), + std::mem::discriminant(&expected_parameters.weight) + ); + assert_eq!(actual_parameters.k, expected_parameters.k); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters: KNNClassifierParameters> = + KNNClassifierParameters::default(); + let mut classifier = + KNNClassifier::, Vec, Euclidian>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/neighbors/knn_regressor.rs b/src/neighbors/knn_regressor.rs index a97c6684..da0e57dd 100644 --- a/src/neighbors/knn_regressor.rs +++ b/src/neighbors/knn_regressor.rs @@ -284,10 +284,9 @@ impl, Y: Array1, D: Distance>> /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &KNNRegressorParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&KNNRegressorParameters> { + self.parameters.as_ref() } } @@ -362,4 +361,57 @@ mod tests { assert_eq!(knn, deserialized_knn); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters: KNNRegressorParameters> = + KNNRegressorParameters::default(); + let expected_parameters = parameters.clone(); + let regressor = + KNNRegressor::, Vec, Euclidian>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regressor + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!( + format!("{:?}", actual_parameters.distance), + format!("{:?}", expected_parameters.distance) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.algorithm), + std::mem::discriminant(&expected_parameters.algorithm) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.weight), + std::mem::discriminant(&expected_parameters.weight) + ); + assert_eq!(actual_parameters.k, expected_parameters.k); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters: KNNRegressorParameters> = + KNNRegressorParameters::default(); + let mut regressor = + KNNRegressor::, Vec, Euclidian>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regressor.parameters = None; + + assert!(regressor.parameters().is_none()); + } } diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 692d3ef7..116dae6c 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -612,10 +612,9 @@ impl<'a, TX: Number + RealNumber, TY: Number + Ord, X: Array2 + 'a, Y: Array /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &SVCParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&SVCParameters> { + self.parameters } } @@ -1438,4 +1437,41 @@ mod tests { assert_eq!(svc, deserialized_svc); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![-1, -1, 1, 1]; + let parameters: SVCParameters, Vec> = + SVCParameters::default().with_kernel(Kernels::linear()); + let classifier: SVC<'_, f64, i32, DenseMatrix, Vec> = + SVC::fit(&matrix, &target, ¶meters).unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.epoch, parameters.epoch); + assert_eq!(actual_parameters.c, parameters.c); + assert_eq!(actual_parameters.tol, parameters.tol); + assert_eq!( + format!("{:?}", actual_parameters.kernel), + format!("{:?}", parameters.kernel) + ); + assert_eq!(actual_parameters.seed, parameters.seed); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![-1, -1, 1, 1]; + let parameters: SVCParameters, Vec> = + SVCParameters::default().with_kernel(Kernels::linear()); + let mut classifier: SVC<'_, f64, i32, DenseMatrix, Vec> = + SVC::fit(&matrix, &target, ¶meters).unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/svm/svr.rs b/src/svm/svr.rs index 679a77a3..1c42abcf 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -284,10 +284,9 @@ impl<'a, T: Number + FloatNumber + PartialOrd, X: Array2, Y: Array1> SVR<' /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &SVRParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&SVRParameters> { + self.parameters } } @@ -717,4 +716,37 @@ mod tests { assert_eq!(svr, deserialized_svr); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters: SVRParameters = + SVRParameters::default().with_kernel(Kernels::linear()); + let regressor: SVR<'_, f64, DenseMatrix, Vec> = + SVR::fit(&matrix, &target, ¶meters).unwrap(); + + let actual_parameters = regressor + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.eps, parameters.eps); + assert_eq!(actual_parameters.c, parameters.c); + assert_eq!(actual_parameters.tol, parameters.tol); + assert_eq!(actual_parameters.kernel, parameters.kernel); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters: SVRParameters = + SVRParameters::default().with_kernel(Kernels::linear()); + let mut regressor: SVR<'_, f64, DenseMatrix, Vec> = + SVR::fit(&matrix, &target, ¶meters).unwrap(); + regressor.parameters = None; + + assert!(regressor.parameters().is_none()); + } } diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 2230e8a1..92de927b 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -66,10 +66,9 @@ impl, Y: Array1> /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - fn parameters(&self) -> &BaseTreeRegressorParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + fn parameters(&self) -> Option<&BaseTreeRegressorParameters> { + self.parameters.as_ref() } /// Get estimate of intercept, return value fn depth(&self) -> u16 { @@ -241,7 +240,7 @@ impl, Y: Array1> visitor_queue.push_back(visitor); } - while base_tree.depth() < base_tree.parameters().max_depth.unwrap_or(u16::MAX) { + while base_tree.depth() < base_tree.parameters().expect("parameters not set — model not fitted").max_depth.unwrap_or(u16::MAX) { match visitor_queue.pop_front() { Some(node) => base_tree.split(node, mtry, &mut visitor_queue, &mut rng), None => break, @@ -302,7 +301,9 @@ impl, Y: Array1> let n: usize = visitor.samples.iter().sum(); - if n < self.parameters().min_samples_split { + let parameters = self.parameters().expect("parameters not set — model not fitted"); + + if n < parameters.min_samples_split { return false; } @@ -317,7 +318,7 @@ impl, Y: Array1> let parent_gain = n as f64 * self.nodes()[visitor.node].output * self.nodes()[visitor.node].output; - let splitter = self.parameters().splitter.clone(); + let splitter = parameters.splitter.clone(); for variable in variables.iter().take(mtry) { match splitter { @@ -384,8 +385,10 @@ impl, Y: Array1> let false_count = n - true_count; - if true_count < self.parameters().min_samples_leaf - || false_count < self.parameters().min_samples_leaf + let parameters = self.parameters().expect("parameters not set — model not fitted"); + + if true_count < parameters.min_samples_leaf + || false_count < parameters.min_samples_leaf { return; } @@ -440,8 +443,10 @@ impl, Y: Array1> let false_count = n - true_count; - if true_count < self.parameters().min_samples_leaf - || false_count < self.parameters().min_samples_leaf + let parameters = self.parameters().expect("parameters not set — model not fitted"); + + if true_count < parameters.min_samples_leaf + || false_count < parameters.min_samples_leaf { prevx = Some(x_ij); true_count += visitor.samples[*i]; @@ -505,7 +510,9 @@ impl, Y: Array1> } } - if tc < self.parameters().min_samples_leaf || fc < self.parameters().min_samples_leaf { + let parameters = self.parameters().expect("parameters not set — model not fitted"); + + if tc < parameters.min_samples_leaf || fc < parameters.min_samples_leaf { self.nodes[visitor.node].split_feature = 0; self.nodes[visitor.node].split_value = Option::None; self.nodes[visitor.node].split_score = Option::None; diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index 2a2e7859..dadac9d8 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -134,10 +134,9 @@ impl, Y: Array1> /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &DecisionTreeClassifierParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&DecisionTreeClassifierParameters> { + self.parameters.as_ref() } /// get classes vector, return a shared reference fn classes(&self) -> &Vec { @@ -619,7 +618,7 @@ impl, Y: Array1> visitor_queue.push_back(visitor); } - while tree.depth() < tree.parameters().max_depth.unwrap_or(u16::MAX) { + while tree.depth() < tree.parameters().expect("parameters not set — model not fitted").max_depth.unwrap_or(u16::MAX) { match visitor_queue.pop_front() { Some(node) => tree.split(node, mtry, &mut visitor_queue, &mut rng), None => break, @@ -704,14 +703,17 @@ impl, Y: Array1> count[visitor.y[i]] += visitor.samples[i]; } } + + let parameters = self.parameters().expect("parameters not set — model not fitted"); + let parameters_min_samples_split = parameters.min_samples_split; - self.nodes[visitor.node].impurity = Some(impurity(&self.parameters().criterion, &count, n)); + self.nodes[visitor.node].impurity = Some(impurity(¶meters.criterion, &count, n)); if is_pure { return false; } - if n <= self.parameters().min_samples_split { + if n <= parameters_min_samples_split { return false; } @@ -753,9 +755,10 @@ impl, Y: Array1> let tc = true_count.iter().sum(); let fc = n - tc; + let parameters = self.parameters().expect("parameters not set — model not fitted"); - if tc < self.parameters().min_samples_leaf - || fc < self.parameters().min_samples_leaf + if tc < parameters.min_samples_leaf + || fc < parameters.min_samples_leaf { prevx = Some(x_ij); prevy = visitor.y[*i]; @@ -772,9 +775,9 @@ impl, Y: Array1> let parent_impurity = self.nodes()[visitor.node].impurity.unwrap(); let gain = parent_impurity - tc as f64 / n as f64 - * impurity(&self.parameters().criterion, &true_count, tc) + * impurity(¶meters.criterion, &true_count, tc) - fc as f64 / n as f64 - * impurity(&self.parameters().criterion, false_count, fc); + * impurity(¶meters.criterion, false_count, fc); if self.nodes()[visitor.node].split_score.is_none() || gain > self.nodes()[visitor.node].split_score.unwrap() @@ -824,8 +827,10 @@ impl, Y: Array1> } } } + + let parameters = self.parameters().expect("parameters not set — model not fitted"); - if tc < self.parameters().min_samples_leaf || fc < self.parameters().min_samples_leaf { + if tc < parameters.min_samples_leaf || fc < parameters.min_samples_leaf { self.nodes[visitor.node].split_feature = 0; self.nodes[visitor.node].split_value = Option::None; self.nodes[visitor.node].split_score = Option::None; @@ -1240,4 +1245,50 @@ mod tests { assert_eq!(tree, deserialized_tree); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters = DecisionTreeClassifierParameters::default(); + let expected_parameters = parameters.clone(); + let classifier = + DecisionTreeClassifier::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = classifier + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!( + std::mem::discriminant(&actual_parameters.criterion), + std::mem::discriminant(&expected_parameters.criterion) + ); + assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); + assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); + assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!(actual_parameters.seed, expected_parameters.seed); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0, 0, 1, 1]; + let parameters = DecisionTreeClassifierParameters::default(); + let mut classifier = + DecisionTreeClassifier::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + classifier.parameters = None; + + assert!(classifier.parameters().is_none()); + } } diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index ebfc1d8f..d65b3822 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -324,10 +324,9 @@ impl, Y: Array1> /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &DecisionTreeRegressorParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&DecisionTreeRegressorParameters> { + self.parameters.as_ref() } } @@ -478,4 +477,44 @@ mod tests { assert_eq!(tree, deserialized_tree); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = DecisionTreeRegressorParameters::default(); + let expected_parameters = parameters.clone(); + let regressor = DecisionTreeRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regressor + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); + assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); + assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!(actual_parameters.seed, expected_parameters.seed); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = DecisionTreeRegressorParameters::default(); + let mut regressor = DecisionTreeRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regressor.parameters = None; + + assert!(regressor.parameters().is_none()); + } } diff --git a/src/xgboost/xgb_regressor.rs b/src/xgboost/xgb_regressor.rs index 92a3f4f0..2d7d1f42 100644 --- a/src/xgboost/xgb_regressor.rs +++ b/src/xgboost/xgb_regressor.rs @@ -613,10 +613,9 @@ impl, Y: Array1> XGRegres /// Getter for parameters used in the model /// /// # Returns - /// Parameters used to setup the model - pub fn parameters(&self) -> &XGRegressorParameters { - assert!(self.parameters.is_some()); - &self.parameters.as_ref().unwrap() + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&XGRegressorParameters> { + self.parameters.as_ref() } } @@ -806,4 +805,53 @@ mod tests { let predictions = predict_result.unwrap(); assert_eq!(predictions.len(), 4); } + + #[test] + fn test_can_get_assigned_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = XGRegressorParameters::default(); + let expected_parameters = parameters.clone(); + let regressor = XGRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + + let actual_parameters = regressor + .parameters() + .expect("parameters should be set after fitting"); + assert_eq!(actual_parameters.n_estimators, expected_parameters.n_estimators); + assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); + assert_eq!(actual_parameters.learning_rate, expected_parameters.learning_rate); + assert_eq!(actual_parameters.min_child_weight, expected_parameters.min_child_weight); + assert_eq!(actual_parameters.lambda, expected_parameters.lambda); + assert_eq!(actual_parameters.gamma, expected_parameters.gamma); + assert_eq!(actual_parameters.base_score, expected_parameters.base_score); + assert_eq!(actual_parameters.subsample, expected_parameters.subsample); + assert_eq!(actual_parameters.seed, expected_parameters.seed); + assert_eq!( + std::mem::discriminant(&actual_parameters.objective), + std::mem::discriminant(&expected_parameters.objective) + ); + } + + #[test] + fn test_returns_none_on_no_parameters() { + let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; + let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); + let target = vec![0.0, 1.0, 1.0, 2.0]; + let parameters = XGRegressorParameters::default(); + let mut regressor = XGRegressor::, Vec>::fit( + &matrix, + &target, + parameters, + ) + .unwrap(); + regressor.parameters = None; + + assert!(regressor.parameters().is_none()); + } } From 7ceae8c8897b9e62b13ff7633618891126d57166 Mon Sep 17 00:00:00 2001 From: Hubert Majewski <36614010+HubertMajewski@users.noreply.github.com> Date: Wed, 22 Jul 2026 05:56:56 +0000 Subject: [PATCH 3/4] Add Deserialize without Default --- src/cluster/dbscan.rs | 60 +++++++++++++++++++++++++++- src/linear/logistic_regression.rs | 42 +++++++++++++++++++- src/naive_bayes/bernoulli.rs | 44 ++++++++++++++++++++- src/neighbors/knn_classifier.rs | 63 ++++++++++++++++++++++++++++- src/neighbors/knn_regressor.rs | 66 +++++++++++++++++++++++++++++-- 5 files changed, 263 insertions(+), 12 deletions(-) diff --git a/src/cluster/dbscan.rs b/src/cluster/dbscan.rs index 5a8c8fef..a2f7caa1 100644 --- a/src/cluster/dbscan.rs +++ b/src/cluster/dbscan.rs @@ -70,7 +70,7 @@ pub struct DBSCAN, Y: Array1, D: Dista _phantom_y: PhantomData, } -#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[cfg_attr(feature = "serde", derive(Serialize))] #[derive(Debug, Clone)] /// DBSCAN clustering algorithm parameters pub struct DBSCANParameters>> { @@ -92,6 +92,43 @@ pub struct DBSCANParameters>> { _phantom_t: PhantomData, } +// A manual implementation avoids adding an unnecessary `Default` bound to +// generic parameter types while preserving their serialized representation. +#[cfg(feature = "serde")] +impl<'de, T, D> Deserialize<'de> for DBSCANParameters +where + T: Number, + D: Distance> + Deserialize<'de>, +{ + fn deserialize(deserializer: De) -> Result + where + De: serde::Deserializer<'de>, + { + #[derive(Deserialize)] + struct DBSCANParametersData { + distance: D, + #[serde(default)] + min_samples: usize, + #[serde(default)] + eps: f64, + #[serde(default)] + algorithm: KNNAlgorithmName, + #[serde(default)] + _phantom_t: PhantomData<()>, + } + + let data = DBSCANParametersData::deserialize(deserializer)?; + + Ok(Self { + distance: data.distance, + min_samples: data.min_samples, + eps: data.eps, + algorithm: data.algorithm, + _phantom_t: PhantomData, + }) + } +} + impl>> DBSCANParameters { /// 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. @@ -501,12 +538,31 @@ mod tests { ]) .unwrap(); - let dbscan = DBSCAN::fit(&x, Default::default()).unwrap(); + let parameters = DBSCANParameters::default() + .with_eps(0.5) + .with_min_samples(2) + .with_algorithm(KNNAlgorithmName::LinearSearch); + let expected_parameters = parameters.clone(); + let dbscan = DBSCAN::fit(&x, parameters).unwrap(); let deserialized_dbscan: DBSCAN, Vec, Euclidian> = serde_json::from_str(&serde_json::to_string(&dbscan).unwrap()).unwrap(); assert_eq!(dbscan, deserialized_dbscan); + + let actual_parameters = deserialized_dbscan + .parameters() + .expect("parameters should survive serialization"); + assert_eq!( + format!("{:?}", actual_parameters.distance), + format!("{:?}", expected_parameters.distance) + ); + assert_eq!(actual_parameters.min_samples, expected_parameters.min_samples); + assert_eq!(actual_parameters.eps, expected_parameters.eps); + assert_eq!( + std::mem::discriminant(&actual_parameters.algorithm), + std::mem::discriminant(&expected_parameters.algorithm) + ); } #[cfg(feature = "datasets")] diff --git a/src/linear/logistic_regression.rs b/src/linear/logistic_regression.rs index e2170768..f69d7c43 100644 --- a/src/linear/logistic_regression.rs +++ b/src/linear/logistic_regression.rs @@ -80,7 +80,7 @@ pub enum LogisticRegressionSolverName { } /// Logistic Regression parameters -#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[cfg_attr(feature = "serde", derive(Serialize))] #[derive(Debug, Clone)] pub struct LogisticRegressionParameters { #[cfg_attr(feature = "serde", serde(default))] @@ -91,6 +91,33 @@ pub struct LogisticRegressionParameters { pub alpha: T, } +// A manual implementation avoids adding an unnecessary `Default` bound to +// generic parameter types while preserving their serialized representation. +#[cfg(feature = "serde")] +impl<'de, T> Deserialize<'de> for LogisticRegressionParameters +where + T: Number + FloatNumber + Deserialize<'de>, +{ + fn deserialize(deserializer: De) -> Result + where + De: serde::Deserializer<'de>, + { + #[derive(Deserialize)] + struct LogisticRegressionParametersData { + #[serde(default)] + solver: LogisticRegressionSolverName, + alpha: T, + } + + let data = LogisticRegressionParametersData::deserialize(deserializer)?; + + Ok(Self { + solver: data.solver, + alpha: data.alpha, + }) + } +} + /// Logistic Regression grid search parameters #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)] @@ -858,12 +885,23 @@ mod tests { .unwrap(); let y: Vec = vec![0, 0, 1, 1, 2, 1, 1, 0, 0, 2, 1, 1, 0, 0, 1]; - let lr = LogisticRegression::fit(&x, &y, Default::default()).unwrap(); + let parameters = LogisticRegressionParameters::default().with_alpha(10.0); + let expected_parameters = parameters.clone(); + let lr = LogisticRegression::fit(&x, &y, parameters).unwrap(); let deserialized_lr: LogisticRegression, Vec> = serde_json::from_str(&serde_json::to_string(&lr).unwrap()).unwrap(); assert_eq!(lr, deserialized_lr); + + let actual_parameters = deserialized_lr + .parameters() + .expect("parameters should survive serialization"); + assert_eq!( + std::mem::discriminant(&actual_parameters.solver), + std::mem::discriminant(&expected_parameters.solver) + ); + assert_eq!(actual_parameters.alpha, expected_parameters.alpha); } #[cfg_attr( diff --git a/src/naive_bayes/bernoulli.rs b/src/naive_bayes/bernoulli.rs index 0aef35b0..556b3946 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -126,7 +126,7 @@ impl NBDistribution } /// `BernoulliNB` parameters. Use `Default::default()` for default values. -#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[cfg_attr(feature = "serde", derive(Serialize))] #[derive(Debug, Clone, PartialEq)] pub struct BernoulliNBParameters { #[cfg_attr(feature = "serde", serde(default))] @@ -140,6 +140,35 @@ pub struct BernoulliNBParameters { pub binarize: Option, } +// A manual implementation avoids adding an unnecessary `Default` bound to +// generic parameter types while preserving their serialized representation. +#[cfg(feature = "serde")] +impl<'de, T> Deserialize<'de> for BernoulliNBParameters +where + T: Number + Deserialize<'de>, +{ + fn deserialize(deserializer: De) -> Result + where + De: serde::Deserializer<'de>, + { + #[derive(Deserialize)] + struct BernoulliNBParametersData { + #[serde(default)] + alpha: f64, + priors: Option>, + binarize: Option, + } + + let data = BernoulliNBParametersData::deserialize(deserializer)?; + + Ok(Self { + alpha: data.alpha, + priors: data.priors, + binarize: data.binarize, + }) + } +} + impl BernoulliNBParameters { /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). pub fn with_alpha(mut self, alpha: f64) -> Self { @@ -660,11 +689,22 @@ mod tests { .unwrap(); let y: Vec = vec![0, 0, 0, 1]; - let bnb = BernoulliNB::fit(&x, &y, Default::default()).unwrap(); + let parameters = BernoulliNBParameters::default() + .with_alpha(0.5) + .with_priors(vec![0.75, 0.25]) + .with_binarize(0); + let expected_parameters = parameters.clone(); + let bnb = BernoulliNB::fit(&x, &y, parameters).unwrap(); let deserialized_bnb: BernoulliNB, Vec> = serde_json::from_str(&serde_json::to_string(&bnb).unwrap()).unwrap(); assert_eq!(bnb, deserialized_bnb); + assert_eq!( + deserialized_bnb + .parameters() + .expect("parameters should survive serialization"), + &expected_parameters + ); } #[test] diff --git a/src/neighbors/knn_classifier.rs b/src/neighbors/knn_classifier.rs index cd6e04cc..df47277b 100644 --- a/src/neighbors/knn_classifier.rs +++ b/src/neighbors/knn_classifier.rs @@ -46,7 +46,7 @@ use crate::neighbors::KNNWeightFunction; use crate::numbers::basenum::Number; /// `KNNClassifier` parameters. Use `Default::default()` for default values. -#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[cfg_attr(feature = "serde", derive(Serialize))] #[derive(Debug, Clone)] pub struct KNNClassifierParameters>> { #[cfg_attr(feature = "serde", serde(default))] @@ -68,6 +68,43 @@ pub struct KNNClassifierParameters>> { t: PhantomData, } +// A manual implementation avoids adding an unnecessary `Default` bound to +// generic parameter types while preserving their serialized representation. +#[cfg(feature = "serde")] +impl<'de, T, D> Deserialize<'de> for KNNClassifierParameters +where + T: Number, + D: Distance> + Deserialize<'de>, +{ + fn deserialize(deserializer: De) -> Result + where + De: serde::Deserializer<'de>, + { + #[derive(Deserialize)] + struct KNNClassifierParametersData { + distance: D, + #[serde(default)] + algorithm: KNNAlgorithmName, + #[serde(default)] + weight: KNNWeightFunction, + #[serde(default)] + k: usize, + #[serde(default, rename = "t")] + _t: PhantomData<()>, + } + + let data = KNNClassifierParametersData::deserialize(deserializer)?; + + Ok(Self { + distance: data.distance, + algorithm: data.algorithm, + weight: data.weight, + k: data.k, + t: PhantomData, + }) + } +} + /// K Nearest Neighbors Classifier #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug)] @@ -731,10 +768,32 @@ mod tests { .unwrap(); let y = vec![2, 2, 2, 3, 3]; - let knn = KNNClassifier::fit(&x, &y, Default::default()).unwrap(); + let parameters = KNNClassifierParameters::default() + .with_k(2) + .with_algorithm(KNNAlgorithmName::LinearSearch) + .with_weight(KNNWeightFunction::Distance); + let expected_parameters = parameters.clone(); + let knn = KNNClassifier::fit(&x, &y, parameters).unwrap(); let deserialized_knn = bincode::deserialize(&bincode::serialize(&knn).unwrap()).unwrap(); assert_eq!(knn, deserialized_knn); + + let actual_parameters = deserialized_knn + .parameters() + .expect("parameters should survive serialization"); + assert_eq!( + format!("{:?}", actual_parameters.distance), + format!("{:?}", expected_parameters.distance) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.algorithm), + std::mem::discriminant(&expected_parameters.algorithm) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.weight), + std::mem::discriminant(&expected_parameters.weight) + ); + assert_eq!(actual_parameters.k, expected_parameters.k); } #[test] diff --git a/src/neighbors/knn_regressor.rs b/src/neighbors/knn_regressor.rs index da0e57dd..c98548c8 100644 --- a/src/neighbors/knn_regressor.rs +++ b/src/neighbors/knn_regressor.rs @@ -49,7 +49,7 @@ use crate::neighbors::KNNWeightFunction; use crate::numbers::basenum::Number; /// `KNNRegressor` parameters. Use `Default::default()` for default values. -#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[cfg_attr(feature = "serde", derive(Serialize))] #[derive(Debug, Clone)] pub struct KNNRegressorParameters>> { #[cfg_attr(feature = "serde", serde(default))] @@ -71,6 +71,43 @@ pub struct KNNRegressorParameters>> { t: PhantomData, } +// A manual implementation avoids adding an unnecessary `Default` bound to +// generic parameter types while preserving their serialized representation. +#[cfg(feature = "serde")] +impl<'de, T, D> Deserialize<'de> for KNNRegressorParameters +where + T: Number, + D: Distance> + Deserialize<'de>, +{ + fn deserialize(deserializer: De) -> Result + where + De: serde::Deserializer<'de>, + { + #[derive(Deserialize)] + struct KNNRegressorParameters { + distance: D, + #[serde(default)] + algorithm: KNNAlgorithmName, + #[serde(default)] + weight: KNNWeightFunction, + #[serde(default)] + k: usize, + #[serde(default, rename = "t")] + _t: PhantomData<()>, + } + + let data = KNNRegressorParameters::deserialize(deserializer)?; + + Ok(Self { + distance: data.distance, + algorithm: data.algorithm, + weight: data.weight, + k: data.k, + t: PhantomData, + }) + } +} + /// K Nearest Neighbors Regressor #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug)] @@ -355,13 +392,34 @@ mod tests { .unwrap(); let y = vec![1., 2., 3., 4., 5.]; - let knn = KNNRegressor::fit(&x, &y, Default::default()).unwrap(); - + let parameters = KNNRegressorParameters::default() + .with_k(2) + .with_algorithm(KNNAlgorithmName::LinearSearch) + .with_weight(KNNWeightFunction::Distance); + let expected_parameters = parameters.clone(); + let knn = KNNRegressor::fit(&x, &y, parameters).unwrap(); let deserialized_knn = bincode::deserialize(&bincode::serialize(&knn).unwrap()).unwrap(); assert_eq!(knn, deserialized_knn); + + let actual_parameters = deserialized_knn + .parameters() + .expect("parameters should survive serialization"); + assert_eq!( + format!("{:?}", actual_parameters.distance), + format!("{:?}", expected_parameters.distance) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.algorithm), + std::mem::discriminant(&expected_parameters.algorithm) + ); + assert_eq!( + std::mem::discriminant(&actual_parameters.weight), + std::mem::discriminant(&expected_parameters.weight) + ); + assert_eq!(actual_parameters.k, expected_parameters.k); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; From 8c8502890c8f023dd07591df1d6053412f773f42 Mon Sep 17 00:00:00 2001 From: Hubert Majewski <36614010+HubertMajewski@users.noreply.github.com> Date: Wed, 22 Jul 2026 05:58:37 +0000 Subject: [PATCH 4/4] Format; cargo fmt --- src/cluster/dbscan.rs | 44 ++++++++------- src/cluster/kmeans.rs | 18 +++---- src/decomposition/pca.rs | 9 ++-- src/decomposition/svd.rs | 7 ++- src/ensemble/extra_trees_regressor.rs | 47 ++++++++-------- src/ensemble/random_forest_classifier.rs | 45 ++++++++-------- src/ensemble/random_forest_regressor.rs | 45 ++++++++-------- src/linear/elastic_net.rs | 22 +++----- src/linear/lasso.rs | 27 +++++----- src/linear/linear_regression.rs | 10 ++-- src/linear/logistic_regression.rs | 18 +++---- src/linear/ridge_regression.rs | 12 ++--- src/naive_bayes/bernoulli.rs | 26 ++++----- src/naive_bayes/categorical.rs | 31 +++++------ src/naive_bayes/gaussian.rs | 29 +++++----- src/naive_bayes/multinomial.rs | 29 +++++----- src/neighbors/knn_classifier.rs | 18 +++---- src/neighbors/knn_regressor.rs | 21 ++++---- src/svm/svc.rs | 4 +- src/svm/svr.rs | 4 +- src/tree/base_tree_regressor.rs | 30 +++++++---- src/tree/decision_tree_classifier.rs | 68 +++++++++++++----------- src/tree/decision_tree_regressor.rs | 28 +++++----- src/xgboost/xgb_regressor.rs | 37 +++++++------ 24 files changed, 316 insertions(+), 313 deletions(-) diff --git a/src/cluster/dbscan.rs b/src/cluster/dbscan.rs index a2f7caa1..8879d77b 100644 --- a/src/cluster/dbscan.rs +++ b/src/cluster/dbscan.rs @@ -431,7 +431,7 @@ impl, Y: Array1, D: Distance>> Ok(result) } - + /// Getter for parameters used in the model /// /// # Returns @@ -557,7 +557,10 @@ mod tests { format!("{:?}", actual_parameters.distance), format!("{:?}", expected_parameters.distance) ); - assert_eq!(actual_parameters.min_samples, expected_parameters.min_samples); + assert_eq!( + actual_parameters.min_samples, + expected_parameters.min_samples + ); assert_eq!(actual_parameters.eps, expected_parameters.eps); assert_eq!( std::mem::discriminant(&actual_parameters.algorithm), @@ -580,20 +583,19 @@ mod tests { println!("{labels:?}"); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.5, 0.5, 10.0, 10.0]; let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); - let parameters: DBSCANParameters> = - DBSCANParameters::default().with_eps(1.0).with_min_samples(1); + let parameters: DBSCANParameters> = DBSCANParameters::default() + .with_eps(1.0) + .with_min_samples(1); let expected_parameters = parameters.clone(); - let clustering = - DBSCAN::, Vec, Euclidian>::fit( - &matrix, - parameters, - ) - .unwrap(); + let clustering = DBSCAN::, Vec, Euclidian>::fit( + &matrix, parameters, + ) + .unwrap(); let actual_parameters = clustering .parameters() @@ -602,7 +604,10 @@ mod tests { format!("{:?}", actual_parameters.distance), format!("{:?}", expected_parameters.distance) ); - assert_eq!(actual_parameters.min_samples, expected_parameters.min_samples); + assert_eq!( + actual_parameters.min_samples, + expected_parameters.min_samples + ); assert_eq!(actual_parameters.eps, expected_parameters.eps); assert_eq!( std::mem::discriminant(&actual_parameters.algorithm), @@ -614,14 +619,13 @@ mod tests { fn test_returns_none_on_no_parameters() { let data = vec![0.0, 0.0, 0.5, 0.5, 10.0, 10.0]; let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); - let parameters: DBSCANParameters> = - DBSCANParameters::default().with_eps(1.0).with_min_samples(1); - let mut clustering = - DBSCAN::, Vec, Euclidian>::fit( - &matrix, - parameters, - ) - .unwrap(); + let parameters: DBSCANParameters> = DBSCANParameters::default() + .with_eps(1.0) + .with_min_samples(1); + let mut clustering = DBSCAN::, Vec, Euclidian>::fit( + &matrix, parameters, + ) + .unwrap(); clustering.parameters = None; assert!(clustering.parameters().is_none()); diff --git a/src/cluster/kmeans.rs b/src/cluster/kmeans.rs index 6b6c7563..c0ac0b9e 100644 --- a/src/cluster/kmeans.rs +++ b/src/cluster/kmeans.rs @@ -413,7 +413,7 @@ impl, Y: Array1> KMeans y } - + /// Getter for parameters used in the model /// /// # Returns @@ -553,18 +553,15 @@ mod tests { assert_eq!(kmeans, deserialized_kmeans); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 5.0, 5.0, 10.0, 10.0]; let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); let parameters = KMeansParameters::default().with_k(2); let expected_parameters = parameters.clone(); - let clustering = KMeans::, Vec>::fit( - &matrix, - parameters, - ) - .unwrap(); + let clustering = + KMeans::, Vec>::fit(&matrix, parameters).unwrap(); let actual_parameters = clustering .parameters() @@ -579,11 +576,8 @@ mod tests { let data = vec![0.0, 0.0, 5.0, 5.0, 10.0, 10.0]; let matrix = DenseMatrix::new(3, 2, data, false).unwrap(); let parameters = KMeansParameters::default().with_k(2); - let mut clustering = KMeans::, Vec>::fit( - &matrix, - parameters, - ) - .unwrap(); + let mut clustering = + KMeans::, Vec>::fit(&matrix, parameters).unwrap(); clustering.parameters = None; assert!(clustering.parameters().is_none()); diff --git a/src/decomposition/pca.rs b/src/decomposition/pca.rs index 8189a730..c86fbc2e 100644 --- a/src/decomposition/pca.rs +++ b/src/decomposition/pca.rs @@ -362,7 +362,7 @@ impl + SVDDecomposable + EVDDecomposable pub fn components(&self) -> &X { &self.projection } - + /// Getter for parameters used in the model /// /// # Returns @@ -757,7 +757,7 @@ mod tests { // assert_eq!(pca, deserialized_pca); // } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -769,7 +769,10 @@ mod tests { let actual_parameters = pca .parameters() .expect("parameters should be set after fitting"); - assert_eq!(actual_parameters.n_components, expected_parameters.n_components); + assert_eq!( + actual_parameters.n_components, + expected_parameters.n_components + ); assert_eq!( actual_parameters.use_correlation_matrix, expected_parameters.use_correlation_matrix diff --git a/src/decomposition/svd.rs b/src/decomposition/svd.rs index 96fd1948..aa8e0916 100644 --- a/src/decomposition/svd.rs +++ b/src/decomposition/svd.rs @@ -362,7 +362,7 @@ mod tests { // assert_eq!(svd, deserialized_svd); // } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 0.0, 1.0]; @@ -374,7 +374,10 @@ mod tests { let actual_parameters = svd .parameters() .expect("parameters should be set after fitting"); - assert_eq!(actual_parameters.n_components, expected_parameters.n_components); + assert_eq!( + actual_parameters.n_components, + expected_parameters.n_components + ); } #[test] diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index 9d94afd9..6cb8e786 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -104,7 +104,7 @@ pub struct ExtraTreesRegressor< Y: Array1, > { forest_regressor: Option>, - parameters: Option + parameters: Option, } impl ExtraTreesRegressorParameters { @@ -166,7 +166,7 @@ impl, Y: Array1 fn new() -> Self { Self { forest_regressor: Option::None, - parameters: Option::None + parameters: Option::None, } } @@ -209,7 +209,7 @@ impl, Y: Array1 Ok(ExtraTreesRegressor { forest_regressor: Some(forest_regressor), - parameters: Some(parameters) + parameters: Some(parameters), }) } @@ -225,7 +225,7 @@ impl, Y: Array1 let forest_regressor = self.forest_regressor.as_ref().unwrap(); forest_regressor.predict_oob(x) } - + /// Getter for parameters used in the model /// /// # Returns @@ -326,7 +326,7 @@ mod tests { assert_eq!(y_hat1, y_hat2); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -334,23 +334,29 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = ExtraTreesRegressorParameters::default(); let expected_parameters = parameters.clone(); - let regressor = - ExtraTreesRegressor::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let regressor = ExtraTreesRegressor::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); let actual_parameters = regressor .parameters() .expect("parameters should be set after fitting"); assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); - assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); - assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!( + actual_parameters.min_samples_leaf, + expected_parameters.min_samples_leaf + ); + assert_eq!( + actual_parameters.min_samples_split, + expected_parameters.min_samples_split + ); assert_eq!(actual_parameters.n_trees, expected_parameters.n_trees); assert_eq!(actual_parameters.m, expected_parameters.m); - assert_eq!(actual_parameters.keep_samples, expected_parameters.keep_samples); + assert_eq!( + actual_parameters.keep_samples, + expected_parameters.keep_samples + ); assert_eq!(actual_parameters.seed, expected_parameters.seed); } @@ -360,13 +366,10 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = ExtraTreesRegressorParameters::default(); - let mut regressor = - ExtraTreesRegressor::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut regressor = ExtraTreesRegressor::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); regressor.parameters = None; assert!(regressor.parameters().is_none()); diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index a457eaeb..2bfe5062 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -107,7 +107,7 @@ pub struct RandomForestClassifier< trees: Option>>, classes: Option>, samples: Option>>, - parameters: Option + parameters: Option, } impl RandomForestClassifierParameters { @@ -508,7 +508,7 @@ impl, Y: Array1, Y: Array1, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let classifier = RandomForestClassifier::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); let actual_parameters = classifier .parameters() @@ -845,11 +842,20 @@ mod tests { std::mem::discriminant(&expected_parameters.criterion) ); assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); - assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); - assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!( + actual_parameters.min_samples_leaf, + expected_parameters.min_samples_leaf + ); + assert_eq!( + actual_parameters.min_samples_split, + expected_parameters.min_samples_split + ); assert_eq!(actual_parameters.n_trees, expected_parameters.n_trees); assert_eq!(actual_parameters.m, expected_parameters.m); - assert_eq!(actual_parameters.keep_samples, expected_parameters.keep_samples); + assert_eq!( + actual_parameters.keep_samples, + expected_parameters.keep_samples + ); assert_eq!(actual_parameters.seed, expected_parameters.seed); } @@ -859,13 +865,10 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0, 0, 1, 1]; let parameters = RandomForestClassifierParameters::default(); - let mut classifier = - RandomForestClassifier::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut classifier = RandomForestClassifier::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); classifier.parameters = None; assert!(classifier.parameters().is_none()); diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index b7bcb508..611b0ade 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -95,7 +95,7 @@ pub struct RandomForestRegressor< Y: Array1, > { forest_regressor: Option>, - parameters: Option + parameters: Option, } impl RandomForestRegressorParameters { @@ -401,7 +401,7 @@ impl, Y: Array1 Ok(RandomForestRegressor { forest_regressor: Some(forest_regressor), - parameters: Some(parameters) + parameters: Some(parameters), }) } @@ -417,7 +417,7 @@ impl, Y: Array1 let forest_regressor = self.forest_regressor.as_ref().unwrap(); forest_regressor.predict_oob(x) } - + /// Getter for parameters used in the model /// /// # Returns @@ -623,7 +623,7 @@ mod tests { assert_eq!(forest, deserialized_forest); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -631,23 +631,29 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = RandomForestRegressorParameters::default(); let expected_parameters = parameters.clone(); - let regressor = - RandomForestRegressor::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let regressor = RandomForestRegressor::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); let actual_parameters = regressor .parameters() .expect("parameters should be set after fitting"); assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); - assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); - assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!( + actual_parameters.min_samples_leaf, + expected_parameters.min_samples_leaf + ); + assert_eq!( + actual_parameters.min_samples_split, + expected_parameters.min_samples_split + ); assert_eq!(actual_parameters.n_trees, expected_parameters.n_trees); assert_eq!(actual_parameters.m, expected_parameters.m); - assert_eq!(actual_parameters.keep_samples, expected_parameters.keep_samples); + assert_eq!( + actual_parameters.keep_samples, + expected_parameters.keep_samples + ); assert_eq!(actual_parameters.seed, expected_parameters.seed); } @@ -657,13 +663,10 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = RandomForestRegressorParameters::default(); - let mut regressor = - RandomForestRegressor::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut regressor = RandomForestRegressor::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); regressor.parameters = None; assert!(regressor.parameters().is_none()); diff --git a/src/linear/elastic_net.rs b/src/linear/elastic_net.rs index 0cafb9c1..f9ea706e 100644 --- a/src/linear/elastic_net.rs +++ b/src/linear/elastic_net.rs @@ -462,7 +462,7 @@ impl, Y: Array1> (x2, y2, gamma) } - + /// Getter for parameters used in the model /// /// # Returns @@ -656,7 +656,7 @@ mod tests { // assert_eq!(lr, deserialized_lr); // } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -664,12 +664,9 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = ElasticNetParameters::default(); let expected_parameters = parameters.clone(); - let regression = ElasticNet::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let regression = + ElasticNet::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); let actual_parameters = regression .parameters() @@ -687,12 +684,9 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = ElasticNetParameters::default(); - let mut regression = ElasticNet::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut regression = + ElasticNet::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); regression.parameters = None; assert!(regression.parameters().is_none()); diff --git a/src/linear/lasso.rs b/src/linear/lasso.rs index 81092cce..de2479ad 100644 --- a/src/linear/lasso.rs +++ b/src/linear/lasso.rs @@ -405,7 +405,7 @@ impl, Y: Array1> Las scaled_x.scale_mut(&col_mean, &col_std, 0); Ok((scaled_x, col_mean, col_std)) } - + /// Getter for parameters used in the model /// /// # Returns @@ -585,7 +585,7 @@ mod tests { // assert_eq!(lr, deserialized_lr); // } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -593,12 +593,9 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = LassoParameters::default(); let expected_parameters = parameters.clone(); - let regression = Lasso::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let regression = + Lasso::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); let actual_parameters = regression .parameters() @@ -607,7 +604,10 @@ mod tests { assert_eq!(actual_parameters.normalize, expected_parameters.normalize); assert_eq!(actual_parameters.tol, expected_parameters.tol); assert_eq!(actual_parameters.max_iter, expected_parameters.max_iter); - assert_eq!(actual_parameters.fit_intercept, expected_parameters.fit_intercept); + assert_eq!( + actual_parameters.fit_intercept, + expected_parameters.fit_intercept + ); } #[test] @@ -616,12 +616,9 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = LassoParameters::default(); - let mut regression = Lasso::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut regression = + Lasso::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); regression.parameters = None; assert!(regression.parameters().is_none()); diff --git a/src/linear/linear_regression.rs b/src/linear/linear_regression.rs index ea97be3e..3d31bd73 100644 --- a/src/linear/linear_regression.rs +++ b/src/linear/linear_regression.rs @@ -423,7 +423,7 @@ mod tests { // let parameters: LinearRegressionParameters = serde_json::from_str("{}").unwrap(); // assert_eq!(parameters.solver, default.solver); // } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -432,9 +432,7 @@ mod tests { let parameters = LinearRegressionParameters::default(); let expected_parameters = parameters.clone(); let regression = LinearRegression::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); @@ -451,9 +449,7 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = LinearRegressionParameters::default(); let mut regression = LinearRegression::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); regression.parameters = None; diff --git a/src/linear/logistic_regression.rs b/src/linear/logistic_regression.rs index f69d7c43..8ef3078f 100644 --- a/src/linear/logistic_regression.rs +++ b/src/linear/logistic_regression.rs @@ -591,7 +591,7 @@ impl, Y: optimizer.optimize(&f, &df, &x0, &ls) } - + /// Getter for parameters used in the model /// /// # Returns @@ -996,19 +996,16 @@ mod tests { assert_eq!(y_hat.shape(), 52181); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0, 0, 1, 1]; - let parameters: LogisticRegressionParameters = - LogisticRegressionParameters::default(); + let parameters: LogisticRegressionParameters = LogisticRegressionParameters::default(); let expected_parameters = parameters.clone(); let regression = LogisticRegression::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); @@ -1024,12 +1021,9 @@ mod tests { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0, 0, 1, 1]; - let parameters: LogisticRegressionParameters = - LogisticRegressionParameters::default(); + let parameters: LogisticRegressionParameters = LogisticRegressionParameters::default(); let mut regression = LogisticRegression::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); regression.parameters = None; diff --git a/src/linear/ridge_regression.rs b/src/linear/ridge_regression.rs index 10c28439..0c68bb47 100644 --- a/src/linear/ridge_regression.rs +++ b/src/linear/ridge_regression.rs @@ -412,7 +412,7 @@ impl< pub fn intercept(&self) -> &TX { self.intercept.as_ref().unwrap() } - + /// Getter for parameters used in the model /// /// # Returns @@ -539,7 +539,7 @@ mod tests { // assert_eq!(lr, deserialized_lr); // } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -548,9 +548,7 @@ mod tests { let parameters: RidgeRegressionParameters = RidgeRegressionParameters::default(); let expected_parameters = parameters.clone(); let regression = RidgeRegression::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); @@ -569,9 +567,7 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters: RidgeRegressionParameters = RidgeRegressionParameters::default(); let mut regression = RidgeRegression::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); regression.parameters = None; diff --git a/src/naive_bayes/bernoulli.rs b/src/naive_bayes/bernoulli.rs index 556b3946..d799cd0a 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -386,7 +386,7 @@ pub struct BernoulliNB< > { inner: Option>>, binarize: Option, - parameters: Option> + parameters: Option>, } impl, Y: Array1> @@ -410,7 +410,7 @@ impl, Y: Arr Self { inner: Option::None, binarize: Option::None, - parameters: Option::None + parameters: Option::None, } } @@ -452,7 +452,7 @@ impl, Y: Arr Ok(Self { inner: Some(inner), binarize: parameters.binarize, - parameters: Some(parameters) + parameters: Some(parameters), }) } @@ -706,7 +706,7 @@ mod tests { &expected_parameters ); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0, 0, 0, 1, 1, 0, 1, 1]; @@ -714,12 +714,9 @@ mod tests { let target = vec![0_u32, 0, 1, 1]; let parameters: BernoulliNBParameters = BernoulliNBParameters::default(); let expected_parameters = parameters.clone(); - let classifier = BernoulliNB::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let classifier = + BernoulliNB::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); let actual_parameters = classifier .parameters() @@ -733,12 +730,9 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0_u32, 0, 1, 1]; let parameters: BernoulliNBParameters = BernoulliNBParameters::default(); - let mut classifier = BernoulliNB::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut classifier = + BernoulliNB::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); classifier.parameters = None; assert!(classifier.parameters().is_none()); diff --git a/src/naive_bayes/categorical.rs b/src/naive_bayes/categorical.rs index 020a0391..30135217 100644 --- a/src/naive_bayes/categorical.rs +++ b/src/naive_bayes/categorical.rs @@ -338,7 +338,7 @@ impl Default for CategoricalNBSearchParameters { #[derive(Debug, PartialEq)] pub struct CategoricalNB, Y: Array1> { inner: Option>>, - parameters: Option + parameters: Option, } impl, Y: Array1> @@ -347,7 +347,7 @@ impl, Y: Array1> fn new() -> Self { Self { inner: Option::None, - parameters: Option::None + parameters: Option::None, } } @@ -372,7 +372,10 @@ impl, Y: Array1> CategoricalNB { let alpha = parameters.alpha; let distribution = CategoricalNBDistribution::fit(x, y, alpha)?; let inner = BaseNaiveBayes::fit(distribution)?; - Ok(Self { inner: Some(inner), parameters: Some(parameters) }) + Ok(Self { + inner: Some(inner), + parameters: Some(parameters), + }) } /// Estimates the class labels for the provided data. @@ -417,7 +420,7 @@ impl, Y: Array1> CategoricalNB { pub fn feature_log_prob(&self) -> &Vec>> { &self.inner.as_ref().unwrap().distribution.coefficients } - + /// Getter for parameters used in the model /// /// # Returns @@ -595,7 +598,7 @@ mod tests { assert_eq!(cnb, deserialized_cnb); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0_u32, 0, 0, 1, 1, 0, 1, 1]; @@ -603,12 +606,9 @@ mod tests { let target = vec![0_u32, 0, 1, 1]; let parameters = CategoricalNBParameters::default(); let expected_parameters = parameters.clone(); - let classifier = CategoricalNB::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let classifier = + CategoricalNB::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); let actual_parameters = classifier .parameters() @@ -622,12 +622,9 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0_u32, 0, 1, 1]; let parameters = CategoricalNBParameters::default(); - let mut classifier = CategoricalNB::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut classifier = + CategoricalNB::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); classifier.parameters = None; assert!(classifier.parameters().is_none()); diff --git a/src/naive_bayes/gaussian.rs b/src/naive_bayes/gaussian.rs index d41028fc..0107be89 100644 --- a/src/naive_bayes/gaussian.rs +++ b/src/naive_bayes/gaussian.rs @@ -266,7 +266,7 @@ pub struct GaussianNB< Y: Array1, > { inner: Option>>, - parameters: Option + parameters: Option, } impl< @@ -324,7 +324,10 @@ impl, Y: Arr pub fn fit(x: &X, y: &Y, parameters: GaussianNBParameters) -> Result { let distribution = GaussianNBDistribution::fit(x, y, parameters.priors.clone())?; let inner = BaseNaiveBayes::fit(distribution)?; - Ok(Self { inner: Some(inner), parameters: Some(parameters) }) + Ok(Self { + inner: Some(inner), + parameters: Some(parameters), + }) } /// Estimates the class labels for the provided data. @@ -364,7 +367,7 @@ impl, Y: Arr pub fn var(&self) -> &Vec> { &self.inner.as_ref().unwrap().distribution.var } - + /// Getter for parameters used in the model /// /// # Returns @@ -484,7 +487,7 @@ mod tests { assert_eq!(gnb, deserialized_gnb); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -492,12 +495,9 @@ mod tests { let target = vec![0_u32, 0, 1, 1]; let parameters = GaussianNBParameters::default(); let expected_parameters = parameters.clone(); - let classifier = GaussianNB::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let classifier = + GaussianNB::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); let actual_parameters = classifier .parameters() @@ -511,12 +511,9 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0_u32, 0, 1, 1]; let parameters = GaussianNBParameters::default(); - let mut classifier = GaussianNB::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut classifier = + GaussianNB::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); classifier.parameters = None; assert!(classifier.parameters().is_none()); diff --git a/src/naive_bayes/multinomial.rs b/src/naive_bayes/multinomial.rs index bcd89ac3..8ab32bef 100644 --- a/src/naive_bayes/multinomial.rs +++ b/src/naive_bayes/multinomial.rs @@ -302,7 +302,7 @@ pub struct MultinomialNB< Y: Array1, > { inner: Option>>, - parameters: Option + parameters: Option, } impl, Y: Array1> fmt::Display @@ -324,7 +324,7 @@ impl, Y: Array fn new() -> Self { Self { inner: Option::None, - parameters: Option::None + parameters: Option::None, } } @@ -351,10 +351,17 @@ impl, Y: Array /// * `parameters` - additional parameters like class priors, alpha for smoothing and /// binarizing threshold. pub fn fit(x: &X, y: &Y, parameters: MultinomialNBParameters) -> Result { - let distribution = - MultinomialNBDistribution::fit(x, y, parameters.alpha.clone(), parameters.priors.clone())?; + let distribution = MultinomialNBDistribution::fit( + x, + y, + parameters.alpha.clone(), + parameters.priors.clone(), + )?; let inner = BaseNaiveBayes::fit(distribution)?; - Ok(Self { inner: Some(inner), parameters: Some(parameters) }) + Ok(Self { + inner: Some(inner), + parameters: Some(parameters), + }) } /// Estimates the class labels for the provided data. @@ -393,7 +400,7 @@ impl, Y: Array pub fn feature_count(&self) -> &Vec> { &self.inner.as_ref().unwrap().distribution.feature_count } - + /// Getter for parameters used in the model /// /// # Returns @@ -576,7 +583,7 @@ mod tests { assert_eq!(mnb, deserialized_mnb); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0_u32, 1, 1, 0, 1, 2, 2, 1]; @@ -585,9 +592,7 @@ mod tests { let parameters = MultinomialNBParameters::default(); let expected_parameters = parameters.clone(); let classifier = MultinomialNB::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); @@ -604,9 +609,7 @@ mod tests { let target = vec![0_u32, 0, 1, 1]; let parameters = MultinomialNBParameters::default(); let mut classifier = MultinomialNB::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); classifier.parameters = None; diff --git a/src/neighbors/knn_classifier.rs b/src/neighbors/knn_classifier.rs index df47277b..7a1d356b 100644 --- a/src/neighbors/knn_classifier.rs +++ b/src/neighbors/knn_classifier.rs @@ -290,7 +290,11 @@ impl, Y: Array1, D: Distance, Y: Array1, D: Distance, Vec, Euclidian>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); @@ -839,9 +841,7 @@ mod tests { KNNClassifierParameters::default(); let mut classifier = KNNClassifier::, Vec, Euclidian>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); classifier.parameters = None; diff --git a/src/neighbors/knn_regressor.rs b/src/neighbors/knn_regressor.rs index c98548c8..09d48bb4 100644 --- a/src/neighbors/knn_regressor.rs +++ b/src/neighbors/knn_regressor.rs @@ -270,7 +270,9 @@ impl, Y: Array1, D: Distance>> ))); } - let knn_algo = parameters.algorithm.fit(data, parameters.distance.clone())?; + let knn_algo = parameters + .algorithm + .fit(data, parameters.distance.clone())?; Ok(KNNRegressor { y: Some(y.clone()), @@ -317,7 +319,7 @@ impl, Y: Array1, D: Distance>> Ok(result) } - + /// Getter for parameters used in the model /// /// # Returns @@ -428,13 +430,10 @@ mod tests { let parameters: KNNRegressorParameters> = KNNRegressorParameters::default(); let expected_parameters = parameters.clone(); - let regressor = - KNNRegressor::, Vec, Euclidian>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let regressor = KNNRegressor::, Vec, Euclidian>::fit( + &matrix, &target, parameters, + ) + .unwrap(); let actual_parameters = regressor .parameters() @@ -463,9 +462,7 @@ mod tests { KNNRegressorParameters::default(); let mut regressor = KNNRegressor::, Vec, Euclidian>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); regressor.parameters = None; diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 116dae6c..84a959ee 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -608,7 +608,7 @@ impl<'a, TX: Number + RealNumber, TY: Number + Ord, X: Array2 + 'a, Y: Array f } - + /// Getter for parameters used in the model /// /// # Returns @@ -1437,7 +1437,7 @@ mod tests { assert_eq!(svc, deserialized_svc); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; diff --git a/src/svm/svr.rs b/src/svm/svr.rs index 1c42abcf..38778617 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -280,7 +280,7 @@ impl<'a, T: Number + FloatNumber + PartialOrd, X: Array2, Y: Array1> SVR<' T::from(f).unwrap() } - + /// Getter for parameters used in the model /// /// # Returns @@ -716,7 +716,7 @@ mod tests { assert_eq!(svr, deserialized_svr); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 92de927b..ac594f7d 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -240,7 +240,13 @@ impl, Y: Array1> visitor_queue.push_back(visitor); } - while base_tree.depth() < base_tree.parameters().expect("parameters not set — model not fitted").max_depth.unwrap_or(u16::MAX) { + while base_tree.depth() + < base_tree + .parameters() + .expect("parameters not set — model not fitted") + .max_depth + .unwrap_or(u16::MAX) + { match visitor_queue.pop_front() { Some(node) => base_tree.split(node, mtry, &mut visitor_queue, &mut rng), None => break, @@ -301,8 +307,10 @@ impl, Y: Array1> let n: usize = visitor.samples.iter().sum(); - let parameters = self.parameters().expect("parameters not set — model not fitted"); - + let parameters = self + .parameters() + .expect("parameters not set — model not fitted"); + if n < parameters.min_samples_split { return false; } @@ -385,11 +393,11 @@ impl, Y: Array1> let false_count = n - true_count; - let parameters = self.parameters().expect("parameters not set — model not fitted"); + let parameters = self + .parameters() + .expect("parameters not set — model not fitted"); - if true_count < parameters.min_samples_leaf - || false_count < parameters.min_samples_leaf - { + if true_count < parameters.min_samples_leaf || false_count < parameters.min_samples_leaf { return; } @@ -443,7 +451,9 @@ impl, Y: Array1> let false_count = n - true_count; - let parameters = self.parameters().expect("parameters not set — model not fitted"); + let parameters = self + .parameters() + .expect("parameters not set — model not fitted"); if true_count < parameters.min_samples_leaf || false_count < parameters.min_samples_leaf @@ -510,7 +520,9 @@ impl, Y: Array1> } } - let parameters = self.parameters().expect("parameters not set — model not fitted"); + let parameters = self + .parameters() + .expect("parameters not set — model not fitted"); if tc < parameters.min_samples_leaf || fc < parameters.min_samples_leaf { self.nodes[visitor.node].split_feature = 0; diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index dadac9d8..768d2d5d 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -618,7 +618,13 @@ impl, Y: Array1> visitor_queue.push_back(visitor); } - while tree.depth() < tree.parameters().expect("parameters not set — model not fitted").max_depth.unwrap_or(u16::MAX) { + while tree.depth() + < tree + .parameters() + .expect("parameters not set — model not fitted") + .max_depth + .unwrap_or(u16::MAX) + { match visitor_queue.pop_front() { Some(node) => tree.split(node, mtry, &mut visitor_queue, &mut rng), None => break, @@ -703,8 +709,10 @@ impl, Y: Array1> count[visitor.y[i]] += visitor.samples[i]; } } - - let parameters = self.parameters().expect("parameters not set — model not fitted"); + + let parameters = self + .parameters() + .expect("parameters not set — model not fitted"); let parameters_min_samples_split = parameters.min_samples_split; self.nodes[visitor.node].impurity = Some(impurity(¶meters.criterion, &count, n)); @@ -755,11 +763,11 @@ impl, Y: Array1> let tc = true_count.iter().sum(); let fc = n - tc; - let parameters = self.parameters().expect("parameters not set — model not fitted"); + let parameters = self + .parameters() + .expect("parameters not set — model not fitted"); - if tc < parameters.min_samples_leaf - || fc < parameters.min_samples_leaf - { + if tc < parameters.min_samples_leaf || fc < parameters.min_samples_leaf { prevx = Some(x_ij); prevy = visitor.y[*i]; true_count[visitor.y[*i]] += visitor.samples[*i]; @@ -774,10 +782,8 @@ impl, Y: Array1> let false_label = which_max(false_count); let parent_impurity = self.nodes()[visitor.node].impurity.unwrap(); let gain = parent_impurity - - tc as f64 / n as f64 - * impurity(¶meters.criterion, &true_count, tc) - - fc as f64 / n as f64 - * impurity(¶meters.criterion, false_count, fc); + - tc as f64 / n as f64 * impurity(¶meters.criterion, &true_count, tc) + - fc as f64 / n as f64 * impurity(¶meters.criterion, false_count, fc); if self.nodes()[visitor.node].split_score.is_none() || gain > self.nodes()[visitor.node].split_score.unwrap() @@ -827,8 +833,10 @@ impl, Y: Array1> } } } - - let parameters = self.parameters().expect("parameters not set — model not fitted"); + + let parameters = self + .parameters() + .expect("parameters not set — model not fitted"); if tc < parameters.min_samples_leaf || fc < parameters.min_samples_leaf { self.nodes[visitor.node].split_feature = 0; @@ -1245,7 +1253,7 @@ mod tests { assert_eq!(tree, deserialized_tree); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -1253,13 +1261,10 @@ mod tests { let target = vec![0, 0, 1, 1]; let parameters = DecisionTreeClassifierParameters::default(); let expected_parameters = parameters.clone(); - let classifier = - DecisionTreeClassifier::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let classifier = DecisionTreeClassifier::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); let actual_parameters = classifier .parameters() @@ -1269,8 +1274,14 @@ mod tests { std::mem::discriminant(&expected_parameters.criterion) ); assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); - assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); - assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!( + actual_parameters.min_samples_leaf, + expected_parameters.min_samples_leaf + ); + assert_eq!( + actual_parameters.min_samples_split, + expected_parameters.min_samples_split + ); assert_eq!(actual_parameters.seed, expected_parameters.seed); } @@ -1280,13 +1291,10 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0, 0, 1, 1]; let parameters = DecisionTreeClassifierParameters::default(); - let mut classifier = - DecisionTreeClassifier::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut classifier = DecisionTreeClassifier::, Vec>::fit( + &matrix, &target, parameters, + ) + .unwrap(); classifier.parameters = None; assert!(classifier.parameters().is_none()); diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index d65b3822..88dcc2dd 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -94,7 +94,7 @@ pub struct DecisionTreeRegressorParameters { pub struct DecisionTreeRegressor, Y: Array1> { tree_regressor: Option>, - parameters: Option + parameters: Option, } impl DecisionTreeRegressorParameters { @@ -273,7 +273,7 @@ impl, Y: Array1> fn new() -> Self { Self { tree_regressor: Option::None, - parameters: Option::None + parameters: Option::None, } } @@ -311,7 +311,7 @@ impl, Y: Array1> let tree = BaseTreeRegressor::fit(x, y, tree_parameters)?; Ok(Self { tree_regressor: Some(tree), - parameters: Some(parameters) + parameters: Some(parameters), }) } @@ -320,7 +320,7 @@ impl, Y: Array1> pub fn predict(&self, x: &X) -> Result { self.tree_regressor.as_ref().unwrap().predict(x) } - + /// Getter for parameters used in the model /// /// # Returns @@ -477,7 +477,7 @@ mod tests { assert_eq!(tree, deserialized_tree); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -486,9 +486,7 @@ mod tests { let parameters = DecisionTreeRegressorParameters::default(); let expected_parameters = parameters.clone(); let regressor = DecisionTreeRegressor::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); @@ -496,8 +494,14 @@ mod tests { .parameters() .expect("parameters should be set after fitting"); assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); - assert_eq!(actual_parameters.min_samples_leaf, expected_parameters.min_samples_leaf); - assert_eq!(actual_parameters.min_samples_split, expected_parameters.min_samples_split); + assert_eq!( + actual_parameters.min_samples_leaf, + expected_parameters.min_samples_leaf + ); + assert_eq!( + actual_parameters.min_samples_split, + expected_parameters.min_samples_split + ); assert_eq!(actual_parameters.seed, expected_parameters.seed); } @@ -508,9 +512,7 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = DecisionTreeRegressorParameters::default(); let mut regressor = DecisionTreeRegressor::, Vec>::fit( - &matrix, - &target, - parameters, + &matrix, &target, parameters, ) .unwrap(); regressor.parameters = None; diff --git a/src/xgboost/xgb_regressor.rs b/src/xgboost/xgb_regressor.rs index 2d7d1f42..612c0fb6 100644 --- a/src/xgboost/xgb_regressor.rs +++ b/src/xgboost/xgb_regressor.rs @@ -609,7 +609,7 @@ impl, Y: Array1> XGRegres indices.truncate((population_size as f64 * subsample_ratio) as usize); indices } - + /// Getter for parameters used in the model /// /// # Returns @@ -805,7 +805,7 @@ mod tests { let predictions = predict_result.unwrap(); assert_eq!(predictions.len(), 4); } - + #[test] fn test_can_get_assigned_parameters() { let data = vec![0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0]; @@ -813,20 +813,26 @@ mod tests { let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = XGRegressorParameters::default(); let expected_parameters = parameters.clone(); - let regressor = XGRegressor::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let regressor = + XGRegressor::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); let actual_parameters = regressor .parameters() .expect("parameters should be set after fitting"); - assert_eq!(actual_parameters.n_estimators, expected_parameters.n_estimators); + assert_eq!( + actual_parameters.n_estimators, + expected_parameters.n_estimators + ); assert_eq!(actual_parameters.max_depth, expected_parameters.max_depth); - assert_eq!(actual_parameters.learning_rate, expected_parameters.learning_rate); - assert_eq!(actual_parameters.min_child_weight, expected_parameters.min_child_weight); + assert_eq!( + actual_parameters.learning_rate, + expected_parameters.learning_rate + ); + assert_eq!( + actual_parameters.min_child_weight, + expected_parameters.min_child_weight + ); assert_eq!(actual_parameters.lambda, expected_parameters.lambda); assert_eq!(actual_parameters.gamma, expected_parameters.gamma); assert_eq!(actual_parameters.base_score, expected_parameters.base_score); @@ -844,12 +850,9 @@ mod tests { let matrix = DenseMatrix::new(4, 2, data, false).unwrap(); let target = vec![0.0, 1.0, 1.0, 2.0]; let parameters = XGRegressorParameters::default(); - let mut regressor = XGRegressor::, Vec>::fit( - &matrix, - &target, - parameters, - ) - .unwrap(); + let mut regressor = + XGRegressor::, Vec>::fit(&matrix, &target, parameters) + .unwrap(); regressor.parameters = None; assert!(regressor.parameters().is_none());