plotters-statistical 0.2.0

Statistical chart primitives (box, violin, ROC, PR, regularization-path, residual) as native plotters series
Documentation
//! All six chart types on one multi-panel figure, rendered to `dashboard.svg`.
//!
//! Doubles as a visual regression test and the README's hero image, and is the
//! natural place to confirm the shared styling looks consistent side by side.

use plotters::coord::Shift;
use plotters::prelude::*;
use plotters_statistical::style::palette_color;
use plotters_statistical::{
    BoxPlotSeries, PrecisionRecallCurve, RegularizationPath, ResidualPlot, RocCurve,
    ViolinPlotSeries,
};

type Res = Result<(), Box<dyn std::error::Error>>;
type Panel<'a> = DrawingArea<SVGBackend<'a>, Shift>;

fn noise(i: usize) -> f64 {
    ((i as f64 * 12.9898).sin() * 43758.5453).rem_euclid(1.0)
}

fn dataset(sep: f64, balanced: bool) -> (Vec<f64>, Vec<bool>) {
    let n = 160;
    let mut scores = Vec::with_capacity(n);
    let mut labels = Vec::with_capacity(n);
    for i in 0..n {
        let pos = if balanced { i % 2 == 0 } else { i % 10 < 3 };
        scores.push(if pos { sep } else { 0.0 } + noise(i));
        labels.push(pos);
    }
    (scores, labels)
}

fn sample(center: f64, spread: f64, n: usize) -> Vec<f64> {
    (0..n)
        .map(|i| {
            let t = i as f64 / n as f64;
            center + spread * ((t * std::f64::consts::TAU * 2.0).sin() + (t - 0.5) * 2.0)
        })
        .collect()
}

fn main() -> Res {
    let root = SVGBackend::new("dashboard.svg", (1280, 720)).into_drawing_area();
    root.fill(&WHITE)?;
    let panels = root.split_evenly((2, 3));

    panel_box(&panels[0])?;
    panel_violin(&panels[1])?;
    panel_roc(&panels[2])?;
    panel_pr(&panels[3])?;
    panel_regpath(&panels[4])?;
    panel_residual(&panels[5])?;

    root.present()?;
    println!("wrote dashboard.svg");
    Ok(())
}

fn panel_box(area: &Panel) -> Res {
    let groups = vec![
        (1.0f64, vec![4.0, 5.0, 5.5, 6.0, 6.5, 7.0, 7.5, 8.0]),
        (2.0f64, vec![5.0, 6.0, 6.5, 7.0, 7.5, 8.0, 9.0, 16.0]),
    ];
    let mut chart = ChartBuilder::on(area)
        .caption("Box plot", ("sans-serif", 18))
        .margin(12)
        .set_label_area_size(LabelAreaPosition::Left, 38)
        .set_label_area_size(LabelAreaPosition::Bottom, 30)
        .build_cartesian_2d(0.5f64..2.5f64, 0f64..18f64)?;
    chart.configure_mesh().draw()?;
    chart.draw_series(BoxPlotSeries::from_samples(groups)?.width(40))?;
    Ok(())
}

fn panel_violin(area: &Panel) -> Res {
    let groups = vec![
        (1.0f64, sample(6.0, 1.5, 100)),
        (2.0f64, sample(3.5, 1.0, 100)),
    ];
    let mut chart = ChartBuilder::on(area)
        .caption("Violin plot", ("sans-serif", 18))
        .margin(12)
        .set_label_area_size(LabelAreaPosition::Left, 38)
        .set_label_area_size(LabelAreaPosition::Bottom, 30)
        .build_cartesian_2d(0.5f64..2.5f64, 0f64..12f64)?;
    chart.configure_mesh().draw()?;
    chart.draw_series(
        ViolinPlotSeries::from_samples(groups)?
            .width(70)
            .show_box(true),
    )?;
    Ok(())
}

fn panel_roc(area: &Panel) -> Res {
    let mut chart = ChartBuilder::on(area)
        .caption("ROC", ("sans-serif", 18))
        .margin(12)
        .set_label_area_size(LabelAreaPosition::Left, 38)
        .set_label_area_size(LabelAreaPosition::Bottom, 30)
        .build_cartesian_2d(0f64..1f64, 0f64..1f64)?;
    chart.configure_mesh().draw()?;
    for (i, sep) in [0.9, 0.35].into_iter().enumerate() {
        let (s, l) = dataset(sep, true);
        let mut c = RocCurve::from_scores(&s, &l)?.color(palette_color(i));
        if i == 0 {
            c = c.with_baseline().shade_area(true);
        }
        chart.draw_series(std::iter::once(c))?;
    }
    Ok(())
}

fn panel_pr(area: &Panel) -> Res {
    let mut chart = ChartBuilder::on(area)
        .caption("Precision-recall", ("sans-serif", 18))
        .margin(12)
        .set_label_area_size(LabelAreaPosition::Left, 38)
        .set_label_area_size(LabelAreaPosition::Bottom, 30)
        .build_cartesian_2d(0f64..1f64, 0f64..1.05f64)?;
    chart.configure_mesh().draw()?;
    for (i, sep) in [0.9, 0.35].into_iter().enumerate() {
        let (s, l) = dataset(sep, false);
        let mut c = PrecisionRecallCurve::from_scores(&s, &l)?.color(palette_color(i));
        if i == 0 {
            c = c.with_baseline();
        }
        chart.draw_series(std::iter::once(c))?;
    }
    Ok(())
}

fn panel_regpath(area: &Panel) -> Res {
    let n = 25;
    let strengths: Vec<f64> = (0..n)
        .map(|i| 10f64.powf(-3.0 + 4.0 * i as f64 / (n as f64 - 1.0)))
        .collect();
    let bases = [3.0f64, -2.0, 1.5, 2.5];
    let thr = [3.0f64, 0.3, 1.0, 0.05];
    let coefs: Vec<Vec<f64>> = strengths
        .iter()
        .map(|&s| {
            bases
                .iter()
                .zip(thr.iter())
                .map(|(&b, &t)| b * (1.0 - s / t).max(0.0))
                .collect()
        })
        .collect();
    let mut chart = ChartBuilder::on(area)
        .caption("Regularization path", ("sans-serif", 18))
        .margin(12)
        .set_label_area_size(LabelAreaPosition::Left, 38)
        .set_label_area_size(LabelAreaPosition::Bottom, 30)
        .build_cartesian_2d((1e-3f64..1e1f64).log_scale(), -2.5f64..3.5f64)?;
    chart.configure_mesh().draw()?;
    chart.draw_series(RegularizationPath::new(&strengths, &coefs)?)?;
    Ok(())
}

fn panel_residual(area: &Panel) -> Res {
    let n = 160;
    let mut fitted = Vec::with_capacity(n);
    let mut residuals = Vec::with_capacity(n);
    for i in 0..n {
        let x = 10.0 * i as f64 / (n as f64 - 1.0);
        let r = 0.45 * (x - 5.0) + (noise(i) * 2.0 - 1.0) * (0.6 + 0.3 * x);
        fitted.push(x);
        residuals.push(r);
    }
    let mut chart = ChartBuilder::on(area)
        .caption("Residual plot", ("sans-serif", 18))
        .margin(12)
        .set_label_area_size(LabelAreaPosition::Left, 38)
        .set_label_area_size(LabelAreaPosition::Bottom, 30)
        .build_cartesian_2d(0f64..10f64, -8f64..8f64)?;
    chart.configure_mesh().draw()?;
    chart.draw_series(std::iter::once(
        ResidualPlot::from_residuals(&fitted, &residuals)?.trend(true),
    ))?;
    Ok(())
}