use ferrolearn_core::traits::{Fit, Predict};
use ferrolearn_tree::Node;
use ferrolearn_tree::decision_tree::RegressionCriterion;
use ferrolearn_tree::random_forest::MaxFeatures;
use ferrolearn_tree::{ExtraTreeClassifier, ExtraTreeRegressor};
use ndarray::{Array2, array};
#[test]
fn pin1_clf_defaults_match_sklearn() {
let m = ExtraTreeClassifier::<f64>::new();
assert_eq!(
m.max_features,
MaxFeatures::Sqrt,
"sklearn ExtraTreeClassifier max_features default is 'sqrt'"
);
use ferrolearn_tree::decision_tree::ClassificationCriterion;
assert_eq!(
m.criterion,
ClassificationCriterion::Gini,
"sklearn ExtraTreeClassifier criterion default is 'gini'"
);
assert_eq!(
m.random_state, None,
"sklearn ExtraTreeClassifier random_state default is None"
);
}
#[test]
fn pin1_reg_defaults_match_sklearn() {
let m = ExtraTreeRegressor::<f64>::new();
assert_eq!(
m.max_features,
MaxFeatures::All,
"sklearn ExtraTreeRegressor max_features default is 1.0 (== all features)"
);
assert_eq!(
m.criterion,
RegressionCriterion::Mse,
"sklearn ExtraTreeRegressor criterion default is 'squared_error'"
);
assert_eq!(
m.random_state, None,
"sklearn ExtraTreeRegressor random_state default is None"
);
}
#[test]
fn pin2_clf_random_state_reproducible() {
let x = Array2::from_shape_vec(
(8, 2),
vec![
1.0, 2.0, 2.0, 3.0, 3.0, 3.0, 4.0, 4.0, 5.0, 6.0, 6.0, 7.0, 7.0, 8.0, 8.0, 9.0,
],
)
.unwrap();
let y = array![0usize, 0, 0, 0, 1, 1, 1, 1];
let f1 = ExtraTreeClassifier::<f64>::new()
.with_random_state(42)
.fit(&x, &y)
.unwrap();
let f2 = ExtraTreeClassifier::<f64>::new()
.with_random_state(42)
.fit(&x, &y)
.unwrap();
let p1 = f1.predict(&x).unwrap();
let p2 = f2.predict(&x).unwrap();
assert_eq!(p1, p2, "same random_state must give identical predictions");
assert_eq!(
f1.nodes().len(),
f2.nodes().len(),
"same random_state must give identical tree structure"
);
for (a, b) in f1.nodes().iter().zip(f2.nodes().iter()) {
match (a, b) {
(
Node::Split {
feature: fa,
threshold: ta,
..
},
Node::Split {
feature: fb,
threshold: tb,
..
},
) => {
assert_eq!(fa, fb, "split feature must match across identical seeds");
assert_eq!(ta, tb, "split threshold must match across identical seeds");
}
(Node::Leaf { .. }, Node::Leaf { .. }) => {}
_ => panic!("node kind diverged across identical seeds"),
}
}
}
fn reg_leaves(f: &ferrolearn_tree::FittedExtraTreeRegressor<f64>) -> Vec<f64> {
let mut v: Vec<f64> = f
.nodes()
.iter()
.filter_map(|n| match n {
Node::Leaf { value, .. } => Some(*value),
Node::Split { .. } => None,
})
.collect();
v.sort_by(|a, b| a.partial_cmp(b).unwrap());
v
}
#[test]
fn pin3_reg_absolute_error_differs_from_mse_same_seed() {
let x = Array2::from_shape_vec((6, 1), vec![0.0, 1.0, 2.0, 100.0, 101.0, 102.0]).unwrap();
let y = array![1.0f64, 1.0, 5.0, 20.0, 20.0, 20.0];
let f_mse = ExtraTreeRegressor::<f64>::new()
.with_max_features(MaxFeatures::All)
.with_criterion(RegressionCriterion::Mse)
.with_max_depth(Some(1))
.with_random_state(7)
.fit(&x, &y)
.unwrap();
let f_mae = ExtraTreeRegressor::<f64>::new()
.with_max_features(MaxFeatures::All)
.with_criterion(RegressionCriterion::AbsoluteError)
.with_max_depth(Some(1))
.with_random_state(7)
.fit(&x, &y)
.unwrap();
let mse_leaves = reg_leaves(&f_mse);
let mae_leaves = reg_leaves(&f_mae);
assert_ne!(
mse_leaves, mae_leaves,
"ExtraTreeRegressor(AbsoluteError) must give median leaves distinct from \
Mse's mean leaves (sklearn _criterion.pyx MAE.node_value); ferrolearn \
hard-wires MSE means and ignores `criterion` (#681). \
Mse leaves={mse_leaves:?}, AbsoluteError leaves={mae_leaves:?}"
);
}
#[test]
fn pin3_reg_absolute_error_yields_median_leaf() {
let x = Array2::from_shape_vec((6, 1), vec![0.0, 1.0, 2.0, 100.0, 101.0, 102.0]).unwrap();
let y = array![1.0f64, 1.0, 5.0, 20.0, 20.0, 20.0];
let f_mae = ExtraTreeRegressor::<f64>::new()
.with_max_features(MaxFeatures::All)
.with_criterion(RegressionCriterion::AbsoluteError)
.with_max_depth(Some(1))
.with_random_state(5)
.fit(&x, &y)
.unwrap();
let leaf_vals = reg_leaves(&f_mae);
let median_of_small = 1.0_f64; let has_median_leaf = leaf_vals
.iter()
.any(|&v| (v - median_of_small).abs() < 1e-9);
assert!(
has_median_leaf,
"AbsoluteError must yield the median leaf value 1.0 for the y={{1,1,5}} group \
(sklearn _criterion.pyx MAE.node_value); ferrolearn returns the mean 2.333 \
because the regression builder ignores `criterion` (#681). leaves={leaf_vals:?}"
);
}