plotters-statistical 0.1.0

Statistical chart primitives (box, violin, ROC, PR, regularization-path, residual) as native plotters series
Documentation
//! Unit tests for the `stats` layer against hand-computed reference values and
//! every edge case called out in Milestone 1.

use approx::assert_abs_diff_eq;
use plotters_statistical::stats::{
    kde_curve, precision_recall_curve, quartiles, roc::auc_trapezoid, roc_curve,
    silverman_bandwidth, StatsError,
};

// ---------------------------------------------------------------------------
// Quartiles (type-7 / linear interpolation)
// ---------------------------------------------------------------------------

#[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() {
    // Hand-computed for [1,2,3,4,5,100] under type-7:
    //   q1 = 2.25, median = 3.5, q3 = 4.75, IQR = 2.5
    //   upper fence = 4.75 + 1.5*2.5 = 8.5  -> 100 is an outlier
    //   upper whisker = 5 (largest value <= 8.5)
    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);
    // Zero IQR collapses the fences onto the value: nothing is an outlier.
    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));
}

// ---------------------------------------------------------------------------
// KDE
// ---------------------------------------------------------------------------

#[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() {
    // A symmetric sample: the density should integrate (trapezoid) to ~1 and
    // peak near the center of mass.
    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));

    // Peak index should sit near x = 0 for this symmetric data.
    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));
}

// ---------------------------------------------------------------------------
// ROC
// ---------------------------------------------------------------------------

#[test]
fn roc_matches_reference() {
    // The scikit-learn documentation example:
    //   y_true  = [0, 0, 1, 1], y_score = [0.1, 0.4, 0.35, 0.8]
    //   fpr = [0, 0, 0.5, 0.5, 1], tpr = [0, 0.5, 0.5, 1, 1], AUC = 0.75
    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() {
    // Straight line from (0,0) to (1,1) -> area 0.5.
    assert_abs_diff_eq!(auc_trapezoid(&[(0.0, 0.0), (1.0, 1.0)]), 0.5);
}

// ---------------------------------------------------------------------------
// Precision / recall
// ---------------------------------------------------------------------------

#[test]
fn pr_matches_reference() {
    // Same data; scikit-learn average_precision_score = 0.8333...
    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);
    // Recall is non-decreasing and reaches 1.0.
    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)
    );
}