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
12 changes: 11 additions & 1 deletion src/ensemble/base_forest_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ use std::fmt::Debug;
use serde::{Deserialize, Serialize};

use crate::error::{Failed, FailedError};
use crate::linalg::basic::arrays::MutArrayView1;
use crate::linalg::basic::arrays::{Array1, Array2};
use crate::numbers::basenum::Number;
use crate::numbers::floatnum::FloatNumber;
Expand Down Expand Up @@ -120,6 +121,14 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1
})
.transpose()?;

// Compute the order of each attribute once
let mut order: Vec<Vec<usize>> = Vec::with_capacity(num_attributes);

for i in 0..num_attributes {
let mut col_i: Vec<TX> = x.get_col(i).iterator(0).copied().collect();
order.push(col_i.argsort_mut());
}

for tree_idx in 0..parameters.n_trees {
if parameters.bootstrap {
samples = BaseForestRegressor::<TX, TY, X, Y>::sample_with_replacement(
Expand All @@ -139,7 +148,7 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1
min_samples_leaf: parameters.min_samples_leaf,
min_samples_split: parameters.min_samples_split,
seed: Some(parameters.seed.wrapping_add(tree_idx as u64)), // give each tree its own fixed seed
splitter: parameters.splitter.clone(),
splitter: parameters.splitter,
};
// Only use sample weights on base tree if not already applied during bootstrapping
let sample_weights_for_base_tree = if parameters.bootstrap {
Expand All @@ -153,6 +162,7 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1
sample_weights_for_base_tree,
samples.clone(),
mtry,
&order,
params,
)?;
trees.push(tree);
Expand Down
25 changes: 25 additions & 0 deletions src/ensemble/extra_trees_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -277,6 +277,7 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1
#[cfg(test)]
mod tests {
use super::*;
use crate::error::FailedError;
use crate::linalg::basic::matrix::DenseMatrix;
use crate::metrics::mean_squared_error;

Expand Down Expand Up @@ -470,6 +471,30 @@ 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 fit_with_empty_x_should_not_panic() {
let parameters = ExtraTreesRegressorParameters::default()
.with_n_trees(5)
.with_seed(42);
// (rows, columns, y): no rows, no columns, or both
let cases: Vec<(usize, usize, Vec<f64>)> =
vec![(0, 0, vec![]), (0, 3, vec![]), (3, 0, vec![1.0, 2.0, 3.0])];
for (nrows, ncols, y) in cases {
let x: DenseMatrix<f64> = DenseMatrix::new(nrows, ncols, vec![], false)
.expect("Construction of empty x should work");
let result = ExtraTreesRegressor::fit(&x, &y, parameters.clone());
let expected = Failed::because(
FailedError::ParametersError,
"Training data must contain at least one sample and one feature.",
);
assert_eq!(result.err(), Some(expected), "shape: ({nrows}, {ncols})");
}
}

mod sklearn_parity {
use super::*;
// sklearn parity tests.
Expand Down
25 changes: 25 additions & 0 deletions src/ensemble/random_forest_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -470,6 +470,7 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1
#[cfg(test)]
mod tests {
use super::*;
use crate::error::FailedError;
use crate::linalg::basic::matrix::DenseMatrix;
use crate::metrics::mean_absolute_error;

Expand Down Expand Up @@ -801,6 +802,30 @@ 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 fit_with_empty_x_should_not_panic() {
let parameters = RandomForestRegressorParameters::default()
.with_n_trees(5)
.with_seed(42);
// (rows, columns, y): no rows, no columns, or both
let cases: Vec<(usize, usize, Vec<f64>)> =
vec![(0, 0, vec![]), (0, 3, vec![]), (3, 0, vec![1.0, 2.0, 3.0])];
for (nrows, ncols, y) in cases {
let x: DenseMatrix<f64> = DenseMatrix::new(nrows, ncols, vec![], false)
.expect("Construction of empty x should work");
let result = RandomForestRegressor::fit(&x, &y, parameters.clone());
let expected = Failed::because(
FailedError::ParametersError,
"Training data must contain at least one sample and one feature.",
);
assert_eq!(result.err(), Some(expected), "shape: ({nrows}, {ncols})");
}
}

mod sklearn_parity {
use super::*;
// sklearn parity tests.
Expand Down
Loading
Loading