Skip to content
Open
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
38 changes: 34 additions & 4 deletions src/ensemble/extra_trees_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -255,14 +255,22 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, 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<Y, Failed> {
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<Y, Failed> {
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'",
)),
}
}
}

Expand Down Expand Up @@ -431,4 +439,26 @@ mod tests {
);
}
}

#[test]
fn predict_without_fit_should_not_panic() {
let forest: ExtraTreesRegressor<f64, f64, DenseMatrix<f64>, Vec<f64>> =
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<f64, f64, DenseMatrix<f64>, Vec<f64>> =
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)));
}
}
102 changes: 69 additions & 33 deletions src/ensemble/random_forest_classifier.rs
Original file line number Diff line number Diff line change
Expand Up @@ -519,18 +519,22 @@ impl<TX: FloatNumber + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY
/// 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<Y, Failed> {
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 {
Expand All @@ -545,35 +549,45 @@ 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> {
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 {
Expand Down Expand Up @@ -865,6 +879,28 @@ mod tests {
);
}

#[test]
fn predict_without_fit_should_not_panic() {
let tree: RandomForestClassifier<f64, u32, DenseMatrix<f64>, Vec<u32>> =
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<f64, u32, DenseMatrix<f64>, Vec<u32>> =
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
Expand Down
38 changes: 34 additions & 4 deletions src/ensemble/random_forest_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -448,14 +448,22 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, 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<Y, Failed> {
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<Y, Failed> {
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'",
)),
}
}
}

Expand Down Expand Up @@ -759,4 +767,26 @@ mod tests {

assert_eq!(forest, deserialized_forest);
}

#[test]
fn predict_without_fit_should_not_panic() {
let forest: RandomForestRegressor<f64, f64, DenseMatrix<f64>, Vec<f64>> =
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<f64, f64, DenseMatrix<f64>, Vec<f64>> =
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)));
}
}
33 changes: 25 additions & 8 deletions src/linear/elastic_net.rs
Original file line number Diff line number Diff line change
Expand Up @@ -395,14 +395,21 @@ impl<TX: FloatNumber + RealNumber, TY: Number, X: Array2<TX>, Y: Array1<TY>>
/// 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<Y, Failed> {
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
Expand Down Expand Up @@ -650,4 +657,14 @@ mod tests {

assert_eq!(lr, deserialized_lr);
}

#[test]
fn predict_without_fit_should_not_panic() {
let model: ElasticNet<f64, f64, DenseMatrix<f64>, Vec<f64>> = 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)));
}
}
39 changes: 31 additions & 8 deletions src/linear/lasso.rs
Original file line number Diff line number Diff line change
Expand Up @@ -362,14 +362,27 @@ impl<TX: FloatNumber + RealNumber, TY: Number, X: Array2<TX>, Y: Array1<TY>> 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<Y, Failed> {
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,
))
if self.coefficients.is_none() {
return Err(Failed::predict(
"'fit' should be called before calling 'predict'",
));
}

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
Expand Down Expand Up @@ -578,4 +591,14 @@ mod tests {

assert_eq!(lr, deserialized_lr);
}

#[test]
fn predict_without_fit_should_not_panic() {
let model: Lasso<f64, f64, DenseMatrix<f64>, Vec<f64>> = 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)));
}
}
Loading
Loading