diff --git a/src/cluster/agglomerative.rs b/src/cluster/agglomerative.rs index 373f6f95..88f468e6 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,21 @@ 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 + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&AgglomerativeClusteringParameters> { + self.parameters.as_ref() + } } impl, Y: Array1> @@ -314,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 2e2aac10..8879d77b 100644 --- a/src/cluster/dbscan.rs +++ b/src/cluster/dbscan.rs @@ -64,12 +64,13 @@ 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, } -#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] +#[cfg_attr(feature = "serde", derive(Serialize))] #[derive(Debug, Clone)] /// DBSCAN clustering algorithm parameters pub struct DBSCANParameters>> { @@ -91,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. @@ -295,7 +333,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 +391,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 +431,14 @@ impl, Y: Array1, D: Distance>> Ok(result) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&DBSCANParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -491,12 +538,34 @@ 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")] @@ -514,4 +583,51 @@ 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 b81ffd7e..c0ac0b9e 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,14 @@ impl, Y: Array1> KMeans y } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&KMeansParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -543,4 +553,33 @@ 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 11853648..c86fbc2e 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,14 @@ impl + SVDDecomposable + EVDDecomposable pub fn components(&self) -> &X { &self.projection } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&PCAParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -747,4 +757,36 @@ 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 259bfbc0..aa8e0916 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,14 @@ impl + SVDDecomposable + EVDDecomposable pub fn components(&self) -> &X { &self.components } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&SVDParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -352,4 +362,32 @@ 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 818ac6c7..6cb8e786 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,14 @@ 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 + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&ExtraTreesRegressorParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -315,4 +326,52 @@ 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 0f86a4df..2bfe5062 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 Option<&RandomForestClassifierParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -810,4 +821,56 @@ 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 0a8a888c..611b0ade 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,14 @@ 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 + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&RandomForestRegressorParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -612,4 +623,52 @@ 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 d5b1d4d5..f9ea706e 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,14 @@ impl, Y: Array1> (x2, y2, gamma) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&ElasticNetParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -645,4 +656,39 @@ 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 59c60ddc..de2479ad 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,14 @@ 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 + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&LassoParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -574,4 +585,42 @@ 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 43410bbb..3d31bd73 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,14 @@ impl< pub fn intercept(&self) -> &TX { self.intercept.as_ref().unwrap() } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&LinearRegressionParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -412,4 +423,37 @@ 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 c28dc347..8ef3078f 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)] @@ -178,6 +205,7 @@ pub struct LogisticRegression< classes: Option>, num_attributes: usize, num_classes: usize, + parameters: Option>, _phantom_tx: PhantomData, _phantom_y: PhantomData, } @@ -389,6 +417,7 @@ impl, Y: classes: Option::None, num_attributes: 0, num_classes: 0, + parameters: Option::None, _phantom_tx: PhantomData, _phantom_y: PhantomData, } @@ -465,6 +494,7 @@ impl, Y: classes: Some(classes), num_attributes, num_classes: k, + parameters: Some(parameters), _phantom_tx: PhantomData, _phantom_y: PhantomData, }) @@ -491,6 +521,7 @@ impl, Y: classes: Some(classes), num_attributes, num_classes: k, + parameters: Some(parameters), _phantom_tx: PhantomData, _phantom_y: PhantomData, }) @@ -560,6 +591,14 @@ impl, Y: optimizer.optimize(&f, &df, &x0, &ls) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&LogisticRegressionParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -846,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( @@ -946,4 +996,38 @@ 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 be2f3d41..0c68bb47 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,14 @@ impl< pub fn intercept(&self) -> &TX { self.intercept.as_ref().unwrap() } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&RidgeRegressionParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -528,4 +539,39 @@ 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 cdd5b83d..d799cd0a 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -126,8 +126,8 @@ impl NBDistribution } /// `BernoulliNB` parameters. Use `Default::default()` for default values. -#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone)] +#[cfg_attr(feature = "serde", derive(Serialize))] +#[derive(Debug, Clone, PartialEq)] pub struct BernoulliNBParameters { #[cfg_attr(feature = "serde", serde(default))] /// Additive (Laplace/Lidstone) smoothing parameter (0 for no smoothing). @@ -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 { @@ -357,6 +386,7 @@ pub struct BernoulliNB< > { inner: Option>>, binarize: Option, + parameters: Option>, } impl, Y: Array1> @@ -380,6 +410,7 @@ impl, Y: Arr Self { inner: Option::None, binarize: Option::None, + parameters: Option::None, } } @@ -410,17 +441,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 +517,14 @@ impl, Y: Arr Self::binarize_mut(&mut new_x, threshold); new_x } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&BernoulliNBParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -649,10 +689,52 @@ 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] + 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 b60ee0d3..30135217 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,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) }) + Ok(Self { + inner: Some(inner), + parameters: Some(parameters), + }) } /// Estimates the class labels for the provided data. @@ -415,6 +420,14 @@ 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 + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&CategoricalNBParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -585,4 +598,35 @@ 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 dbf3fd81..0107be89 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,12 @@ 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 +367,14 @@ impl, Y: Arr pub fn var(&self) -> &Vec> { &self.inner.as_ref().unwrap().distribution.var } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&GaussianNBParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -474,4 +487,35 @@ 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 ad873943..8ab32bef 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, } } @@ -349,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, parameters.priors)?; + let distribution = 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 +400,14 @@ 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 + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&MultinomialNBParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -566,4 +583,37 @@ 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 264ab0e2..7a1d356b 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)] @@ -83,6 +120,7 @@ pub struct KNNClassifier< knn_algorithm: Option>, weight: Option, k: Option, + parameters: Option>, _phantom_tx: PhantomData, _phantom_x: PhantomData, _phantom_y: PhantomData, @@ -188,6 +226,7 @@ impl, Y: Array1, D: Distance, Y: Array1, D: Distance, Y: Array1, D: Distance Option<&KNNClassifierParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -720,9 +772,80 @@ 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] + 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 b49743f8..09d48bb4 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)] @@ -80,6 +117,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 +217,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 +270,16 @@ 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 +319,14 @@ impl, Y: Array1, D: Distance>> Ok(result) } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&KNNRegressorParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -344,10 +394,79 @@ 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]; + 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 d72ecdac..84a959ee 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,14 @@ impl<'a, TX: Number + RealNumber, TY: Number + Ord, X: Array2 + 'a, Y: Array f } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&SVCParameters> { + self.parameters + } } impl, Y: Array1> PartialEq @@ -1429,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 e912743b..38778617 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -280,6 +280,14 @@ impl<'a, T: Number + FloatNumber + PartialOrd, X: Array2, Y: Array1> SVR<' T::from(f).unwrap() } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&SVRParameters> { + self.parameters + } } impl, Y: Array1> PartialEq @@ -708,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 f84ae7e9..ac594f7d 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -63,9 +63,12 @@ impl, Y: Array1> fn nodes(&self) -> &Vec { self.nodes.as_ref() } - /// Get parameters, return a shared reference - fn parameters(&self) -> &BaseTreeRegressorParameters { - self.parameters.as_ref().unwrap() + /// Getter for parameters used in the model + /// + /// # Returns + /// `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 { @@ -237,7 +240,13 @@ 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, @@ -298,7 +307,11 @@ 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; } @@ -313,7 +326,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 { @@ -380,9 +393,11 @@ 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; } @@ -436,8 +451,12 @@ 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]; @@ -501,7 +520,11 @@ 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 96007677..768d2d5d 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -131,9 +131,12 @@ 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 + /// `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 { @@ -615,7 +618,13 @@ 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, @@ -701,13 +710,18 @@ impl, Y: Array1> } } - self.nodes[visitor.node].impurity = Some(impurity(&self.parameters().criterion, &count, n)); + 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)); if is_pure { return false; } - if n <= self.parameters().min_samples_split { + if n <= parameters_min_samples_split { return false; } @@ -749,10 +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"); - 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]; true_count[visitor.y[*i]] += visitor.samples[*i]; @@ -767,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(&self.parameters().criterion, &true_count, tc) - - fc as f64 / n as f64 - * impurity(&self.parameters().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() @@ -821,7 +834,11 @@ 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; @@ -1236,4 +1253,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 86b99343..88dcc2dd 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,14 @@ 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 + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&DecisionTreeRegressorParameters> { + self.parameters.as_ref() + } } #[cfg(test)] @@ -466,4 +477,46 @@ 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 a3b9bf0a..612c0fb6 100644 --- a/src/xgboost/xgb_regressor.rs +++ b/src/xgboost/xgb_regressor.rs @@ -609,6 +609,14 @@ impl, Y: Array1> XGRegres indices.truncate((population_size as f64 * subsample_ratio) as usize); indices } + + /// Getter for parameters used in the model + /// + /// # Returns + /// `Some` with the parameters used to configure the model, or `None` if unavailable. + pub fn parameters(&self) -> Option<&XGRegressorParameters> { + self.parameters.as_ref() + } } // Boilerplate implementation for the smartcore traits @@ -797,4 +805,56 @@ 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()); + } }