use super::*;
use crate::primitives::Matrix;
#[test]
fn falsify_dt_001_predictions_in_label_range() {
let x = Matrix::from_vec(
6,
2,
vec![0.0, 0.0, 1.0, 0.0, 2.0, 0.0, 0.0, 1.0, 1.0, 1.0, 2.0, 1.0],
)
.expect("valid matrix");
let y = vec![0_usize, 0, 1, 1, 2, 2];
let mut dt = DecisionTreeClassifier::new();
dt.fit(&x, &y).expect("fit succeeds");
let preds = dt.predict(&x);
for (i, &p) in preds.iter().enumerate() {
assert!(
p <= 2,
"FALSIFIED DT-001: prediction[{i}] = {p}, not in [0, 2]"
);
}
}
#[test]
fn falsify_dt_002_deterministic() {
let x =
Matrix::from_vec(4, 2, vec![0.0, 0.0, 1.0, 1.0, 2.0, 2.0, 3.0, 3.0]).expect("valid matrix");
let y = vec![0_usize, 0, 1, 1];
let mut dt = DecisionTreeClassifier::new();
dt.fit(&x, &y).expect("fit");
let p1 = dt.predict(&x);
let p2 = dt.predict(&x);
assert_eq!(p1, p2, "FALSIFIED DT-002: predictions differ on same input");
}
#[test]
fn falsify_dt_003_perfect_separable() {
let x = Matrix::from_vec(4, 1, vec![0.0, 1.0, 10.0, 11.0]).expect("valid matrix");
let y = vec![0_usize, 0, 1, 1];
let mut dt = DecisionTreeClassifier::new();
dt.fit(&x, &y).expect("fit");
let preds = dt.predict(&x);
assert_eq!(
preds, y,
"FALSIFIED DT-003: tree cannot perfectly fit separable data"
);
}
#[test]
fn falsify_dt_004_prediction_count() {
let x_train =
Matrix::from_vec(4, 2, vec![0.0, 0.0, 1.0, 1.0, 2.0, 2.0, 3.0, 3.0]).expect("valid");
let y_train = vec![0_usize, 0, 1, 1];
let mut dt = DecisionTreeClassifier::new();
dt.fit(&x_train, &y_train).expect("fit");
let x_test = Matrix::from_vec(3, 2, vec![0.5, 0.5, 1.5, 1.5, 2.5, 2.5]).expect("valid");
let preds = dt.predict(&x_test);
assert_eq!(
preds.len(),
3,
"FALSIFIED DT-004: {} predictions for 3 inputs",
preds.len()
);
}
#[test]
fn falsify_dt_008_split_search_matches_rescan_bit_for_bit() {
use crate::tree::helpers::{find_best_split_for_feature, find_best_split_for_feature_rescan};
fn check(x: &[f32], y: &[usize], case: &str) {
let fast = find_best_split_for_feature(x, y);
let oracle = find_best_split_for_feature_rescan(x, y);
match (fast, oracle) {
(None, None) => {}
(Some((tf, gf)), Some((to, go))) => assert_eq!(
(tf.to_bits(), gf.to_bits()),
(to.to_bits(), go.to_bits()),
"FALSIFIED DT-008 [{case}]: single pass chose ({tf}, {gf}), rescan chose ({to}, {go})"
),
(f, o) => panic!("FALSIFIED DT-008 [{case}]: single pass {f:?} vs rescan {o:?}"),
}
}
check(&[1.0], &[0], "single sample");
check(&[1.0, 1.0, 1.0], &[0, 1, 0], "all values identical");
check(&[0.0, 1.0, 2.0, 3.0], &[1, 1, 1, 1], "pure labels");
check(
&[0.0, 1.0, 10.0, 11.0],
&[0, 0, 1, 1],
"clean two-way split",
);
check(
&[1.0, 1.0 + 1e-12, 2.0, 2.0],
&[0, 1, 1, 0],
"sub-tolerance neighbours",
);
check(
&[-3.5, -3.5, 0.0, 7.25, 7.25],
&[2, 2, 0, 5, 5],
"sparse class labels",
);
check(
&[0.0, 1e-10, 2e-10, 3e-10],
&[0, 1, 0, 1],
"sample on the threshold, alternating",
);
check(
&[0.0, 1e-10, 2e-10, 3e-10],
&[0, 0, 1, 1],
"sample on the threshold, separable",
);
check(
&[0.0, 1e-10, 1e-10, 2e-10],
&[1, 0, 0, 1],
"threshold sample duplicated",
);
check(
&[-2e-10, -1e-10, 0.0, 1e-10],
&[0, 0, 1, 1],
"sample on the threshold, negative",
);
let mut state: u64 = 0x9E37_79B9_7F4A_7C15;
let mut next = move || {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(state >> 33) as usize
};
for case in 0..500 {
let n = 2 + next() % 64;
let n_classes = 2 + next() % 4;
let spread = 1 + next() % 10;
let quantum = if case % 2 == 0 { 0.25 } else { 1e-10 };
let x: Vec<f32> = (0..n).map(|_| (next() % spread) as f32 * quantum).collect();
let y: Vec<usize> = (0..n).map(|_| next() % n_classes).collect();
check(
&x,
&y,
&format!("random case {case} (n={n}, k={n_classes}, q={quantum})"),
);
}
let n = 400;
let x: Vec<f32> = (0..n).map(|_| next() as f32 / usize::MAX as f32).collect();
let y: Vec<usize> = (0..n).map(|i| usize::from(x[i] > 0.4)).collect();
check(&x, &y, "continuous column");
}
#[test]
fn falsify_dt_009_root_split_search_is_subquadratic() {
use std::time::{Duration, Instant};
const N: usize = 200_000;
const BUDGET: Duration = Duration::from_secs(60);
let mut state: u64 = 0x2545_F491_4F6C_DD1D;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
(state >> 11) as f32 / (1u64 << 53) as f32
};
let values: Vec<f32> = (0..N).map(|_| next()).collect();
let y: Vec<usize> = values.iter().map(|&v| usize::from(v > 0.6)).collect();
let x = Matrix::from_vec(N, 1, values).expect("valid matrix");
let mut dt = DecisionTreeClassifier::new().with_max_depth(1);
let started = Instant::now();
dt.fit(&x, &y).expect("fit");
let elapsed = started.elapsed();
assert!(
elapsed < BUDGET,
"FALSIFIED DT-009: one root split over {N} rows took {elapsed:?} (budget {BUDGET:?}) \
— split search is quadratic in row count again"
);
}