use approx::assert_abs_diff_eq;
use plotters_statistical::stats::{
kde_curve, precision_recall_curve, quartiles, roc::auc_trapezoid, roc_curve,
silverman_bandwidth, StatsError,
};
#[test]
fn quartiles_simple_odd() {
let q = quartiles(&[1.0, 2.0, 3.0, 4.0, 5.0]).unwrap();
assert_abs_diff_eq!(q.q1, 2.0);
assert_abs_diff_eq!(q.median, 3.0);
assert_abs_diff_eq!(q.q3, 4.0);
assert_abs_diff_eq!(q.iqr, 2.0);
assert_abs_diff_eq!(q.lower_whisker, 1.0);
assert_abs_diff_eq!(q.upper_whisker, 5.0);
assert!(q.outliers.is_empty());
}
#[test]
fn quartiles_with_high_outlier() {
let q = quartiles(&[1.0, 2.0, 3.0, 4.0, 5.0, 100.0]).unwrap();
assert_abs_diff_eq!(q.q1, 2.25);
assert_abs_diff_eq!(q.median, 3.5);
assert_abs_diff_eq!(q.q3, 4.75);
assert_abs_diff_eq!(q.iqr, 2.5);
assert_abs_diff_eq!(q.upper_fence(), 8.5);
assert_abs_diff_eq!(q.upper_whisker, 5.0);
assert_abs_diff_eq!(q.lower_whisker, 1.0);
assert_eq!(q.outliers, vec![100.0]);
}
#[test]
fn quartiles_single_point() {
let q = quartiles(&[7.0]).unwrap();
assert_abs_diff_eq!(q.q1, 7.0);
assert_abs_diff_eq!(q.median, 7.0);
assert_abs_diff_eq!(q.q3, 7.0);
assert_abs_diff_eq!(q.iqr, 0.0);
assert!(q.outliers.is_empty());
}
#[test]
fn quartiles_all_identical_zero_iqr() {
let q = quartiles(&[3.0, 3.0, 3.0, 3.0]).unwrap();
assert_abs_diff_eq!(q.iqr, 0.0);
assert!(q.outliers.is_empty());
assert_abs_diff_eq!(q.lower_whisker, 3.0);
assert_abs_diff_eq!(q.upper_whisker, 3.0);
}
#[test]
fn quartiles_empty_errors() {
assert_eq!(quartiles(&[]), Err(StatsError::EmptyInput));
}
#[test]
fn silverman_positive_for_spread_data() {
let h = silverman_bandwidth(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
assert!(h > 0.0 && h.is_finite());
}
#[test]
fn silverman_errors_on_degenerate() {
assert_eq!(
silverman_bandwidth(&[5.0]),
Err(StatsError::InvalidBandwidth)
);
assert_eq!(
silverman_bandwidth(&[2.0, 2.0, 2.0]),
Err(StatsError::InvalidBandwidth)
);
}
#[test]
fn kde_integrates_to_one() {
let data = [-2.0, -1.0, 0.0, 1.0, 2.0];
let curve = kde_curve(&data, Some(0.5), 400, 4.0).unwrap();
let area: f64 = curve
.xs
.windows(2)
.zip(curve.density.windows(2))
.map(|(x, d)| (x[1] - x[0]) * (d[0] + d[1]) / 2.0)
.sum();
assert_abs_diff_eq!(area, 1.0, epsilon = 1e-2);
assert!(curve.density.iter().all(|&d| d >= 0.0));
let (peak_i, _) = curve
.density
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.unwrap();
assert_abs_diff_eq!(curve.xs[peak_i], 0.0, epsilon = 0.1);
}
#[test]
fn kde_rejects_bad_bandwidth() {
assert_eq!(
kde_curve(&[1.0, 2.0, 3.0], Some(0.0), 100, 3.0),
Err(StatsError::InvalidBandwidth)
);
assert_eq!(kde_curve(&[], None, 100, 3.0), Err(StatsError::EmptyInput));
}
#[test]
fn roc_matches_reference() {
let scores = [0.1, 0.4, 0.35, 0.8];
let labels = [false, false, true, true];
let roc = roc_curve(&scores, &labels).unwrap();
assert_abs_diff_eq!(roc.auc, 0.75);
let fprs: Vec<f64> = roc.points.iter().map(|p| p.fpr).collect();
let tprs: Vec<f64> = roc.points.iter().map(|p| p.tpr).collect();
assert_eq!(fprs, vec![0.0, 0.0, 0.5, 0.5, 1.0]);
assert_eq!(tprs, vec![0.0, 0.5, 0.5, 1.0, 1.0]);
}
#[test]
fn roc_perfect_separation_auc_one() {
let scores = [0.9, 0.8, 0.2, 0.1];
let labels = [true, true, false, false];
let roc = roc_curve(&scores, &labels).unwrap();
assert_abs_diff_eq!(roc.auc, 1.0);
}
#[test]
fn roc_edge_cases_error() {
assert_eq!(
roc_curve(&[0.1, 0.2], &[true]),
Err(StatsError::LengthMismatch {
scores: 2,
labels: 1
})
);
assert_eq!(roc_curve(&[], &[]), Err(StatsError::EmptyInput));
assert_eq!(
roc_curve(&[0.1, 0.2], &[true, true]),
Err(StatsError::NoNegativeLabels)
);
assert_eq!(
roc_curve(&[0.1, 0.2], &[false, false]),
Err(StatsError::NoPositiveLabels)
);
}
#[test]
fn auc_trapezoid_unit_triangle() {
assert_abs_diff_eq!(auc_trapezoid(&[(0.0, 0.0), (1.0, 1.0)]), 0.5);
}
#[test]
fn pr_matches_reference() {
let scores = [0.1, 0.4, 0.35, 0.8];
let labels = [false, false, true, true];
let pr = precision_recall_curve(&scores, &labels).unwrap();
assert_abs_diff_eq!(pr.average_precision, 5.0 / 6.0, epsilon = 1e-9);
assert_abs_diff_eq!(pr.baseline, 0.5);
assert!(pr.points.windows(2).all(|w| w[1].recall >= w[0].recall));
assert_abs_diff_eq!(pr.points.last().unwrap().recall, 1.0);
}
#[test]
fn pr_no_positives_errors() {
assert_eq!(
precision_recall_curve(&[0.1, 0.2], &[false, false]),
Err(StatsError::NoPositiveLabels)
);
}