1use crate::aes::Aesthetic;
12use crate::data::{DataFrame, Value};
13use crate::scale::ScaleSet;
14
15use super::distribution as d;
16use super::Stat;
17
18fn quantile_type7(sorted: &[f64], p: f64) -> f64 {
20 let n = sorted.len();
21 if n == 0 {
22 return 0.0;
23 }
24 if n == 1 {
25 return sorted[0];
26 }
27 let h = (n - 1) as f64 * p;
28 let lo = h.floor() as usize;
29 let hi = (lo + 1).min(n - 1);
30 let frac = h - lo as f64;
31 sorted[lo] + frac * (sorted[hi] - sorted[lo])
32}
33
34pub fn ppoints(n: usize) -> Vec<f64> {
36 let a = if n <= 10 { 3.0 / 8.0 } else { 0.5 };
37 (0..n)
38 .map(|i| (i as f64 + 1.0 - a) / (n as f64 + 1.0 - 2.0 * a))
39 .collect()
40}
41
42#[derive(Clone, Copy, Debug, PartialEq)]
45pub enum QQDistribution {
46 Normal { mean: f64, sd: f64 },
48 StudentT { df: f64 },
50 Exponential { rate: f64 },
52 HalfNormal { sd: f64 },
55}
56
57impl Default for QQDistribution {
58 fn default() -> Self {
59 QQDistribution::normal()
60 }
61}
62
63impl QQDistribution {
64 pub fn normal() -> Self {
66 QQDistribution::Normal { mean: 0.0, sd: 1.0 }
67 }
68 pub fn t(df: f64) -> Self {
70 QQDistribution::StudentT { df }
71 }
72 pub fn exponential() -> Self {
74 QQDistribution::Exponential { rate: 1.0 }
75 }
76 pub fn half_normal() -> Self {
78 QQDistribution::HalfNormal { sd: 1.0 }
79 }
80
81 pub fn quantile(&self, p: f64) -> f64 {
83 match *self {
84 QQDistribution::Normal { mean, sd } if sd > 0.0 => mean + sd * d::qnorm(p),
85 QQDistribution::StudentT { df } if df > 0.0 => d::qt(p, df),
86 QQDistribution::Exponential { rate } if rate > 0.0 => {
87 if !(0.0..=1.0).contains(&p) {
88 f64::NAN
89 } else {
90 -(-p).ln_1p() / rate
91 }
92 }
93 QQDistribution::HalfNormal { sd } if sd > 0.0 => {
94 if !(0.0..=1.0).contains(&p) {
95 f64::NAN
96 } else {
97 sd * d::qnorm((1.0 + p) / 2.0)
98 }
99 }
100 _ => f64::NAN,
101 }
102 }
103
104 pub fn density(&self, x: f64) -> f64 {
106 match *self {
107 QQDistribution::Normal { mean, sd } if sd > 0.0 => d::dnorm((x - mean) / sd) / sd,
108 QQDistribution::StudentT { df } if df > 0.0 => d::dt(x, df),
109 QQDistribution::Exponential { rate } if rate > 0.0 => {
110 if x < 0.0 {
111 0.0
112 } else {
113 rate * (-rate * x).exp()
114 }
115 }
116 QQDistribution::HalfNormal { sd } if sd > 0.0 => {
117 if x < 0.0 {
118 0.0
119 } else {
120 2.0 * d::dnorm(x / sd) / sd
121 }
122 }
123 _ => f64::NAN,
124 }
125 }
126
127 fn label(&self) -> &'static str {
128 match self {
129 QQDistribution::Normal { .. } => "norm",
130 QQDistribution::StudentT { .. } => "t",
131 QQDistribution::Exponential { .. } => "exp",
132 QQDistribution::HalfNormal { .. } => "halfnorm",
133 }
134 }
135}
136
137fn sorted_sample(data: &DataFrame) -> Vec<f64> {
139 let mut v: Vec<f64> = data
140 .column("y")
141 .map(|c| {
142 c.iter()
143 .filter_map(|v| v.as_f64())
144 .filter(|v| v.is_finite())
145 .collect()
146 })
147 .unwrap_or_default();
148 v.sort_by(|a, b| a.total_cmp(b));
149 v
150}
151
152fn carry_groups(data: &DataFrame, out: &mut DataFrame, n: usize) {
154 for col_name in ["color", "fill", "group", "linetype"] {
155 if let Some(first) = data.column(col_name).and_then(|c| c.first()) {
156 out.add_column(col_name.to_string(), vec![first.clone(); n]);
157 }
158 }
159}
160
161fn floats(v: impl IntoIterator<Item = f64>) -> Vec<Value> {
162 v.into_iter().map(Value::Float).collect()
163}
164
165fn qq_line_coef(sorted: &[f64], dist: &QQDistribution, line_p: (f64, f64)) -> (f64, f64) {
168 let (x1, x2) = (dist.quantile(line_p.0), dist.quantile(line_p.1));
169 let (y1, y2) = (
170 quantile_type7(sorted, line_p.0),
171 quantile_type7(sorted, line_p.1),
172 );
173 let slope = (y2 - y1) / (x2 - x1);
174 (slope, y1 - slope * x1)
175}
176
177#[derive(Clone, Debug, Default)]
180pub struct StatQQDist {
181 pub distribution: QQDistribution,
182}
183
184impl StatQQDist {
185 pub fn new(distribution: QQDistribution) -> Self {
186 StatQQDist { distribution }
187 }
188}
189
190impl Stat for StatQQDist {
191 fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
192 let values = sorted_sample(data);
193 if values.is_empty() {
194 return DataFrame::new();
195 }
196 let n = values.len();
197 let theo: Vec<f64> = ppoints(n)
198 .into_iter()
199 .map(|p| self.distribution.quantile(p))
200 .collect();
201 let mut result = DataFrame::new();
202 result.add_column("x".to_string(), floats(theo.iter().copied()));
203 result.add_column("y".to_string(), floats(values.iter().copied()));
204 result.add_column("theoretical".to_string(), floats(theo));
205 result.add_column("sample".to_string(), floats(values));
206 carry_groups(data, &mut result, n);
207 result
208 }
209
210 fn required_aes(&self) -> Vec<Aesthetic> {
211 vec![Aesthetic::Y]
212 }
213
214 fn name(&self) -> &str {
215 "qq"
216 }
217}
218
219#[derive(Clone, Debug)]
223pub struct StatQQLineDist {
224 pub distribution: QQDistribution,
225 pub line_p: (f64, f64),
228}
229
230impl Default for StatQQLineDist {
231 fn default() -> Self {
232 StatQQLineDist {
233 distribution: QQDistribution::default(),
234 line_p: (0.25, 0.75),
235 }
236 }
237}
238
239impl StatQQLineDist {
240 pub fn new(distribution: QQDistribution) -> Self {
241 StatQQLineDist {
242 distribution,
243 ..Default::default()
244 }
245 }
246}
247
248impl Stat for StatQQLineDist {
249 fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
250 let values = sorted_sample(data);
251 let n = values.len();
252 if n < 2 {
253 return DataFrame::new();
254 }
255 let (slope, intercept) = qq_line_coef(&values, &self.distribution, self.line_p);
256 let pp = ppoints(n);
257 let x_min = self.distribution.quantile(pp[0]);
258 let x_max = self.distribution.quantile(pp[n - 1]);
259 let mut result = DataFrame::new();
260 result.add_column("x".to_string(), floats([x_min, x_max]));
261 result.add_column(
262 "y".to_string(),
263 floats([intercept + slope * x_min, intercept + slope * x_max]),
264 );
265 result.add_column("slope".to_string(), floats([slope, slope]));
266 result.add_column("intercept".to_string(), floats([intercept, intercept]));
267 carry_groups(data, &mut result, 2);
268 result
269 }
270
271 fn required_aes(&self) -> Vec<Aesthetic> {
272 vec![Aesthetic::Y]
273 }
274
275 fn name(&self) -> &str {
276 "qq_line"
277 }
278}
279
280#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
282pub enum QQBandType {
283 #[default]
285 Pointwise,
286 Ks,
290}
291
292#[derive(Clone, Debug)]
296pub struct StatQQBand {
297 pub distribution: QQDistribution,
298 pub band: QQBandType,
299 pub level: f64,
301 pub line_p: (f64, f64),
303}
304
305impl Default for StatQQBand {
306 fn default() -> Self {
307 StatQQBand {
308 distribution: QQDistribution::default(),
309 band: QQBandType::Pointwise,
310 level: 0.95,
311 line_p: (0.25, 0.75),
312 }
313 }
314}
315
316impl StatQQBand {
317 pub fn new(distribution: QQDistribution) -> Self {
318 StatQQBand {
319 distribution,
320 ..Default::default()
321 }
322 }
323 pub fn band(mut self, band: QQBandType) -> Self {
325 self.band = band;
326 self
327 }
328 pub fn level(mut self, level: f64) -> Self {
330 self.level = level;
331 self
332 }
333}
334
335impl Stat for StatQQBand {
336 fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
337 let values = sorted_sample(data);
338 let n = values.len();
339 if n < 2 || !(self.level > 0.0 && self.level < 1.0) {
340 return DataFrame::new();
341 }
342 let dist = &self.distribution;
343 let (slope, intercept) = qq_line_coef(&values, dist, self.line_p);
344 let probs = ppoints(n);
345 let nf = n as f64;
346 let mut xs = Vec::with_capacity(n);
347 let mut fit = Vec::with_capacity(n);
348 let mut lo = Vec::with_capacity(n);
349 let mut hi = Vec::with_capacity(n);
350 let z = d::qnorm(1.0 - (1.0 - self.level) / 2.0);
351 let eps = ((2.0 / (1.0 - self.level)).ln() / (2.0 * nf)).sqrt();
352 for &p in &probs {
353 let q = dist.quantile(p);
354 let fitted = intercept + slope * q;
355 let (l, u) = match self.band {
356 QQBandType::Pointwise => {
357 let se = slope / dist.density(q) * (p * (1.0 - p) / nf).sqrt();
358 (fitted - z * se, fitted + z * se)
359 }
360 QQBandType::Ks => {
361 let lp = (p - eps).max(0.0);
362 let up = (p + eps).min(1.0);
363 (
364 intercept + slope * dist.quantile(lp),
365 intercept + slope * dist.quantile(up),
366 )
367 }
368 };
369 let (l, u) = if l <= u { (l, u) } else { (u, l) };
371 xs.push(q);
372 fit.push(fitted);
373 lo.push(l);
374 hi.push(u);
375 }
376 let mut result = DataFrame::new();
377 result.add_column("x".to_string(), floats(xs));
378 result.add_column("y".to_string(), floats(fit));
379 result.add_column("ymin".to_string(), floats(lo));
380 result.add_column("ymax".to_string(), floats(hi));
381 carry_groups(data, &mut result, n);
382 result
383 }
384
385 fn required_aes(&self) -> Vec<Aesthetic> {
386 vec![Aesthetic::Y]
387 }
388
389 fn name(&self) -> &str {
390 match self.band {
391 QQBandType::Pointwise => "qq_band",
392 QQBandType::Ks => "qq_band_ks",
393 }
394 }
395}
396
397pub struct StatQQ;
401
402impl Stat for StatQQ {
403 fn compute_group(&self, data: &DataFrame, scales: &ScaleSet) -> DataFrame {
404 StatQQDist::default().compute_group(data, scales)
405 }
406
407 fn required_aes(&self) -> Vec<Aesthetic> {
408 vec![Aesthetic::Y]
409 }
410
411 fn name(&self) -> &str {
412 "qq"
413 }
414}
415
416pub struct StatQQLine;
419
420impl Stat for StatQQLine {
421 fn compute_group(&self, data: &DataFrame, scales: &ScaleSet) -> DataFrame {
422 StatQQLineDist::default().compute_group(data, scales)
423 }
424
425 fn required_aes(&self) -> Vec<Aesthetic> {
426 vec![Aesthetic::Y]
427 }
428
429 fn name(&self) -> &str {
430 "qq_line"
431 }
432}
433
434impl std::fmt::Display for QQDistribution {
435 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
436 f.write_str(self.label())
437 }
438}
439
440#[cfg(test)]
441mod tests {
442 use super::*;
443
444 fn sample(vals: &[f64]) -> DataFrame {
445 let mut data = DataFrame::new();
446 data.add_column("y".to_string(), floats(vals.iter().copied()));
447 data
448 }
449
450 #[test]
451 fn test_stat_qq() {
452 let data = sample(&(0..100).map(|i| i as f64).collect::<Vec<_>>());
453 let result = StatQQ.compute_group(&data, &ScaleSet::new());
454 assert_eq!(result.nrows(), 100);
455 let x = result.column("x").unwrap();
456 let y = result.column("y").unwrap();
457 for i in 1..y.len() {
458 assert!(y[i].as_f64().unwrap() >= y[i - 1].as_f64().unwrap());
459 assert!(x[i].as_f64().unwrap() >= x[i - 1].as_f64().unwrap());
460 }
461 }
462
463 #[test]
464 fn test_stat_qq_line() {
465 let data = sample(&(0..100).map(|i| i as f64).collect::<Vec<_>>());
466 let result = StatQQLine.compute_group(&data, &ScaleSet::new());
467 assert_eq!(result.nrows(), 2);
468 }
469
470 #[test]
471 fn quantiles_of_each_distribution() {
472 assert!((QQDistribution::exponential().quantile(0.5) - 2f64.ln()).abs() < 1e-15);
473 let hn = QQDistribution::half_normal();
474 assert!((hn.quantile(0.5) - d::qnorm(0.75)).abs() < 1e-15);
475 assert_eq!(hn.quantile(0.0), 0.0);
476 assert!(QQDistribution::t(0.0).quantile(0.3).is_nan());
477 let n = QQDistribution::Normal { mean: 2.0, sd: 3.0 };
478 assert!((n.quantile(0.975) - (2.0 + 3.0 * 1.959_963_984_540_054)).abs() < 1e-12);
479 }
480
481 #[test]
482 fn band_contains_line_and_ks_reaches_infinity() {
483 let vals: Vec<f64> = (1..=30).map(|i| (i as f64 * 0.37).sin() * 2.0).collect();
484 let data = sample(&vals);
485 let pw = StatQQBand::default().compute_group(&data, &ScaleSet::new());
486 let (lo, mid, hi) = (
487 pw.column("ymin").unwrap(),
488 pw.column("y").unwrap(),
489 pw.column("ymax").unwrap(),
490 );
491 for i in 0..pw.nrows() {
492 let (l, m, h) = (
493 lo[i].as_f64().unwrap(),
494 mid[i].as_f64().unwrap(),
495 hi[i].as_f64().unwrap(),
496 );
497 assert!(l < m && m < h);
498 }
499 let ks = StatQQBand::default()
500 .band(QQBandType::Ks)
501 .compute_group(&data, &ScaleSet::new());
502 let lo = ks.column("ymin").unwrap();
503 assert_eq!(lo[0].as_f64(), Some(f64::NEG_INFINITY));
504 }
505
506 #[test]
507 fn non_finite_and_tiny_samples() {
508 let data = sample(&[f64::NAN, 1.0, f64::INFINITY]);
509 assert_eq!(
510 StatQQDist::default()
511 .compute_group(&data, &ScaleSet::new())
512 .nrows(),
513 1
514 );
515 assert_eq!(
516 StatQQBand::default()
517 .compute_group(&data, &ScaleSet::new())
518 .nrows(),
519 0
520 );
521 }
522}