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
43 changes: 34 additions & 9 deletions src/tree/base_tree_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -68,10 +68,6 @@ impl<TX: Number + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>>
fn parameters(&self) -> &BaseTreeRegressorParameters {
self.parameters.as_ref().unwrap()
}
/// Get estimate of intercept, return value
fn depth(&self) -> u16 {
self.depth
}
}

#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
Expand Down Expand Up @@ -244,11 +240,11 @@ impl<TX: Number + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>>
visitor_queue.push_back(visitor);
}

while base_tree.depth() < base_tree.parameters().max_depth.unwrap_or(u16::MAX) {
match visitor_queue.pop_front() {
Some(node) => base_tree.split(node, mtry, &mut visitor_queue, &mut rng),
None => break,
};
let max_depth = base_tree.parameters().max_depth.unwrap_or(u16::MAX);
while let Some(node) = visitor_queue.pop_front() {
if node.level < max_depth {
base_tree.split(node, mtry, &mut visitor_queue, &mut rng);
}
}

Ok(base_tree)
Expand Down Expand Up @@ -553,6 +549,7 @@ mod tests {
use super::*;
use crate::linalg::basic::arrays::Array;
use crate::linalg::basic::matrix::DenseMatrix;
use crate::metrics::mean_absolute_error;

#[test]
fn test_fit_on_empty_data_returns_error() {
Expand Down Expand Up @@ -599,4 +596,32 @@ mod tests {
assert!(result.is_err());
assert_eq!(result.err().unwrap().error(), FailedError::ParametersError);
}

#[test]
fn full_depth() {
let x = DenseMatrix::from_2d_vec(&vec![
vec![1.0_f64],
vec![2.0],
vec![3.0],
vec![4.0],
vec![5.0],
vec![6.0],
])
.unwrap();
let y = vec![1.0f64, 2.0, 6.0, 7.0, 11., 12.];

let parameters = BaseTreeRegressorParameters {
max_depth: Some(3),
min_samples_leaf: 1,
min_samples_split: 2,
seed: None,
splitter: Splitter::Best,
};

let tree = BaseTreeRegressor::fit(&x, &y, parameters).expect("Fit should work");
let y_expected = vec![1.0, 2.0, 6.5, 6.5, 11.50, 11.50];
let y_hat = tree.predict(&x).expect("Predict should work");
assert_eq!(tree.nodes().len(), 7);
assert!(mean_absolute_error(&y_expected, &y_hat) < 1e-9);
}
}
42 changes: 36 additions & 6 deletions src/tree/decision_tree_classifier.rs
Original file line number Diff line number Diff line change
Expand Up @@ -624,11 +624,11 @@ impl<TX: Number + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY>>
visitor_queue.push_back(visitor);
}

while tree.depth() < tree.parameters().max_depth.unwrap_or(u16::MAX) {
match visitor_queue.pop_front() {
Some(node) => tree.split(node, mtry, &mut visitor_queue, &mut rng),
None => break,
};
let max_depth = tree.parameters().max_depth.unwrap_or(u16::MAX);
while let Some(node) = visitor_queue.pop_front() {
if node.level < max_depth {
tree.split(node, mtry, &mut visitor_queue, &mut rng);
}
}

Ok(tree)
Expand Down Expand Up @@ -707,7 +707,7 @@ impl<TX: Number + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY>>
return false;
}

if n <= self.parameters().min_samples_split {
if n < self.parameters().min_samples_split {
return false;
}

Expand Down Expand Up @@ -967,6 +967,7 @@ mod tests {
use super::*;
use crate::linalg::basic::arrays::Array;
use crate::linalg::basic::matrix::DenseMatrix;
use crate::metrics::accuracy;

#[test]
fn search_parameters() {
Expand Down Expand Up @@ -1053,6 +1054,35 @@ mod tests {
}
}

#[test]
fn full_depth() {
let x = DenseMatrix::from_2d_vec(&vec![
vec![1.0_f64],
vec![2.0],
vec![3.0],
vec![4.0],
vec![5.0],
vec![6.0],
])
.unwrap();
let y = vec![0, 1, 2, 2, 3, 4];

let parameters = DecisionTreeClassifierParameters {
max_depth: Some(3),
min_samples_leaf: 1,
min_samples_split: 2,
seed: None,
criterion: SplitCriterion::Gini,
};

let tree = DecisionTreeClassifier::fit(&x, &y, parameters).expect("Fit should work");
let y_hat = tree.predict(&x).expect("Predict should work");
assert_eq!(tree.nodes().len(), 7);

// Tree should have 5 out of 6 examples correct
assert!((accuracy(&y, &y_hat) - 5.0 / 6.0).abs() < 1e-9);
}

#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
Expand Down
Loading