Skip to main content

ggplot_rs/stat/
qq.rs

1use crate::aes::Aesthetic;
2use crate::data::{DataFrame, Value};
3use crate::scale::ScaleSet;
4
5use super::Stat;
6use crate::stat::dist::qnorm;
7
8/// R-compatible type-7 quantile interpolation (R's default `quantile()` method).
9fn quantile_type7(sorted: &[f64], p: f64) -> f64 {
10    let n = sorted.len();
11    if n == 0 {
12        return 0.0;
13    }
14    if n == 1 {
15        return sorted[0];
16    }
17    let h = (n - 1) as f64 * p;
18    let lo = h.floor() as usize;
19    let hi = (lo + 1).min(n - 1);
20    let frac = h - lo as f64;
21    sorted[lo] + frac * (sorted[hi] - sorted[lo])
22}
23
24/// StatQQ: sort sample, compute theoretical normal quantiles.
25/// Output: x (theoretical quantiles), y (sample sorted).
26pub struct StatQQ;
27
28impl Stat for StatQQ {
29    fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
30        let y_col = match data.column("y") {
31            Some(c) => c,
32            None => return DataFrame::new(),
33        };
34
35        let mut values: Vec<f64> = y_col.iter().filter_map(|v| v.as_f64()).collect();
36        if values.is_empty() {
37            return DataFrame::new();
38        }
39
40        values.sort_by(|a, b| a.total_cmp(b));
41        let n = values.len();
42
43        let mut x_vals = Vec::with_capacity(n);
44        let mut y_vals = Vec::with_capacity(n);
45
46        for (i, &val) in values.iter().enumerate() {
47            // R's ppoints(): (i + 1 - a) / (n + 1 - 2*a) where a = 3/8 for n <= 10 (matches R's ppoints)
48            let a = if n <= 10 { 3.0 / 8.0 } else { 0.5 };
49            let p = (i as f64 + 1.0 - a) / (n as f64 + 1.0 - 2.0 * a);
50            let theoretical = qnorm(p);
51            x_vals.push(Value::Float(theoretical));
52            y_vals.push(Value::Float(val));
53        }
54
55        let mut result = DataFrame::new();
56        result.add_column("x".to_string(), x_vals);
57        result.add_column("y".to_string(), y_vals);
58
59        // Carry over grouping columns
60        for col_name in &["color", "fill", "group"] {
61            if let Some(col) = data.column(col_name) {
62                if let Some(first) = col.first() {
63                    result.add_column(col_name.to_string(), vec![first.clone(); n]);
64                }
65            }
66        }
67
68        result
69    }
70
71    fn required_aes(&self) -> Vec<Aesthetic> {
72        vec![Aesthetic::Y]
73    }
74
75    fn name(&self) -> &str {
76        "qq"
77    }
78}
79
80/// StatQQLine: fit line through Q1/Q3 of sample vs theoretical.
81/// Output: x, y (two points defining the reference line).
82pub struct StatQQLine;
83
84impl Stat for StatQQLine {
85    fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
86        let y_col = match data.column("y") {
87            Some(c) => c,
88            None => return DataFrame::new(),
89        };
90
91        let mut values: Vec<f64> = y_col.iter().filter_map(|v| v.as_f64()).collect();
92        if values.len() < 4 {
93            return DataFrame::new();
94        }
95
96        values.sort_by(|a, b| a.total_cmp(b));
97        let n = values.len();
98
99        // Sample Q1 and Q3 using R-compatible type-7 quantile interpolation
100        let sample_q1 = quantile_type7(&values, 0.25);
101        let sample_q3 = quantile_type7(&values, 0.75);
102
103        // Theoretical Q1 and Q3
104        let theo_q1 = qnorm(0.25);
105        let theo_q3 = qnorm(0.75);
106
107        // Line through (theo_q1, sample_q1) and (theo_q3, sample_q3)
108        let slope = (sample_q3 - sample_q1) / (theo_q3 - theo_q1);
109        let intercept = sample_q1 - slope * theo_q1;
110
111        // Extend line to cover full theoretical range using R's ppoints formula
112        let a = if n <= 10 { 3.0 / 8.0 } else { 0.5 };
113        let x_min = qnorm((1.0 - a) / (n as f64 + 1.0 - 2.0 * a));
114        let x_max = qnorm((n as f64 - a) / (n as f64 + 1.0 - 2.0 * a));
115
116        let mut result = DataFrame::new();
117        result.add_column(
118            "x".to_string(),
119            vec![Value::Float(x_min), Value::Float(x_max)],
120        );
121        result.add_column(
122            "y".to_string(),
123            vec![
124                Value::Float(intercept + slope * x_min),
125                Value::Float(intercept + slope * x_max),
126            ],
127        );
128
129        // Carry over grouping columns
130        for col_name in &["color", "fill", "group"] {
131            if let Some(col) = data.column(col_name) {
132                if let Some(first) = col.first() {
133                    result.add_column(col_name.to_string(), vec![first.clone(); 2]);
134                }
135            }
136        }
137
138        result
139    }
140
141    fn required_aes(&self) -> Vec<Aesthetic> {
142        vec![Aesthetic::Y]
143    }
144
145    fn name(&self) -> &str {
146        "qq_line"
147    }
148}
149
150#[cfg(test)]
151mod tests {
152    use super::*;
153
154    #[test]
155    fn test_qnorm_symmetry() {
156        // `qnorm` now lives in stat::dist (exact under `regression`, A&S
157        // approximation otherwise); check it stays symmetric in either build.
158        let q = qnorm(0.5);
159        assert!((q).abs() < 0.01, "qnorm(0.5) should be ~0, got {q}");
160
161        let q1 = qnorm(0.25);
162        let q3 = qnorm(0.75);
163        assert!((q1 + q3).abs() < 0.01, "qnorm should be symmetric");
164        assert!(q1 < 0.0);
165        assert!(q3 > 0.0);
166    }
167
168    #[test]
169    fn test_stat_qq() {
170        let mut data = DataFrame::new();
171        let y_vals: Vec<Value> = (0..100).map(|i| Value::Float(i as f64)).collect();
172        data.add_column("y".to_string(), y_vals);
173
174        let stat = StatQQ;
175        let scales = ScaleSet::new();
176        let result = stat.compute_group(&data, &scales);
177
178        assert_eq!(result.nrows(), 100);
179        let x = result.column("x").unwrap();
180        let y = result.column("y").unwrap();
181        // y should be sorted
182        for i in 1..y.len() {
183            assert!(y[i].as_f64().unwrap() >= y[i - 1].as_f64().unwrap());
184        }
185        // x should be sorted (theoretical quantiles)
186        for i in 1..x.len() {
187            assert!(x[i].as_f64().unwrap() >= x[i - 1].as_f64().unwrap());
188        }
189    }
190
191    #[test]
192    fn test_stat_qq_line() {
193        let mut data = DataFrame::new();
194        let y_vals: Vec<Value> = (0..100).map(|i| Value::Float(i as f64)).collect();
195        data.add_column("y".to_string(), y_vals);
196
197        let stat = StatQQLine;
198        let scales = ScaleSet::new();
199        let result = stat.compute_group(&data, &scales);
200
201        assert_eq!(result.nrows(), 2);
202    }
203}