From df4c11973244b137ad250b4db1e10c7a46c0a6ac Mon Sep 17 00:00:00 2001 From: "Lorenzo (Mec-iS)" Date: Fri, 2 Oct 2026 09:07:20 +0100 Subject: [PATCH 1/2] fix: address review nits from #470 - RandomForestClassifier::predict_oob: fold the leftover is_none() guard into a single match on (&self.trees, &self.classes), matching Lasso. - XGRegressor::predict: match on (&self.parameters, &self.regressors) to remove the latent unwrap when the two fields are out of sync. - DecisionTreeClassifier: comment why nodes.is_empty() is the guard. - Add wasm_bindgen_test attributes to the new should_not_panic tests. - CHANGELOG: record the unfitted-predict fix and the MultiClassSVC error-message change. --- CHANGELOG.md | 7 +++++++ src/ensemble/extra_trees_regressor.rs | 8 ++++++++ src/ensemble/random_forest_classifier.rs | 20 +++++++++++--------- src/ensemble/random_forest_regressor.rs | 8 ++++++++ src/linear/elastic_net.rs | 4 ++++ src/linear/lasso.rs | 4 ++++ src/linear/linear_regression.rs | 8 ++++++++ src/linear/logistic_regression.rs | 4 ++++ src/linear/ridge_regression.rs | 4 ++++ src/naive_bayes/bernoulli.rs | 4 ++++ src/naive_bayes/categorical.rs | 4 ++++ src/naive_bayes/gaussian.rs | 4 ++++ src/naive_bayes/multinomial.rs | 4 ++++ src/neighbors/knn_classifier.rs | 8 ++++++++ src/neighbors/knn_regressor.rs | 4 ++++ src/svm/svc.rs | 12 ++++++++++++ src/svm/svr.rs | 4 ++++ src/tree/decision_tree_classifier.rs | 11 +++++++++++ src/tree/decision_tree_regressor.rs | 4 ++++ src/xgboost/xgb_regressor.rs | 13 +++++++++---- 20 files changed, 126 insertions(+), 13 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index ad3be010..bcc0e24c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,13 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## [Unreleased] +### Fixed +- `predict`, `predict_proba`, `predict_oob`, `predict_matrix` and `decision_function` now return `Err(Failed)` instead of panicking when `fit` has not run (#469, #470). Models return the same unfitted error even after deserialization, when the state fields are `None`. The guarded methods cover the decision trees, the random forests, the extra-trees regressor, the KNN classifier and regressor, `SVC`, `MultiClassSVC`, `SVR`, the linear models, the naive Bayes classifiers and `XGRegressor`. + +### Changed +- The error text for an unfitted `MultiClassSVC::predict` changed from "MultiClassSVC is not fitted" to the common "'fit' should be called before calling 'predict'". Code that matched the old string must be updated. + ## [0.6.15] ### Fixed - `tree`: regression tests for the tree growth fix from #464. New `min_samples_split_boundary` tests pin the split rule for `DecisionTreeClassifier` and `BaseTreeRegressor`: a node that holds exactly `min_samples_split` samples must still split, and a node with fewer samples must stay a leaf. The `full_depth` tests now also assert the tree `depth` (3), which guards the public `DecisionTreeClassifier::depth` accessor against silent regressions. Library code is unchanged. diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index 02748eea..f445e3bf 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -440,6 +440,10 @@ mod tests { } } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let forest: ExtraTreesRegressor, Vec> = @@ -451,6 +455,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_oob_without_fit_should_not_panic() { let forest: ExtraTreesRegressor, Vec> = diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index 5b43e221..ff7a4329 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -549,14 +549,8 @@ impl, Y: Array1 Result { - if self.trees.is_none() { - return Err(Failed::predict( - "'fit' should be called before calling 'predict'", - )); - } - - match &self.classes { - Some(classes) => { + match (&self.trees, &self.classes) { + (Some(_), Some(classes)) => { let (n, _) = x.shape(); let samples = match &self.samples { @@ -584,7 +578,7 @@ impl, Y: Array1 Err(Failed::predict( + _ => Err(Failed::predict( "'fit' should be called before calling 'predict'", )), } @@ -879,6 +873,10 @@ mod tests { ); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let tree: RandomForestClassifier, Vec> = @@ -890,6 +888,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_oob_without_fit_should_not_panic() { let tree: RandomForestClassifier, Vec> = diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index e56ec52f..95eeb483 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -768,6 +768,10 @@ mod tests { assert_eq!(forest, deserialized_forest); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let forest: RandomForestRegressor, Vec> = @@ -779,6 +783,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_oob_without_fit_should_not_panic() { let forest: RandomForestRegressor, Vec> = diff --git a/src/linear/elastic_net.rs b/src/linear/elastic_net.rs index 4fb5ef32..30544bcd 100644 --- a/src/linear/elastic_net.rs +++ b/src/linear/elastic_net.rs @@ -658,6 +658,10 @@ mod tests { assert_eq!(lr, deserialized_lr); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let model: ElasticNet, Vec> = ElasticNet::new(); diff --git a/src/linear/lasso.rs b/src/linear/lasso.rs index 7f1ea9d3..45493ea3 100644 --- a/src/linear/lasso.rs +++ b/src/linear/lasso.rs @@ -586,6 +586,10 @@ mod tests { assert_eq!(lr, deserialized_lr); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let model: Lasso, Vec> = Lasso::new(); diff --git a/src/linear/linear_regression.rs b/src/linear/linear_regression.rs index 56eb6d6d..7f5bf326 100644 --- a/src/linear/linear_regression.rs +++ b/src/linear/linear_regression.rs @@ -733,6 +733,10 @@ mod tests { } } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let model: LinearRegression, Vec> = LinearRegression::new(); @@ -743,6 +747,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_matrix_without_fit_should_not_panic() { let model: LinearRegression, Vec> = LinearRegression::new(); diff --git a/src/linear/logistic_regression.rs b/src/linear/logistic_regression.rs index ce42cc54..031dfd9d 100644 --- a/src/linear/logistic_regression.rs +++ b/src/linear/logistic_regression.rs @@ -1003,6 +1003,10 @@ mod tests { assert_eq!(y_hat.shape(), 52181); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let model: LogisticRegression, Vec> = diff --git a/src/linear/ridge_regression.rs b/src/linear/ridge_regression.rs index 15e07647..cdd77242 100644 --- a/src/linear/ridge_regression.rs +++ b/src/linear/ridge_regression.rs @@ -541,6 +541,10 @@ mod tests { assert_eq!(lr, deserialized_lr); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let model: RidgeRegression, Vec> = RidgeRegression::new(); diff --git a/src/naive_bayes/bernoulli.rs b/src/naive_bayes/bernoulli.rs index 2ff997df..8da68d01 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -662,6 +662,10 @@ mod tests { assert_eq!(bnb, deserialized_bnb); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let bnb: BernoulliNB, Vec> = BernoulliNB::new(); diff --git a/src/naive_bayes/categorical.rs b/src/naive_bayes/categorical.rs index c299c564..4f3f78c6 100644 --- a/src/naive_bayes/categorical.rs +++ b/src/naive_bayes/categorical.rs @@ -599,6 +599,10 @@ mod tests { assert_eq!(cnb, deserialized_cnb); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let cnb: CategoricalNB, Vec> = CategoricalNB::new(); diff --git a/src/naive_bayes/gaussian.rs b/src/naive_bayes/gaussian.rs index 6d8e6f2c..08fc31d0 100644 --- a/src/naive_bayes/gaussian.rs +++ b/src/naive_bayes/gaussian.rs @@ -475,6 +475,10 @@ mod tests { assert_eq!(gnb, deserialized_gnb); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let gnb: GaussianNB, Vec> = GaussianNB::new(); diff --git a/src/naive_bayes/multinomial.rs b/src/naive_bayes/multinomial.rs index dd130e4e..5668840c 100644 --- a/src/naive_bayes/multinomial.rs +++ b/src/naive_bayes/multinomial.rs @@ -579,6 +579,10 @@ mod tests { assert_eq!(mnb, deserialized_mnb); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let mnb: MultinomialNB, Vec> = MultinomialNB::new(); diff --git a/src/neighbors/knn_classifier.rs b/src/neighbors/knn_classifier.rs index 140ad6ef..b1c57e14 100644 --- a/src/neighbors/knn_classifier.rs +++ b/src/neighbors/knn_classifier.rs @@ -741,6 +741,10 @@ mod tests { assert_eq!(knn, deserialized_knn); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let knn: KNNClassifier, Vec, Euclidian> = @@ -752,6 +756,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_proba_without_fit_should_not_panic() { let knn: KNNClassifier, Vec, Euclidian> = diff --git a/src/neighbors/knn_regressor.rs b/src/neighbors/knn_regressor.rs index 3ea95ae1..2e8720bd 100644 --- a/src/neighbors/knn_regressor.rs +++ b/src/neighbors/knn_regressor.rs @@ -359,6 +359,10 @@ mod tests { assert_eq!(knn, deserialized_knn); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let knn: KNNRegressor, Vec, Euclidian> = diff --git a/src/svm/svc.rs b/src/svm/svc.rs index b246aa23..1052c0ff 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -1455,6 +1455,10 @@ mod tests { assert_eq!(svc, deserialized_svc); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); @@ -1465,6 +1469,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn multiclass_predict_without_fit_should_not_panic() { let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); @@ -1475,6 +1483,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn decision_function_without_fit_should_not_panic() { let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); diff --git a/src/svm/svr.rs b/src/svm/svr.rs index f6eb933b..2a325be2 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -717,6 +717,10 @@ mod tests { assert_eq!(svr, deserialized_svr); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index 9683730b..cd8ea5b2 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -637,6 +637,8 @@ impl, Y: Array1> /// Predict class value for `x`. /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &X) -> Result { + // Guard style differs from the Option-based models: `nodes` is a plain + // Vec that `fit` fills, so an empty vector marks an unfitted tree. if self.nodes.is_empty() { return Err(Failed::predict( "'fit' should be called before calling 'predict'", @@ -916,6 +918,7 @@ impl, Y: Array1> /// /// Returns an error if at least one row prediction process fails. pub fn predict_proba(&self, x: &X) -> Result, Failed> { + // Same guard as `predict`: an empty `nodes` vector means `fit` never ran. if self.nodes.is_empty() { return Err(Failed::predict( "'fit' should be called before calling 'predict'", @@ -1297,6 +1300,10 @@ mod tests { ); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let tree: DecisionTreeClassifier, Vec> = @@ -1308,6 +1315,10 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_proba_without_fit_should_not_panic() { let tree: DecisionTreeClassifier, Vec> = diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index 89646637..839cfdac 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -614,6 +614,10 @@ mod tests { } } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let tree: DecisionTreeRegressor, Vec> = diff --git a/src/xgboost/xgb_regressor.rs b/src/xgboost/xgb_regressor.rs index 8a78d7e3..86901339 100644 --- a/src/xgboost/xgb_regressor.rs +++ b/src/xgboost/xgb_regressor.rs @@ -589,12 +589,13 @@ impl, Y: Array1> XGRegres /// Predicts target values for the given input data. pub fn predict(&self, data: &X) -> Result, Failed> { - match &self.parameters { - Some(parameters) => { + // Match on both fields: after deserialization, 'parameters' and + // 'regressors' could be out of sync, so a single check must cover them. + match (&self.parameters, &self.regressors) { + (Some(parameters), Some(regressors)) => { let (n_samples, _) = data.shape(); let mut predictions = vec![parameters.base_score; n_samples]; - let regressors = self.regressors.as_ref().unwrap(); for regressor in regressors.iter() { let corrections = regressor.predict(data); @@ -608,7 +609,7 @@ impl, Y: Array1> XGRegres .map(|p| TX::from_f64(p).unwrap()) .collect()) } - None => Err(Failed::predict( + _ => Err(Failed::predict( "'fit' should be called before calling 'predict'", )), } @@ -940,6 +941,10 @@ mod tests { assert_eq!(predictions.len(), 4); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] #[test] fn predict_without_fit_should_not_panic() { let tree: XGRegressor, Vec> = XGRegressor::new(); From b6644c1720baaa0a777996eee4fc17549f115a41 Mon Sep 17 00:00:00 2001 From: "Lorenzo (Mec-iS)" Date: Fri, 2 Oct 2026 09:11:29 +0100 Subject: [PATCH 2/2] chore: bump patch 0.6.15 -> 0.6.16 and stamp the changelog entry --- CHANGELOG.md | 2 +- Cargo.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index bcc0e24c..c250a5db 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,7 +4,7 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). -## [Unreleased] +## [0.6.16] ### Fixed - `predict`, `predict_proba`, `predict_oob`, `predict_matrix` and `decision_function` now return `Err(Failed)` instead of panicking when `fit` has not run (#469, #470). Models return the same unfitted error even after deserialization, when the state fields are `None`. The guarded methods cover the decision trees, the random forests, the extra-trees regressor, the KNN classifier and regressor, `SVC`, `MultiClassSVC`, `SVR`, the linear models, the naive Bayes classifiers and `XGRegressor`. diff --git a/Cargo.toml b/Cargo.toml index 017358be..4ec3a4b5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -2,7 +2,7 @@ name = "smartcore" description = "Machine Learning in Rust." homepage = "https://smartcorelib.github.io/" -version = "0.6.15" +version = "0.6.16" authors = ["smartcore Developers"] edition = "2024" rust-version = "1.85"