use ferrolearn_core::{Fit, Predict};
use ferrolearn_tree::{DecisionTreeClassifier, DecisionTreeRegressor, Node};
use ndarray::{Array1, Array2, array};
fn nan() -> f64 {
f64::NAN
}
fn root_split(nodes: &[Node<f64>]) -> (usize, f64) {
match nodes[0] {
Node::Split {
feature, threshold, ..
} => (feature, threshold),
Node::Leaf { .. } => panic!("root is a leaf, expected a split"),
}
}
#[test]
fn clf_max_depth1_missing_go_right_oracle() {
let x = Array2::from_shape_vec((6, 1), vec![1.0, 2.0, nan(), 8.0, 9.0, nan()]).unwrap();
let y = array![0, 0, 0, 1, 1, 1];
let fitted = DecisionTreeClassifier::<f64>::new()
.with_max_depth(Some(1))
.fit(&x, &y)
.expect("fit must accept NaN (sklearn force_all_finite=False)");
assert_eq!(fitted.nodes().len(), 3);
let (feat, thr) = root_split(fitted.nodes());
assert_eq!(feat, 0);
assert!((thr - 5.0).abs() < 1e-12, "threshold {thr} != 5.0");
let mut q = Array2::zeros((1, 1));
q[[0, 0]] = nan();
assert_eq!(fitted.predict(&q).unwrap()[0], 1);
assert_eq!(fitted.predict(&x).unwrap(), array![0, 0, 1, 1, 1, 1]);
}
#[test]
fn clf_max_depth1_missing_go_left_oracle() {
let x =
Array2::from_shape_vec((8, 1), vec![1.0, 2.0, 3.0, 8.0, 9.0, 10.0, nan(), nan()]).unwrap();
let y = array![0, 0, 0, 1, 1, 1, 0, 0];
let fitted = DecisionTreeClassifier::<f64>::new()
.with_max_depth(Some(1))
.fit(&x, &y)
.expect("fit must accept NaN");
assert_eq!(fitted.nodes().len(), 3);
let (feat, thr) = root_split(fitted.nodes());
assert_eq!(feat, 0);
assert!((thr - 5.5).abs() < 1e-12, "threshold {thr} != 5.5");
let mut q = Array2::zeros((1, 1));
q[[0, 0]] = nan();
assert_eq!(fitted.predict(&q).unwrap()[0], 0);
assert_eq!(fitted.predict(&x).unwrap(), array![0, 0, 0, 1, 1, 1, 0, 0]);
}
#[test]
fn clf_deep_tree_threshold_inf_candidate_oracle() {
let x = Array2::from_shape_vec((6, 1), vec![1.0, 2.0, nan(), 8.0, 9.0, nan()]).unwrap();
let y = array![0, 0, 0, 1, 1, 1];
let fitted = DecisionTreeClassifier::<f64>::new()
.fit(&x, &y)
.expect("fit must accept NaN");
assert_eq!(fitted.nodes().len(), 5);
assert_eq!(fitted.predict(&x).unwrap(), array![0, 0, 0, 1, 1, 0]);
let mut q = Array2::zeros((1, 1));
q[[0, 0]] = nan();
let proba = fitted.predict_proba(&q).unwrap();
assert!(
(proba[[0, 0]] - 0.5).abs() < 1e-12,
"proba0 {}",
proba[[0, 0]]
);
assert!(
(proba[[0, 1]] - 0.5).abs() < 1e-12,
"proba1 {}",
proba[[0, 1]]
);
}
#[test]
fn reg_max_depth1_missing_go_right_oracle() {
let x = Array2::from_shape_vec((6, 1), vec![1.0, 2.0, nan(), 8.0, 9.0, nan()]).unwrap();
let y: Array1<f64> = array![1.0, 1.0, 1.0, 5.0, 5.0, 5.0];
let fitted = DecisionTreeRegressor::<f64>::new()
.with_max_depth(Some(1))
.fit(&x, &y)
.expect("fit must accept NaN");
assert_eq!(fitted.nodes().len(), 3);
let (feat, thr) = root_split(fitted.nodes());
assert_eq!(feat, 0);
assert!((thr - 5.0).abs() < 1e-12, "threshold {thr} != 5.0");
let mut q = Array2::zeros((1, 1));
q[[0, 0]] = nan();
assert!((fitted.predict(&q).unwrap()[0] - 4.0).abs() < 1e-12);
let preds = fitted.predict(&x).unwrap();
for (p, e) in preds.iter().zip([1.0, 1.0, 4.0, 4.0, 4.0, 4.0]) {
assert!((p - e).abs() < 1e-12, "{p} != {e}");
}
}
#[test]
fn reg_max_depth1_missing_go_left_oracle() {
let x =
Array2::from_shape_vec((8, 1), vec![1.0, 2.0, 3.0, 8.0, 9.0, 10.0, nan(), nan()]).unwrap();
let y: Array1<f64> = array![1.0, 1.0, 1.0, 5.0, 5.0, 5.0, 1.0, 1.0];
let fitted = DecisionTreeRegressor::<f64>::new()
.with_max_depth(Some(1))
.fit(&x, &y)
.expect("fit must accept NaN");
assert_eq!(fitted.nodes().len(), 3);
let (feat, thr) = root_split(fitted.nodes());
assert_eq!(feat, 0);
assert!((thr - 5.5).abs() < 1e-12, "threshold {thr} != 5.5");
let mut q = Array2::zeros((1, 1));
q[[0, 0]] = nan();
assert!((fitted.predict(&q).unwrap()[0] - 1.0).abs() < 1e-12);
}
#[test]
fn clf_multi_feature_missing_in_one_feature_oracle() {
let x = Array2::from_shape_vec(
(6, 2),
vec![
nan(),
1.0,
2.0,
1.0,
nan(),
1.0,
8.0,
9.0,
9.0,
9.0,
10.0,
9.0,
],
)
.unwrap();
let y = array![0, 0, 0, 1, 1, 1];
let fitted = DecisionTreeClassifier::<f64>::new()
.with_max_depth(Some(1))
.fit(&x, &y)
.expect("fit must accept NaN");
assert_eq!(fitted.nodes().len(), 3);
let (feat, thr) = root_split(fitted.nodes());
assert_eq!(feat, 0);
assert!((thr - 5.0).abs() < 1e-12, "threshold {thr} != 5.0");
assert_eq!(fitted.predict(&x).unwrap(), array![0, 0, 0, 1, 1, 1]);
let q = Array2::from_shape_vec((3, 2), vec![nan(), 1.0, nan(), 9.0, 5.0, 5.0]).unwrap();
assert_eq!(fitted.predict(&q).unwrap(), array![0, 0, 0]);
}
#[test]
fn clf_2277_fixture_fit_succeeds_and_predicts_oracle() {
let mut x = Array2::from_shape_vec(
(6, 2),
vec![1.0, 2.0, 2.0, 3.0, 3.0, 3.0, 5.0, 6.0, 6.0, 7.0, 7.0, 8.0],
)
.unwrap();
x[[0, 0]] = nan();
let y = array![0, 0, 0, 1, 1, 1];
let fitted = DecisionTreeClassifier::<f64>::new()
.fit(&x, &y)
.expect("fit must accept NaN (#2277), not abort");
assert_eq!(fitted.nodes().len(), 3);
assert!(matches!(fitted.nodes()[0], Node::Split { .. }));
assert_eq!(fitted.predict(&x).unwrap(), array![0, 0, 0, 1, 1, 1]);
}
#[test]
fn reg_2277_fixture_fit_succeeds_and_predicts_oracle() {
let mut x = Array2::from_shape_vec(
(6, 2),
vec![1.0, 2.0, 2.0, 3.0, 3.0, 3.0, 5.0, 6.0, 6.0, 7.0, 7.0, 8.0],
)
.unwrap();
x[[0, 0]] = nan();
let y: Array1<f64> = array![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let fitted = DecisionTreeRegressor::<f64>::new()
.fit(&x, &y)
.expect("fit must accept NaN (#2277), not abort");
let preds = fitted.predict(&x).unwrap();
for (p, e) in preds.iter().zip([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]) {
assert!((p - e).abs() < 1e-12, "{p} != {e}");
}
}
#[test]
fn clf_all_finite_byte_identical_oracle() {
let x = Array2::from_shape_vec(
(9, 2),
vec![
1.0, 2.0, 2.0, 3.0, 3.0, 3.0, 5.0, 6.0, 6.0, 7.0, 7.0, 8.0, 1.5, 5.0, 6.5, 2.0, 3.0,
1.0,
],
)
.unwrap();
let y = array![0, 0, 0, 1, 1, 1, 2, 2, 0];
let fitted = DecisionTreeClassifier::<f64>::new().fit(&x, &y).unwrap();
assert_eq!(fitted.nodes().len(), 7);
let (feat, thr) = root_split(fitted.nodes());
assert_eq!(feat, 1);
assert!((thr - 5.5).abs() < 1e-12, "threshold {thr} != 5.5");
assert_eq!(
fitted.predict(&x).unwrap(),
array![0, 0, 0, 1, 1, 1, 2, 2, 0]
);
}