diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index 5ae33ebe..02748eea 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -255,14 +255,22 @@ impl, Y: Array1 /// Predict class for `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &X) -> Result { - let forest_regressor = self.forest_regressor.as_ref().unwrap(); - forest_regressor.predict(x) + match &self.forest_regressor { + Some(forest) => forest.predict(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// 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 { - let forest_regressor = self.forest_regressor.as_ref().unwrap(); - forest_regressor.predict_oob(x) + match &self.forest_regressor { + Some(forest) => forest.predict_oob(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } } @@ -431,4 +439,26 @@ mod tests { ); } } + + #[test] + fn predict_without_fit_should_not_panic() { + let forest: ExtraTreesRegressor, Vec> = + ExtraTreesRegressor::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = forest.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn predict_oob_without_fit_should_not_panic() { + let forest: ExtraTreesRegressor, Vec> = + ExtraTreesRegressor::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = forest.predict_oob(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index 676450d2..5b43e221 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -519,18 +519,22 @@ impl, Y: Array1 Result { - let mut result = Y::zeros(x.shape().0); + match &self.classes { + Some(classes) => { + let mut result = Y::zeros(x.shape().0); - let (n, _) = x.shape(); + let (n, _) = x.shape(); - for i in 0..n { - result.set( - i, - self.classes.as_ref().unwrap()[self.predict_for_row(x, i)], - ); - } + for i in 0..n { + result.set(i, classes[self.predict_for_row(x, i)]); + } - Ok(result) + Ok(result) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } fn predict_for_row(&self, x: &X, row: usize) -> usize { @@ -545,35 +549,45 @@ impl, Y: Array1 Result { - let (n, _) = x.shape(); - - let samples = match &self.samples { - Some(s) => s, - None => { - return Err(Failed::because( - FailedError::PredictFailed, - "Need samples=true for OOB predictions.", - )); - } - }; - - if samples[0].len() != n { - return Err(Failed::because( - FailedError::PredictFailed, - "Prediction matrix must match matrix used in training for OOB predictions.", + if self.trees.is_none() { + return Err(Failed::predict( + "'fit' should be called before calling 'predict'", )); } - let mut result = Y::zeros(n); + match &self.classes { + Some(classes) => { + let (n, _) = x.shape(); + + let samples = match &self.samples { + Some(s) => s, + None => { + return Err(Failed::because( + FailedError::PredictFailed, + "Need samples=true for OOB predictions.", + )); + } + }; + + if samples[0].len() != n { + return Err(Failed::because( + FailedError::PredictFailed, + "Prediction matrix must match matrix used in training for OOB predictions.", + )); + } - for i in 0..n { - result.set( - i, - self.classes.as_ref().unwrap()[self.predict_for_row_oob(x, i)], - ); - } + let mut result = Y::zeros(n); - Ok(result) + for i in 0..n { + result.set(i, classes[self.predict_for_row_oob(x, i)]); + } + + Ok(result) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } fn predict_for_row_oob(&self, x: &X, row: usize) -> usize { @@ -865,6 +879,28 @@ mod tests { ); } + #[test] + fn predict_without_fit_should_not_panic() { + let tree: RandomForestClassifier, Vec> = + RandomForestClassifier::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = tree.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn predict_oob_without_fit_should_not_panic() { + let tree: RandomForestClassifier, Vec> = + RandomForestClassifier::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = tree.predict_oob(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index d04fa760..e56ec52f 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -448,14 +448,22 @@ impl, Y: Array1 /// Predict class for `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &X) -> Result { - let forest_regressor = self.forest_regressor.as_ref().unwrap(); - forest_regressor.predict(x) + match &self.forest_regressor { + Some(forest) => forest.predict(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// 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 { - let forest_regressor = self.forest_regressor.as_ref().unwrap(); - forest_regressor.predict_oob(x) + match &self.forest_regressor { + Some(forest) => forest.predict_oob(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } } @@ -759,4 +767,26 @@ mod tests { assert_eq!(forest, deserialized_forest); } + + #[test] + fn predict_without_fit_should_not_panic() { + let forest: RandomForestRegressor, Vec> = + RandomForestRegressor::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = forest.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn predict_oob_without_fit_should_not_panic() { + let forest: RandomForestRegressor, Vec> = + RandomForestRegressor::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = forest.predict_oob(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/linear/elastic_net.rs b/src/linear/elastic_net.rs index 804b9800..4fb5ef32 100644 --- a/src/linear/elastic_net.rs +++ b/src/linear/elastic_net.rs @@ -395,14 +395,21 @@ impl, Y: Array1> /// Predict target values from `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &X) -> Result { - let (nrows, _) = x.shape(); - let mut y_hat = x.matmul(self.coefficients.as_ref().unwrap()); - let bias = X::fill(nrows, 1, self.intercept.unwrap()); - y_hat.add_mut(&bias); - Ok(Y::from_iterator( - y_hat.iterator(0).map(|&v| TY::from(v).unwrap()), - nrows, - )) + match (&self.coefficients, &self.intercept) { + (Some(coefficients), Some(intercept)) => { + let (nrows, _) = x.shape(); + let mut y_hat = x.matmul(coefficients); + let bias = X::fill(nrows, 1, *intercept); + y_hat.add_mut(&bias); + Ok(Y::from_iterator( + y_hat.iterator(0).map(|&v| TY::from(v).unwrap()), + nrows, + )) + } + (_, _) => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Get estimates regression coefficients @@ -650,4 +657,14 @@ mod tests { assert_eq!(lr, deserialized_lr); } + + #[test] + fn predict_without_fit_should_not_panic() { + let model: ElasticNet, Vec> = ElasticNet::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = model.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/linear/lasso.rs b/src/linear/lasso.rs index c427bb41..7f1ea9d3 100644 --- a/src/linear/lasso.rs +++ b/src/linear/lasso.rs @@ -362,14 +362,21 @@ impl, Y: Array1> Las /// Predict target values from `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &X) -> Result { - let (nrows, _) = x.shape(); - let mut y_hat = x.matmul(self.coefficients()); - let bias = X::fill(nrows, 1, self.intercept.unwrap()); - y_hat.add_mut(&bias); - Ok(Y::from_iterator( - y_hat.iterator(0).map(|&v| TY::from(v).unwrap()), - nrows, - )) + match (&self.coefficients, &self.intercept) { + (Some(coefficients), Some(intercept)) => { + let (nrows, _) = x.shape(); + let mut y_hat = x.matmul(coefficients); + let bias = X::fill(nrows, 1, *intercept); + y_hat.add_mut(&bias); + Ok(Y::from_iterator( + y_hat.iterator(0).map(|&v| TY::from(v).unwrap()), + nrows, + )) + } + (_, _) => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Get estimates regression coefficients @@ -578,4 +585,14 @@ mod tests { assert_eq!(lr, deserialized_lr); } + + #[test] + fn predict_without_fit_should_not_panic() { + let model: Lasso, Vec> = Lasso::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = model.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/linear/linear_regression.rs b/src/linear/linear_regression.rs index bb541878..56eb6d6d 100644 --- a/src/linear/linear_regression.rs +++ b/src/linear/linear_regression.rs @@ -335,23 +335,29 @@ impl< /// Predict target values from `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict_matrix(&self, x: &X) -> Result { - let (nrows, _) = x.shape(); + match (&self.coefficients, &self.intercept) { + (Some(coefficients), Some(intercept)) => { + let (nrows, _) = x.shape(); - let intercept = self.intercept_matrix(); - let (_, num_targets) = intercept.shape(); + let (_, num_targets) = intercept.shape(); - let mut y_hat = x.matmul(self.coefficients()); + let mut y_hat = x.matmul(coefficients); - // Tile the 1xK intercept across all rows, then add in one pass - let bias = X::from_iterator( - (0..nrows).flat_map(|_| intercept.iterator(0).copied()), - nrows, - num_targets, - 0, - ); - y_hat.add_mut(&bias); + // Tile the 1xK intercept across all rows, then add in one pass + let bias = X::from_iterator( + (0..nrows).flat_map(|_| intercept.iterator(0).copied()), + nrows, + num_targets, + 0, + ); + y_hat.add_mut(&bias); - Ok(y_hat) + Ok(y_hat) + } + (_, _) => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Get estimates regression coefficients @@ -726,4 +732,24 @@ mod tests { assert!((*model.intercept() - 1.0).abs() < 1e-8); } } + + #[test] + fn predict_without_fit_should_not_panic() { + let model: LinearRegression, Vec> = LinearRegression::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = model.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn predict_matrix_without_fit_should_not_panic() { + let model: LinearRegression, Vec> = LinearRegression::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = model.predict_matrix(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/linear/logistic_regression.rs b/src/linear/logistic_regression.rs index 92ea277a..ce42cc54 100644 --- a/src/linear/logistic_regression.rs +++ b/src/linear/logistic_regression.rs @@ -503,32 +503,39 @@ impl, Y: /// Predict class labels for samples in `x`. /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &X) -> Result { - let n = x.shape().0; - let mut result = Y::zeros(n); - if self.num_classes == 2 { - let y_hat = x.ab(false, self.coefficients(), true); - let intercept = *self.intercept().get((0, 0)); - for (i, y_hat_i) in y_hat.iterator(0).enumerate().take(n) { - result.set( - i, - self.classes()[usize::from( - RealNumber::sigmoid(*y_hat_i + intercept) > RealNumber::half(), - )], - ); - } - } else { - let mut y_hat = x.matmul(&self.coefficients().transpose()); - for r in 0..n { - for c in 0..self.num_classes { - y_hat.set((r, c), *y_hat.get((r, c)) + *self.intercept().get((c, 0))); + match (&self.coefficients, &self.intercept, &self.classes) { + (Some(coefficients), Some(intercept), Some(classes)) => { + let n = x.shape().0; + let mut result = Y::zeros(n); + if self.num_classes == 2 { + let y_hat = x.ab(false, coefficients, true); + let intercept = *intercept.get((0, 0)); + for (i, y_hat_i) in y_hat.iterator(0).enumerate().take(n) { + result.set( + i, + classes[usize::from( + RealNumber::sigmoid(*y_hat_i + intercept) > RealNumber::half(), + )], + ); + } + } else { + let mut y_hat = x.matmul(&coefficients.transpose()); + for r in 0..n { + for c in 0..self.num_classes { + y_hat.set((r, c), *y_hat.get((r, c)) + *intercept.get((c, 0))); + } + } + let class_idxs = y_hat.argmax(1); + for (i, class_i) in class_idxs.iter().enumerate().take(n) { + result.set(i, classes[*class_i]); + } } + Ok(result) } - let class_idxs = y_hat.argmax(1); - for (i, class_i) in class_idxs.iter().enumerate().take(n) { - result.set(i, self.classes()[*class_i]); - } + (_, _, _) => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), } - Ok(result) } /// Get estimates regression coefficients, this create a sharable reference @@ -995,4 +1002,15 @@ mod tests { assert_eq!(y_hat.shape(), 52181); } + + #[test] + fn predict_without_fit_should_not_panic() { + let model: LogisticRegression, Vec> = + LogisticRegression::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = model.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/linear/ridge_regression.rs b/src/linear/ridge_regression.rs index f1015240..15e07647 100644 --- a/src/linear/ridge_regression.rs +++ b/src/linear/ridge_regression.rs @@ -393,13 +393,20 @@ impl< /// Predict target values from `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &X) -> Result { - let (nrows, _) = x.shape(); - let mut y_hat = x.matmul(self.coefficients()); - y_hat.add_mut(&X::fill(nrows, 1, self.intercept.unwrap())); - Ok(Y::from_iterator( - y_hat.iterator(0).map(|&v| TY::from(v).unwrap()), - nrows, - )) + match (&self.coefficients, &self.intercept) { + (Some(coefficients), Some(intercept)) => { + let (nrows, _) = x.shape(); + let mut y_hat = x.matmul(coefficients); + y_hat.add_mut(&X::fill(nrows, 1, *intercept)); + Ok(Y::from_iterator( + y_hat.iterator(0).map(|&v| TY::from(v).unwrap()), + nrows, + )) + } + (_, _) => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Get estimates regression coefficients @@ -533,4 +540,14 @@ mod tests { assert_eq!(lr, deserialized_lr); } + + #[test] + fn predict_without_fit_should_not_panic() { + let model: RidgeRegression, Vec> = RidgeRegression::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = model.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/naive_bayes/bernoulli.rs b/src/naive_bayes/bernoulli.rs index 1a552d42..2ff997df 100644 --- a/src/naive_bayes/bernoulli.rs +++ b/src/naive_bayes/bernoulli.rs @@ -431,13 +431,17 @@ impl, Y: Arr /// /// Returns a vector of size N with class estimates. pub fn predict(&self, x: &X) -> Result { - if let Some(threshold) = self.binarize { - self.inner - .as_ref() - .unwrap() - .predict(&Self::binarize(x, threshold)) - } else { - self.inner.as_ref().unwrap().predict(x) + match &self.inner { + Some(inner) => { + if let Some(threshold) = self.binarize { + inner.predict(&Self::binarize(x, threshold)) + } else { + inner.predict(x) + } + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), } } @@ -657,4 +661,14 @@ mod tests { assert_eq!(bnb, deserialized_bnb); } + + #[test] + fn predict_without_fit_should_not_panic() { + let bnb: BernoulliNB, Vec> = BernoulliNB::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = bnb.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/naive_bayes/categorical.rs b/src/naive_bayes/categorical.rs index b5882206..c299c564 100644 --- a/src/naive_bayes/categorical.rs +++ b/src/naive_bayes/categorical.rs @@ -380,7 +380,12 @@ impl, Y: Array1> CategoricalNB { /// /// Returns a vector of size N with class estimates. pub fn predict(&self, x: &X) -> Result { - self.inner.as_ref().unwrap().predict(x) + match &self.inner { + Some(inner) => inner.predict(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Class labels known to the classifier. @@ -593,4 +598,14 @@ mod tests { assert_eq!(cnb, deserialized_cnb); } + + #[test] + fn predict_without_fit_should_not_panic() { + let cnb: CategoricalNB, Vec> = CategoricalNB::new(); + let x = DenseMatrix::from_2d_array(&[&[1u32]]).expect("Construction of x should work"); + let yhat = cnb.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/naive_bayes/gaussian.rs b/src/naive_bayes/gaussian.rs index a5a96b88..6d8e6f2c 100644 --- a/src/naive_bayes/gaussian.rs +++ b/src/naive_bayes/gaussian.rs @@ -320,7 +320,12 @@ impl, Y: Arr /// /// Returns a vector of size N with class estimates. pub fn predict(&self, x: &X) -> Result { - self.inner.as_ref().unwrap().predict(x) + match &self.inner { + Some(inner) => inner.predict(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Class labels known to the classifier. @@ -469,4 +474,14 @@ mod tests { assert_eq!(gnb, deserialized_gnb); } + + #[test] + fn predict_without_fit_should_not_panic() { + let gnb: GaussianNB, Vec> = GaussianNB::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = gnb.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/naive_bayes/multinomial.rs b/src/naive_bayes/multinomial.rs index 0e21c75c..dd130e4e 100644 --- a/src/naive_bayes/multinomial.rs +++ b/src/naive_bayes/multinomial.rs @@ -362,7 +362,12 @@ impl, Y: Array /// /// Returns a vector of size N with class estimates. pub fn predict(&self, x: &X) -> Result { - self.inner.as_ref().unwrap().predict(x) + match &self.inner { + Some(inner) => inner.predict(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Class labels known to the classifier. @@ -573,4 +578,14 @@ mod tests { assert_eq!(mnb, deserialized_mnb); } + + #[test] + fn predict_without_fit_should_not_panic() { + let mnb: MultinomialNB, Vec> = MultinomialNB::new(); + let x = DenseMatrix::from_2d_array(&[&[1u32]]).expect("Construction of x should work"); + let yhat = mnb.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/neighbors/knn_classifier.rs b/src/neighbors/knn_classifier.rs index b7970f7d..140ad6ef 100644 --- a/src/neighbors/knn_classifier.rs +++ b/src/neighbors/knn_classifier.rs @@ -265,17 +265,24 @@ impl, Y: Array1, D: Distance Result { - let mut result = Y::zeros(x.shape().0); - - let mut row_vec = vec![TX::zero(); x.shape().1]; - for (i, row) in x.row_iter().enumerate() { - row.iterator(0) - .zip(row_vec.iter_mut()) - .for_each(|(&s, v)| *v = s); - result.set(i, self.classes()[self.predict_for_row(&row_vec)?]); - } + match &self.knn_algorithm { + Some(_) => { + let mut result = Y::zeros(x.shape().0); + + let mut row_vec = vec![TX::zero(); x.shape().1]; + for (i, row) in x.row_iter().enumerate() { + row.iterator(0) + .zip(row_vec.iter_mut()) + .for_each(|(&s, v)| *v = s); + result.set(i, self.classes()[self.predict_for_row(&row_vec)?]); + } - Ok(result) + Ok(result) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Compute class probabilities for a single row. All the rest functions will use it @@ -334,16 +341,23 @@ impl, Y: Array1, D: Distance Result>, Failed> { - let mut result = Vec::with_capacity(x.shape().0); - let mut row_vec = vec![TX::zero(); x.shape().1]; - for row in x.row_iter() { - row.iterator(0) - .zip(row_vec.iter_mut()) - .for_each(|(&s, v)| *v = s); - result.push(self.predict_proba_for_row(&row_vec)?); - } + match &self.knn_algorithm { + Some(_) => { + let mut result = Vec::with_capacity(x.shape().0); + let mut row_vec = vec![TX::zero(); x.shape().1]; + for row in x.row_iter() { + row.iterator(0) + .zip(row_vec.iter_mut()) + .for_each(|(&s, v)| *v = s); + result.push(self.predict_proba_for_row(&row_vec)?); + } - Ok(result) + Ok(result) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } } @@ -726,4 +740,26 @@ mod tests { assert_eq!(knn, deserialized_knn); } + + #[test] + fn predict_without_fit_should_not_panic() { + let knn: KNNClassifier, Vec, Euclidian> = + KNNClassifier::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = knn.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn predict_proba_without_fit_should_not_panic() { + let knn: KNNClassifier, Vec, Euclidian> = + KNNClassifier::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = knn.predict_proba(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/neighbors/knn_regressor.rs b/src/neighbors/knn_regressor.rs index 9edb5f95..3ea95ae1 100644 --- a/src/neighbors/knn_regressor.rs +++ b/src/neighbors/knn_regressor.rs @@ -250,17 +250,24 @@ impl, Y: Array1, D: Distance>> /// /// Returns a vector of size N with estimates. pub fn predict(&self, x: &X) -> Result { - let mut result = Y::zeros(x.shape().0); - - let mut row_vec = vec![TX::zero(); x.shape().1]; - for (i, row) in x.row_iter().enumerate() { - row.iterator(0) - .zip(row_vec.iter_mut()) - .for_each(|(&s, v)| *v = s); - result.set(i, self.predict_for_row(&row_vec)?); - } + match &self.knn_algorithm { + Some(_) => { + let mut result = Y::zeros(x.shape().0); + + let mut row_vec = vec![TX::zero(); x.shape().1]; + for (i, row) in x.row_iter().enumerate() { + row.iterator(0) + .zip(row_vec.iter_mut()) + .for_each(|(&s, v)| *v = s); + result.set(i, self.predict_for_row(&row_vec)?); + } - Ok(result) + Ok(result) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } fn predict_for_row(&self, row: &Vec) -> Result { @@ -351,4 +358,15 @@ mod tests { assert_eq!(knn, deserialized_knn); } + + #[test] + fn predict_without_fit_should_not_panic() { + let knn: KNNRegressor, Vec, Euclidian> = + KNNRegressor::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = knn.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 2a4ea13b..b246aa23 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -188,7 +188,7 @@ impl<'a, TX: Number + RealNumber, TY: Number + Ord, X: Array2, Y: Array1 /// * `x` - A reference to the input features (2D array). /// * `y` - A reference to the target labels (1D array). /// * `parameters` - A reference to the `SVCParameters` controlling the SVM training for each individual binary classifier. - /// + /// /// /// # Returns /// A `Result` indicating success (`MultiClassSVC`) or failure (`Failed`). @@ -239,46 +239,49 @@ impl<'a, TX: Number + RealNumber, TY: Number + Ord, X: Array2, Y: Array1 /// A `Result` containing a `Vec` of predicted class labels (`TX`) or a `Failed` error. /// pub fn predict(&self, x: &X) -> Result, Failed> { - // Initialize a HashMap for each data point to store votes for each class - let mut polls = vec![HashMap::new(); x.shape().0]; - // Retrieve the trained binary classifiers. The field is None only before - // fit() has run, e.g. on a deserialized model. - let classifiers = self.classifiers.as_ref().ok_or_else(|| { - Failed::because(FailedError::PredictFailed, "MultiClassSVC is not fitted") - })?; - - // Iterate through each binary classifier - for svc in classifiers { - let predictions = svc.predict(x)?; // call SVC::predict for each binary classifier - - // For each prediction from the current binary classifier - for (j, prediction) in predictions.iter().enumerate() { - let prediction = prediction.to_i32().unwrap(); - let poll = polls.get_mut(j).unwrap(); // Get the poll for the current data point - // Increment the vote for the predicted class - if let Some(count) = poll.get_mut(&prediction) { - *count += 1 - } else { - poll.insert(prediction, 1); + match &self.classifiers { + Some(classifiers) => { + // Initialize a HashMap for each data point to store votes for each class + let mut polls = vec![HashMap::new(); x.shape().0]; + + // Iterate through each binary classifier + for svc in classifiers { + let predictions = svc.predict(x)?; // call SVC::predict for each binary classifier + + // For each prediction from the current binary classifier + for (j, prediction) in predictions.iter().enumerate() { + let prediction = prediction.to_i32().unwrap(); + let poll = polls.get_mut(j).unwrap(); // Get the poll for the current data point + // Increment the vote for the predicted class + if let Some(count) = poll.get_mut(&prediction) { + *count += 1 + } else { + poll.insert(prediction, 1); + } + } } + + // Determine the final prediction for each data point based on majority vote. + // A poll stays empty when fit() ran on data with fewer than two classes. + polls + .iter() + .map(|v| { + // Find the class with the maximum votes for each data point + let (class, _) = + v.iter().max_by_key(|(_, class)| *class).ok_or_else(|| { + Failed::because( + FailedError::PredictFailed, + "MultiClassSVC must be fitted on at least two classes", + ) + })?; + Ok(TX::from(*class).unwrap()) + }) + .collect() } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), } - - // Determine the final prediction for each data point based on majority vote. - // A poll stays empty when fit() ran on data with fewer than two classes. - polls - .iter() - .map(|v| { - // Find the class with the maximum votes for each data point - let (class, _) = v.iter().max_by_key(|(_, class)| *class).ok_or_else(|| { - Failed::because( - FailedError::PredictFailed, - "MultiClassSVC must be fitted on at least two classes", - ) - })?; - Ok(TX::from(*class).unwrap()) - }) - .collect() } } @@ -560,35 +563,49 @@ impl<'a, TX: Number + RealNumber, TY: Number + Ord, X: Array2 + 'a, Y: Array /// Predicts estimated class labels from `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &'a X) -> Result, Failed> { - let mut y_hat: Vec = self.decision_function(x)?; + match &self.classes { + Some(classes) => { + let mut y_hat: Vec = self.decision_function(x)?; - for i in 0..y_hat.len() { - let cls_idx = match *y_hat.get(i) > TX::zero() { - false => TX::from(self.classes.as_ref().unwrap().0).unwrap(), - true => TX::from(self.classes.as_ref().unwrap().1).unwrap(), - }; + for i in 0..y_hat.len() { + let cls_idx = match *y_hat.get(i) > TX::zero() { + false => TX::from(classes.0).unwrap(), + true => TX::from(classes.1).unwrap(), + }; - y_hat.set(i, cls_idx); - } + y_hat.set(i, cls_idx); + } - Ok(y_hat) + Ok(y_hat) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } /// Evaluates the decision function for the rows in `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn decision_function(&self, x: &'a X) -> Result, Failed> { - let (n, _) = x.shape(); - let mut y_hat: Vec = Array1::zeros(n); - - let mut row = Vec::with_capacity(n); - for i in 0..n { - row.clear(); - row.extend(x.get_row(i).iterator(0).copied()); - let row_pred: TX = self.predict_for_row(&row); - y_hat.set(i, row_pred); - } + match &self.classes { + Some(_) => { + let (n, _) = x.shape(); + let mut y_hat: Vec = Array1::zeros(n); + + let mut row = Vec::with_capacity(n); + for i in 0..n { + row.clear(); + row.extend(x.get_row(i).iterator(0).copied()); + let row_pred: TX = self.predict_for_row(&row); + y_hat.set(i, row_pred); + } - Ok(y_hat) + Ok(y_hat) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'decision_function'", + )), + } } fn predict_for_row(&self, x: &[TX]) -> TX { @@ -1437,4 +1454,34 @@ mod tests { assert_eq!(svc, deserialized_svc); } + + #[test] + fn predict_without_fit_should_not_panic() { + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let svc: SVC<'_, f64, i32, DenseMatrix, Vec> = SVC::new(); + let yhat = svc.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn multiclass_predict_without_fit_should_not_panic() { + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let svc: MultiClassSVC<'_, f64, i32, DenseMatrix, Vec> = MultiClassSVC::new(); + let yhat = svc.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn decision_function_without_fit_should_not_panic() { + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let svc: SVC<'_, f64, i32, DenseMatrix, Vec> = SVC::new(); + let dec_fn = svc.decision_function(&x); + assert!(dec_fn.is_err()); + let msg = "'fit' should be called before calling 'decision_function'"; + assert_eq!(dec_fn.err(), Some(Failed::predict(msg))); + } } diff --git a/src/svm/svr.rs b/src/svm/svr.rs index c8c1267a..f6eb933b 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -246,18 +246,25 @@ impl<'a, T: Number + FloatNumber + PartialOrd, X: Array2, Y: Array1> SVR<' /// Predict target values from `x` /// * `x` - _KxM_ data where _K_ is number of observations and _M_ is number of features. pub fn predict(&self, x: &'a X) -> Result, Failed> { - let (n, _) = x.shape(); + match &self.instances { + Some(_) => { + let (n, _) = x.shape(); - let mut y_hat: Vec = Vec::::zeros(n); + let mut y_hat: Vec = Vec::::zeros(n); - let mut x_i = Vec::with_capacity(n); - for i in 0..n { - x_i.clear(); - x_i.extend(x.get_row(i).iterator(0).copied()); - y_hat.set(i, self.predict_for_row(&x_i)); - } + let mut x_i = Vec::with_capacity(n); + for i in 0..n { + x_i.clear(); + x_i.extend(x.get_row(i).iterator(0).copied()); + y_hat.set(i, self.predict_for_row(&x_i)); + } - Ok(y_hat) + Ok(y_hat) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } pub(crate) fn predict_for_row(&self, x: &[T]) -> T { @@ -709,4 +716,14 @@ mod tests { assert_eq!(svr, deserialized_svr); } + + #[test] + fn predict_without_fit_should_not_panic() { + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let svr: SVR<'_, f64, DenseMatrix, Vec> = SVR::new(); + let yhat = svr.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } } diff --git a/src/tree/decision_tree_classifier.rs b/src/tree/decision_tree_classifier.rs index 7e2bf959..9683730b 100644 --- a/src/tree/decision_tree_classifier.rs +++ b/src/tree/decision_tree_classifier.rs @@ -637,6 +637,11 @@ 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 { + if self.nodes.is_empty() { + return Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )); + } let mut result = Y::zeros(x.shape().0); let (n, _) = x.shape(); @@ -911,6 +916,11 @@ impl, Y: Array1> /// /// Returns an error if at least one row prediction process fails. pub fn predict_proba(&self, x: &X) -> Result, Failed> { + if self.nodes.is_empty() { + return Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )); + } let (n_samples, _) = x.shape(); let n_classes = self.classes().len(); let mut result = DenseMatrix::::zeros(n_samples, n_classes); @@ -1287,6 +1297,28 @@ mod tests { ); } + #[test] + fn predict_without_fit_should_not_panic() { + let tree: DecisionTreeClassifier, Vec> = + DecisionTreeClassifier::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = tree.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + + #[test] + fn predict_proba_without_fit_should_not_panic() { + let tree: DecisionTreeClassifier, Vec> = + DecisionTreeClassifier::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = tree.predict_proba(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index d1c91ec8..89646637 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -358,7 +358,12 @@ impl, Y: Array1> /// Predict regression 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 { - self.tree_regressor.as_ref().unwrap().predict(x) + match &self.tree_regressor { + Some(tree) => tree.predict(x), + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), + } } } @@ -609,6 +614,17 @@ mod tests { } } + #[test] + fn predict_without_fit_should_not_panic() { + let tree: DecisionTreeRegressor, Vec> = + DecisionTreeRegressor::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = tree.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test diff --git a/src/xgboost/xgb_regressor.rs b/src/xgboost/xgb_regressor.rs index e71c83db..8a78d7e3 100644 --- a/src/xgboost/xgb_regressor.rs +++ b/src/xgboost/xgb_regressor.rs @@ -589,23 +589,29 @@ impl, Y: Array1> XGRegres /// Predicts target values for the given input data. pub fn predict(&self, data: &X) -> Result, Failed> { - let (n_samples, _) = data.shape(); - - let parameters = self.parameters.as_ref().unwrap(); - 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); - predictions = zip(predictions, corrections) - .map(|(pred, correction)| pred + (parameters.learning_rate * correction)) - .collect(); + match &self.parameters { + Some(parameters) => { + 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); + predictions = zip(predictions, corrections) + .map(|(pred, correction)| pred + (parameters.learning_rate * correction)) + .collect(); + } + + Ok(predictions + .into_iter() + .map(|p| TX::from_f64(p).unwrap()) + .collect()) + } + None => Err(Failed::predict( + "'fit' should be called before calling 'predict'", + )), } - - Ok(predictions - .into_iter() - .map(|p| TX::from_f64(p).unwrap()) - .collect()) } /// Creates a random sample of indices without replacement. @@ -933,4 +939,14 @@ mod tests { let predictions = predict_result.unwrap(); assert_eq!(predictions.len(), 4); } + + #[test] + fn predict_without_fit_should_not_panic() { + let tree: XGRegressor, Vec> = XGRegressor::new(); + let x = DenseMatrix::from_2d_array(&[&[1.0f64]]).expect("Construction of x should work"); + let yhat = tree.predict(&x); + assert!(yhat.is_err()); + let msg = "'fit' should be called before calling 'predict'"; + assert_eq!(yhat.err(), Some(Failed::predict(msg))); + } }