Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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).

## [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`.

### 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.
Expand Down
2 changes: 1 addition & 1 deletion Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
8 changes: 8 additions & 0 deletions src/ensemble/extra_trees_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> =
Expand All @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> =
Expand Down
20 changes: 11 additions & 9 deletions src/ensemble/random_forest_classifier.rs
Original file line number Diff line number Diff line change
Expand Up @@ -549,14 +549,8 @@ impl<TX: FloatNumber + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY

/// Predict OOB classes for `x`. `x` is expected to be equal to the dataset used in training.
pub fn predict_oob(&self, x: &X) -> Result<Y, Failed> {
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 {
Expand Down Expand Up @@ -584,7 +578,7 @@ impl<TX: FloatNumber + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY

Ok(result)
}
None => Err(Failed::predict(
_ => Err(Failed::predict(
"'fit' should be called before calling 'predict'",
)),
}
Expand Down Expand Up @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>> =
Expand All @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>> =
Expand Down
8 changes: 8 additions & 0 deletions src/ensemble/random_forest_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> =
Expand All @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> =
Expand Down
4 changes: 4 additions & 0 deletions src/linear/elastic_net.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> = ElasticNet::new();
Expand Down
4 changes: 4 additions & 0 deletions src/linear/lasso.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> = Lasso::new();
Expand Down
8 changes: 8 additions & 0 deletions src/linear/linear_regression.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> = LinearRegression::new();
Expand All @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> = LinearRegression::new();
Expand Down
4 changes: 4 additions & 0 deletions src/linear/logistic_regression.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, i32, DenseMatrix<f64>, Vec<i32>> =
Expand Down
4 changes: 4 additions & 0 deletions src/linear/ridge_regression.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> = RidgeRegression::new();
Expand Down
4 changes: 4 additions & 0 deletions src/naive_bayes/bernoulli.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>> = BernoulliNB::new();
Expand Down
4 changes: 4 additions & 0 deletions src/naive_bayes/categorical.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<u32, DenseMatrix<u32>, Vec<u32>> = CategoricalNB::new();
Expand Down
4 changes: 4 additions & 0 deletions src/naive_bayes/gaussian.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>> = GaussianNB::new();
Expand Down
4 changes: 4 additions & 0 deletions src/naive_bayes/multinomial.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<u32, u32, DenseMatrix<u32>, Vec<u32>> = MultinomialNB::new();
Expand Down
8 changes: 8 additions & 0 deletions src/neighbors/knn_classifier.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>, Euclidian<f64>> =
Expand All @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>, Euclidian<f64>> =
Expand Down
4 changes: 4 additions & 0 deletions src/neighbors/knn_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>, Euclidian<f64>> =
Expand Down
12 changes: 12 additions & 0 deletions src/svm/svc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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");
Expand All @@ -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");
Expand All @@ -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");
Expand Down
4 changes: 4 additions & 0 deletions src/svm/svr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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");
Expand Down
11 changes: 11 additions & 0 deletions src/tree/decision_tree_classifier.rs
Original file line number Diff line number Diff line change
Expand Up @@ -637,6 +637,8 @@ impl<TX: Number + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY>>
/// 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<Y, Failed> {
// 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'",
Expand Down Expand Up @@ -916,6 +918,7 @@ impl<TX: Number + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY>>
///
/// Returns an error if at least one row prediction process fails.
pub fn predict_proba(&self, x: &X) -> Result<DenseMatrix<f64>, 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'",
Expand Down Expand Up @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>> =
Expand All @@ -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<f64, u32, DenseMatrix<f64>, Vec<u32>> =
Expand Down
4 changes: 4 additions & 0 deletions src/tree/decision_tree_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> =
Expand Down
13 changes: 9 additions & 4 deletions src/xgboost/xgb_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -589,12 +589,13 @@ impl<TX: Number + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>> XGRegres

/// Predicts target values for the given input data.
pub fn predict(&self, data: &X) -> Result<Vec<TX>, 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);
Expand All @@ -608,7 +609,7 @@ impl<TX: Number + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>> XGRegres
.map(|p| TX::from_f64(p).unwrap())
.collect())
}
None => Err(Failed::predict(
_ => Err(Failed::predict(
"'fit' should be called before calling 'predict'",
)),
}
Expand Down Expand Up @@ -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<f64, f64, DenseMatrix<f64>, Vec<f64>> = XGRegressor::new();
Expand Down
Loading