use super::*;
#[test]
fn variance_threshold_drops_constant_feature() -> Result<(), FeatureSelectionError> {
let x = vec![
vec![0.0, 1.0, 2.0],
vec![0.0, 4.0, 3.0],
vec![0.0, 7.0, 10.0],
];
let sel = variance_threshold(&x, 1.0)?;
assert_eq!(
sel.mask(),
&[false, true, true],
"mask was {:?}",
sel.mask()
);
let scores = sel.scores();
assert!(
scores.first().copied().unwrap_or(f64::NAN).abs() < 1e-12,
"feature 0 variance was {:?}",
scores.first()
);
let f1 = scores.get(1).copied().unwrap_or(f64::NAN);
assert!((f1 - 6.0).abs() < 1e-12, "feature 1 variance was {f1}");
Ok(())
}
#[test]
fn anova_f_scores_match_hand_value() -> Result<(), FeatureSelectionError> {
let x = vec![
vec![1.0],
vec![2.0],
vec![3.0],
vec![5.0],
vec![6.0],
vec![7.0],
vec![10.0],
vec![11.0],
vec![12.0],
];
let labels = [0usize, 0, 0, 1, 1, 1, 2, 2, 2];
let f = anova_f_scores(&x, &labels)?;
assert_eq!(f.len(), 1, "one feature expected, got {}", f.len());
let f0 = f.first().copied().unwrap_or(f64::NAN);
assert!((f0 - 61.0).abs() < 1e-9, "F was {f0}");
Ok(())
}
#[test]
fn anova_f_select_keeps_top_k() -> Result<(), FeatureSelectionError> {
let x = vec![
vec![0.0, 1.0],
vec![0.1, 0.9],
vec![9.0, 1.0],
vec![9.1, 1.1],
];
let labels = [0usize, 0, 1, 1];
let sel = anova_f_select(&x, &labels, 1)?;
assert_eq!(sel.selected_count(), 1, "exactly one feature kept");
assert_eq!(
sel.selected_indices(),
vec![0],
"the separating feature 0 must win, got {:?}",
sel.selected_indices()
);
Ok(())
}
#[test]
fn anova_f_select_k_bounds() -> Result<(), FeatureSelectionError> {
let x = vec![
vec![0.0, 1.0],
vec![0.1, 0.9],
vec![9.0, 5.0],
vec![9.1, 5.2],
];
let labels = [0usize, 0, 1, 1];
let all = anova_f_select(&x, &labels, 9)?;
assert_eq!(all.selected_count(), 2, "k beyond width selects all");
let none = anova_f_select(&x, &labels, 0)?;
assert_eq!(none.selected_count(), 0, "k = 0 selects none");
Ok(())
}
#[test]
fn anova_f_pvalues_track_scores() -> Result<(), FeatureSelectionError> {
let x = vec![
vec![0.0, 1.0],
vec![0.1, 0.9],
vec![9.0, 1.0],
vec![9.1, 1.1],
];
let labels = [0usize, 0, 1, 1];
let f = anova_f_scores(&x, &labels)?;
let p = anova_f_pvalues(&x, &labels)?;
assert_eq!(p.len(), 2, "one p-value per feature");
for (i, &pi) in p.iter().enumerate() {
assert!((0.0..=1.0).contains(&pi), "p[{i}] out of range: {pi}");
}
let (f0, f1) = (
f.first().copied().unwrap_or(0.0),
f.get(1).copied().unwrap_or(0.0),
);
let (p0, p1) = (
p.first().copied().unwrap_or(1.0),
p.get(1).copied().unwrap_or(1.0),
);
assert!(f0 > f1, "feature 0 should have larger F: {f0} vs {f1}");
assert!(p0 < p1, "feature 0 should have smaller p: {p0} vs {p1}");
Ok(())
}
#[test]
fn empty_matrix_is_rejected() {
let x: Vec<Vec<f64>> = Vec::new();
assert_eq!(
variance_threshold(&x, 0.0),
Err(FeatureSelectionError::EmptyInput)
);
assert_eq!(
anova_f_scores(&x, &[0usize]),
Err(FeatureSelectionError::EmptyInput)
);
}
#[test]
fn no_features_is_rejected() {
let x = vec![Vec::new(), Vec::new()];
assert_eq!(
variance_threshold(&x, 0.0),
Err(FeatureSelectionError::NoFeatures)
);
}
#[test]
fn ragged_rows_are_rejected() {
let x = vec![vec![1.0, 2.0], vec![3.0]];
assert_eq!(
variance_threshold(&x, 0.0),
Err(FeatureSelectionError::RaggedRows)
);
}
#[test]
fn non_finite_is_rejected() {
let x = vec![vec![1.0, f64::NAN], vec![3.0, 4.0]];
assert_eq!(
variance_threshold(&x, 0.0),
Err(FeatureSelectionError::NonFinite)
);
}
#[test]
fn non_finite_threshold_is_rejected() {
let x = vec![vec![1.0], vec![2.0]];
assert_eq!(
variance_threshold(&x, f64::NAN),
Err(FeatureSelectionError::InvalidThreshold)
);
}
#[test]
fn label_length_mismatch_is_rejected() {
let x = vec![vec![1.0], vec![2.0], vec![3.0]];
assert_eq!(
anova_f_scores(&x, &[0usize, 1]),
Err(FeatureSelectionError::LabelLengthMismatch)
);
}
#[test]
fn too_few_classes_is_rejected() {
let x = vec![vec![1.0], vec![2.0], vec![3.0]];
assert_eq!(
anova_f_scores(&x, &[0usize, 0, 0]),
Err(FeatureSelectionError::TooFewClasses)
);
}
#[test]
fn degenerate_classes_is_rejected() {
let x = vec![vec![5.0], vec![5.0], vec![5.0], vec![5.0]];
let labels = [0usize, 0, 1, 1];
assert_eq!(
anova_f_scores(&x, &labels),
Err(FeatureSelectionError::DegenerateClasses)
);
}