Skip to main content

ggplot_rs/stat/
summary.rs

1use crate::aes::Aesthetic;
2use crate::data::{DataFrame, Value};
3use crate::scale::ScaleSet;
4
5use super::Stat;
6
7/// Aggregation function type for StatSummary (a single scalar per group).
8#[derive(Clone)]
9pub enum SummaryFun {
10    Mean,
11    Median,
12    Min,
13    Max,
14    Sum,
15}
16
17impl SummaryFun {
18    pub fn apply(&self, values: &[f64]) -> f64 {
19        if values.is_empty() {
20            return 0.0;
21        }
22        match self {
23            SummaryFun::Mean => mean(values),
24            SummaryFun::Median => {
25                let mut sorted = values.to_vec();
26                sorted.sort_by(|a, b| a.total_cmp(b));
27                quantile_type7(&sorted, 0.5)
28            }
29            SummaryFun::Min => values.iter().cloned().fold(f64::INFINITY, f64::min),
30            SummaryFun::Max => values.iter().cloned().fold(f64::NEG_INFINITY, f64::max),
31            SummaryFun::Sum => values.iter().sum(),
32        }
33    }
34}
35
36/// A `fun.data`-style summary: from a group's values it returns `(y, ymin, ymax)`
37/// together — needed for measures where the interval depends on the centre (mean
38/// ± CI), which the scalar [`SummaryFun`]s cannot express. These mirror R's
39/// Hmisc helpers used by `ggplot2::stat_summary`.
40#[derive(Clone)]
41pub enum SummaryData {
42    /// Mean ± standard error of the mean (`mean_se`).
43    MeanSe,
44    /// Mean with a normal-theory *t* confidence interval (`mean_cl_normal`).
45    MeanClNormal { level: f64 },
46    /// Mean with a bootstrap percentile CI (`mean_cl_boot`), `b` resamples.
47    MeanClBoot { level: f64, b: usize },
48    /// Mean ± `mult` × sd (`mean_sdl`; ggplot2 default `mult = 2`).
49    MeanSdl { mult: f64 },
50    /// Median with outer sample quantiles (`median_hilow`; `level = 0.95` →
51    /// median with the 2.5% / 97.5% quantiles).
52    MedianHilow { level: f64 },
53}
54
55impl SummaryData {
56    /// `(y, ymin, ymax)` for a group's `values`.
57    pub fn apply3(&self, values: &[f64]) -> (f64, f64, f64) {
58        let n = values.len();
59        if n == 0 {
60            return (0.0, 0.0, 0.0);
61        }
62        let m = mean(values);
63        match self {
64            SummaryData::MeanSe => {
65                let se = sd(values) / (n as f64).sqrt();
66                (m, m - se, m + se)
67            }
68            SummaryData::MeanClNormal { level } => {
69                if n < 2 {
70                    return (m, m, m);
71                }
72                let se = sd(values) / (n as f64).sqrt();
73                let t = crate::stat::dist::qt(0.5 + level / 2.0, n as f64 - 1.0);
74                (m, m - t * se, m + t * se)
75            }
76            SummaryData::MeanSdl { mult } => {
77                let s = sd(values);
78                (m, m - mult * s, m + mult * s)
79            }
80            SummaryData::MedianHilow { level } => {
81                let mut s = values.to_vec();
82                s.sort_by(|a, b| a.total_cmp(b));
83                (
84                    quantile_type7(&s, 0.5),
85                    quantile_type7(&s, (1.0 - level) / 2.0),
86                    quantile_type7(&s, (1.0 + level) / 2.0),
87                )
88            }
89            SummaryData::MeanClBoot { level, b } => {
90                if n < 2 {
91                    return (m, m, m);
92                }
93                // Percentile bootstrap of the mean, seeded for reproducibility;
94                // centre stays the observed mean (as in Hmisc::smean.cl.boot).
95                let mut rng = crate::rng::SplitMix64::new(0x5EED_B007);
96                let mut means = Vec::with_capacity(*b);
97                for _ in 0..*b {
98                    let mut acc = 0.0;
99                    for _ in 0..n {
100                        acc += values[rng.below(n)];
101                    }
102                    means.push(acc / n as f64);
103                }
104                means.sort_by(|a, b| a.total_cmp(b));
105                (
106                    m,
107                    quantile_type7(&means, (1.0 - level) / 2.0),
108                    quantile_type7(&means, (1.0 + level) / 2.0),
109                )
110            }
111        }
112    }
113}
114
115fn mean(v: &[f64]) -> f64 {
116    v.iter().sum::<f64>() / v.len() as f64
117}
118
119/// Sample standard deviation (denominator n − 1), matching R's `sd()`.
120fn sd(v: &[f64]) -> f64 {
121    let n = v.len();
122    if n < 2 {
123        return 0.0;
124    }
125    let m = mean(v);
126    (v.iter().map(|x| (x - m).powi(2)).sum::<f64>() / (n as f64 - 1.0)).sqrt()
127}
128
129/// Type-7 (R default) quantile of an ascending-sorted slice.
130fn quantile_type7(sorted: &[f64], p: f64) -> f64 {
131    let n = sorted.len();
132    if n == 0 {
133        return f64::NAN;
134    }
135    if n == 1 {
136        return sorted[0];
137    }
138    let h = (n as f64 - 1.0) * p;
139    let lo = h.floor() as usize;
140    let hi = (lo + 1).min(n - 1);
141    sorted[lo] + (h - lo as f64) * (sorted[hi] - sorted[lo])
142}
143
144/// Summarize y values for each unique x. Either a `fun.data` measure (mean ± CI,
145/// median hilow, …) or three independent scalar [`SummaryFun`]s (y / ymin / ymax).
146pub struct StatSummary {
147    pub fun_y: SummaryFun,
148    pub fun_ymin: SummaryFun,
149    pub fun_ymax: SummaryFun,
150    /// When set, overrides the three scalar functions and computes y/ymin/ymax
151    /// together (needed for centre-dependent intervals like mean ± CI).
152    pub fun_data: Option<SummaryData>,
153}
154
155impl Default for StatSummary {
156    fn default() -> Self {
157        StatSummary {
158            fun_y: SummaryFun::Mean,
159            fun_ymin: SummaryFun::Min,
160            fun_ymax: SummaryFun::Max,
161            fun_data: None,
162        }
163    }
164}
165
166impl StatSummary {
167    fn with_data(d: SummaryData) -> Self {
168        StatSummary {
169            fun_data: Some(d),
170            ..Default::default()
171        }
172    }
173    /// Mean ± standard error.
174    pub fn mean_se() -> Self {
175        Self::with_data(SummaryData::MeanSe)
176    }
177    /// Mean with a 95% normal-theory (t) confidence interval.
178    pub fn mean_cl_normal() -> Self {
179        Self::with_data(SummaryData::MeanClNormal { level: 0.95 })
180    }
181    /// Mean with a 95% bootstrap-percentile confidence interval.
182    pub fn mean_cl_boot() -> Self {
183        Self::with_data(SummaryData::MeanClBoot {
184            level: 0.95,
185            b: 1000,
186        })
187    }
188    /// Mean ± 2 sd.
189    pub fn mean_sdl() -> Self {
190        Self::with_data(SummaryData::MeanSdl { mult: 2.0 })
191    }
192    /// Median with the 2.5% / 97.5% sample quantiles.
193    pub fn median_hilow() -> Self {
194        Self::with_data(SummaryData::MedianHilow { level: 0.95 })
195    }
196}
197
198impl Stat for StatSummary {
199    fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
200        let x_col = match data.column("x") {
201            Some(c) => c,
202            None => return DataFrame::new(),
203        };
204        let y_col = match data.column("y") {
205            Some(c) => c,
206            None => return DataFrame::new(),
207        };
208
209        // Group y values by x
210        let mut groups: Vec<(String, Value, Vec<f64>)> = Vec::new();
211        for (x, y) in x_col.iter().zip(y_col.iter()) {
212            let key = x.to_group_key();
213            let y_val = y.as_f64().unwrap_or(0.0);
214            if let Some(entry) = groups.iter_mut().find(|(k, _, _)| k == &key) {
215                entry.2.push(y_val);
216            } else {
217                groups.push((key, x.clone(), vec![y_val]));
218            }
219        }
220
221        let n = groups.len();
222        let mut x_vals = Vec::with_capacity(n);
223        let mut y_vals = Vec::with_capacity(n);
224        let mut ymin_vals = Vec::with_capacity(n);
225        let mut ymax_vals = Vec::with_capacity(n);
226
227        for (_, x_val, ys) in &groups {
228            x_vals.push(x_val.clone());
229            let (y, ymin, ymax) = match &self.fun_data {
230                Some(fd) => fd.apply3(ys),
231                None => (
232                    self.fun_y.apply(ys),
233                    self.fun_ymin.apply(ys),
234                    self.fun_ymax.apply(ys),
235                ),
236            };
237            y_vals.push(Value::Float(y));
238            ymin_vals.push(Value::Float(ymin));
239            ymax_vals.push(Value::Float(ymax));
240        }
241
242        let mut result = DataFrame::new();
243        result.add_column("x".to_string(), x_vals);
244        result.add_column("y".to_string(), y_vals);
245        result.add_column("ymin".to_string(), ymin_vals);
246        result.add_column("ymax".to_string(), ymax_vals);
247
248        // Carry over grouping columns
249        for col_name in &["color", "fill", "group"] {
250            if let Some(col) = data.column(col_name) {
251                if let Some(first) = col.first() {
252                    result.add_column(col_name.to_string(), vec![first.clone(); n]);
253                }
254            }
255        }
256
257        result
258    }
259
260    fn required_aes(&self) -> Vec<Aesthetic> {
261        vec![Aesthetic::X, Aesthetic::Y]
262    }
263
264    fn name(&self) -> &str {
265        "summary"
266    }
267}
268
269#[cfg(test)]
270mod tests {
271    use super::SummaryData;
272
273    // Reference values from R (mean_se / mean_cl_normal / mean_sdl / median_hilow)
274    // on v = c(2,4,4,4,5,5,7,9).
275    #[test]
276    fn summary_data_matches_r() {
277        let v = [2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0];
278        let close = |a: f64, b: f64| (a - b).abs() < 1e-5;
279
280        let (y, lo, hi) = SummaryData::MeanSe.apply3(&v);
281        assert!(close(y, 5.0) && close(lo, 4.244071) && close(hi, 5.755929));
282
283        // MeanClNormal uses the exact Student-t quantile, which is only
284        // available with the `regression` feature (otherwise a normal
285        // approximation is used); check the R-fidelity value only there.
286        #[cfg(feature = "regression")]
287        {
288            let (y, lo, hi) = SummaryData::MeanClNormal { level: 0.95 }.apply3(&v);
289            assert!(close(y, 5.0) && close(lo, 3.212512) && close(hi, 6.787488));
290        }
291
292        let (y, lo, hi) = SummaryData::MeanSdl { mult: 2.0 }.apply3(&v);
293        assert!(close(y, 5.0) && close(lo, 0.723820) && close(hi, 9.276180));
294
295        let (y, lo, hi) = SummaryData::MedianHilow { level: 0.95 }.apply3(&v);
296        assert!(close(y, 4.5) && close(lo, 2.35) && close(hi, 8.65));
297
298        // Bootstrap CI is seeded → deterministic; it should centre on the mean
299        // and bracket it within a sane range.
300        let (y, lo, hi) = SummaryData::MeanClBoot {
301            level: 0.95,
302            b: 1000,
303        }
304        .apply3(&v);
305        assert!(close(y, 5.0) && lo < 5.0 && hi > 5.0 && lo > 3.0 && hi < 7.0);
306    }
307}