1use crate::aes::Aesthetic;
2use crate::data::{DataFrame, Value};
3use crate::scale::ScaleSet;
4
5use super::Stat;
6
7#[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.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
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#[derive(Clone)]
41pub enum SummaryData {
42 MeanSe,
44 MeanClNormal { level: f64 },
46 MeanClBoot { level: f64, b: usize },
48 MeanSdl { mult: f64 },
50 MedianHilow { level: f64 },
53}
54
55impl SummaryData {
56 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.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
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 use rand::{Rng, SeedableRng};
96 let mut rng = rand::rngs::StdRng::seed_from_u64(0x5EED_B007);
97 let mut means = Vec::with_capacity(*b);
98 for _ in 0..*b {
99 let mut acc = 0.0;
100 for _ in 0..n {
101 acc += values[rng.gen_range(0..n)];
102 }
103 means.push(acc / n as f64);
104 }
105 means.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
106 (
107 m,
108 quantile_type7(&means, (1.0 - level) / 2.0),
109 quantile_type7(&means, (1.0 + level) / 2.0),
110 )
111 }
112 }
113 }
114}
115
116fn mean(v: &[f64]) -> f64 {
117 v.iter().sum::<f64>() / v.len() as f64
118}
119
120fn sd(v: &[f64]) -> f64 {
122 let n = v.len();
123 if n < 2 {
124 return 0.0;
125 }
126 let m = mean(v);
127 (v.iter().map(|x| (x - m).powi(2)).sum::<f64>() / (n as f64 - 1.0)).sqrt()
128}
129
130fn quantile_type7(sorted: &[f64], p: f64) -> f64 {
132 let n = sorted.len();
133 if n == 0 {
134 return f64::NAN;
135 }
136 if n == 1 {
137 return sorted[0];
138 }
139 let h = (n as f64 - 1.0) * p;
140 let lo = h.floor() as usize;
141 let hi = (lo + 1).min(n - 1);
142 sorted[lo] + (h - lo as f64) * (sorted[hi] - sorted[lo])
143}
144
145pub struct StatSummary {
148 pub fun_y: SummaryFun,
149 pub fun_ymin: SummaryFun,
150 pub fun_ymax: SummaryFun,
151 pub fun_data: Option<SummaryData>,
154}
155
156impl Default for StatSummary {
157 fn default() -> Self {
158 StatSummary {
159 fun_y: SummaryFun::Mean,
160 fun_ymin: SummaryFun::Min,
161 fun_ymax: SummaryFun::Max,
162 fun_data: None,
163 }
164 }
165}
166
167impl StatSummary {
168 fn with_data(d: SummaryData) -> Self {
169 StatSummary {
170 fun_data: Some(d),
171 ..Default::default()
172 }
173 }
174 pub fn mean_se() -> Self {
176 Self::with_data(SummaryData::MeanSe)
177 }
178 pub fn mean_cl_normal() -> Self {
180 Self::with_data(SummaryData::MeanClNormal { level: 0.95 })
181 }
182 pub fn mean_cl_boot() -> Self {
184 Self::with_data(SummaryData::MeanClBoot {
185 level: 0.95,
186 b: 1000,
187 })
188 }
189 pub fn mean_sdl() -> Self {
191 Self::with_data(SummaryData::MeanSdl { mult: 2.0 })
192 }
193 pub fn median_hilow() -> Self {
195 Self::with_data(SummaryData::MedianHilow { level: 0.95 })
196 }
197}
198
199impl Stat for StatSummary {
200 fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
201 let x_col = match data.column("x") {
202 Some(c) => c,
203 None => return DataFrame::new(),
204 };
205 let y_col = match data.column("y") {
206 Some(c) => c,
207 None => return DataFrame::new(),
208 };
209
210 let mut groups: Vec<(String, Value, Vec<f64>)> = Vec::new();
212 for (x, y) in x_col.iter().zip(y_col.iter()) {
213 let key = x.to_group_key();
214 let y_val = y.as_f64().unwrap_or(0.0);
215 if let Some(entry) = groups.iter_mut().find(|(k, _, _)| k == &key) {
216 entry.2.push(y_val);
217 } else {
218 groups.push((key, x.clone(), vec![y_val]));
219 }
220 }
221
222 let n = groups.len();
223 let mut x_vals = Vec::with_capacity(n);
224 let mut y_vals = Vec::with_capacity(n);
225 let mut ymin_vals = Vec::with_capacity(n);
226 let mut ymax_vals = Vec::with_capacity(n);
227
228 for (_, x_val, ys) in &groups {
229 x_vals.push(x_val.clone());
230 let (y, ymin, ymax) = match &self.fun_data {
231 Some(fd) => fd.apply3(ys),
232 None => (
233 self.fun_y.apply(ys),
234 self.fun_ymin.apply(ys),
235 self.fun_ymax.apply(ys),
236 ),
237 };
238 y_vals.push(Value::Float(y));
239 ymin_vals.push(Value::Float(ymin));
240 ymax_vals.push(Value::Float(ymax));
241 }
242
243 let mut result = DataFrame::new();
244 result.add_column("x".to_string(), x_vals);
245 result.add_column("y".to_string(), y_vals);
246 result.add_column("ymin".to_string(), ymin_vals);
247 result.add_column("ymax".to_string(), ymax_vals);
248
249 for col_name in &["color", "fill", "group"] {
251 if let Some(col) = data.column(col_name) {
252 if let Some(first) = col.first() {
253 result.add_column(col_name.to_string(), vec![first.clone(); n]);
254 }
255 }
256 }
257
258 result
259 }
260
261 fn required_aes(&self) -> Vec<Aesthetic> {
262 vec![Aesthetic::X, Aesthetic::Y]
263 }
264
265 fn name(&self) -> &str {
266 "summary"
267 }
268}
269
270#[cfg(test)]
271mod tests {
272 use super::SummaryData;
273
274 #[test]
277 fn summary_data_matches_r() {
278 let v = [2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0];
279 let close = |a: f64, b: f64| (a - b).abs() < 1e-5;
280
281 let (y, lo, hi) = SummaryData::MeanSe.apply3(&v);
282 assert!(close(y, 5.0) && close(lo, 4.244071) && close(hi, 5.755929));
283
284 #[cfg(feature = "regression")]
288 {
289 let (y, lo, hi) = SummaryData::MeanClNormal { level: 0.95 }.apply3(&v);
290 assert!(close(y, 5.0) && close(lo, 3.212512) && close(hi, 6.787488));
291 }
292
293 let (y, lo, hi) = SummaryData::MeanSdl { mult: 2.0 }.apply3(&v);
294 assert!(close(y, 5.0) && close(lo, 0.723820) && close(hi, 9.276180));
295
296 let (y, lo, hi) = SummaryData::MedianHilow { level: 0.95 }.apply3(&v);
297 assert!(close(y, 4.5) && close(lo, 2.35) && close(hi, 8.65));
298
299 let (y, lo, hi) = SummaryData::MeanClBoot {
302 level: 0.95,
303 b: 1000,
304 }
305 .apply3(&v);
306 assert!(close(y, 5.0) && lo < 5.0 && hi > 5.0 && lo > 3.0 && hi < 7.0);
307 }
308}