Skip to main content

ggplot_rs/stat/
ellipse.rs

1use crate::aes::Aesthetic;
2use crate::data::{DataFrame, Value};
3use crate::scale::ScaleSet;
4
5use super::Stat;
6
7/// Confidence ellipse for a 2-D point cloud (analogous to R's `stat_ellipse`).
8///
9/// Assumes a bivariate normal distribution: the ellipse is the covariance
10/// eigen-decomposition scaled by ggplot2's radius `sqrt(2·F⁻¹(level; 2, n−1))`
11/// (see [`crate::stat::dist::ellipse_radius`]).
12/// Emits `segments + 1` boundary points forming a closed path per group.
13pub struct StatEllipse {
14    /// Confidence level in (0, 1). Default 0.95.
15    pub level: f64,
16    /// Number of segments used to draw the ellipse. Default 51.
17    pub segments: usize,
18}
19
20impl Default for StatEllipse {
21    fn default() -> Self {
22        StatEllipse {
23            level: 0.95,
24            segments: 51,
25        }
26    }
27}
28
29impl StatEllipse {
30    pub fn new(level: f64) -> Self {
31        StatEllipse {
32            level,
33            ..Default::default()
34        }
35    }
36}
37
38impl Stat for StatEllipse {
39    fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
40        let (xs, ys) = match (data.column("x"), data.column("y")) {
41            (Some(x), Some(y)) => (x, y),
42            _ => return DataFrame::new(),
43        };
44        let pts: Vec<(f64, f64)> = xs
45            .iter()
46            .zip(ys.iter())
47            .filter_map(|(a, b)| Some((a.as_f64()?, b.as_f64()?)))
48            .collect();
49        if pts.len() < 3 {
50            return DataFrame::new();
51        }
52
53        let n = pts.len() as f64;
54        let mx = pts.iter().map(|p| p.0).sum::<f64>() / n;
55        let my = pts.iter().map(|p| p.1).sum::<f64>() / n;
56
57        // Sample covariance (n - 1 denominator).
58        let mut sxx = 0.0;
59        let mut syy = 0.0;
60        let mut sxy = 0.0;
61        for &(x, y) in &pts {
62            sxx += (x - mx) * (x - mx);
63            syy += (y - my) * (y - my);
64            sxy += (x - mx) * (y - my);
65        }
66        let d = n - 1.0;
67        let (sxx, syy, sxy) = (sxx / d, syy / d, sxy / d);
68
69        // Eigen-decomposition of the symmetric 2x2 [[sxx, sxy], [sxy, syy]].
70        let trace = sxx + syy;
71        let det = sxx * syy - sxy * sxy;
72        let disc = ((trace * 0.5).powi(2) - det).max(0.0).sqrt();
73        let l1 = (trace * 0.5 + disc).max(0.0);
74        let l2 = (trace * 0.5 - disc).max(0.0);
75        let (v1x, v1y) = if sxy.abs() > 1e-12 {
76            let vx = l1 - syy;
77            let vy = sxy;
78            let norm = (vx * vx + vy * vy).sqrt();
79            (vx / norm, vy / norm)
80        } else if sxx >= syy {
81            (1.0, 0.0)
82        } else {
83            (0.0, 1.0)
84        };
85        // Second axis is perpendicular to the first.
86        let (v2x, v2y) = (-v1y, v1x);
87
88        // Ellipse radius scaling (ggplot2 stat_ellipse: sqrt(2·F⁻¹(level;2,n−1)),
89        // or the χ²₂ limit without the `regression` feature). Distribution math
90        // lives in stat::dist, not here.
91        let radius = crate::stat::dist::ellipse_radius(self.level, pts.len());
92        let a = radius * l1.sqrt();
93        let b = radius * l2.sqrt();
94
95        let steps = self.segments.max(3);
96        let mut x_vals = Vec::with_capacity(steps + 1);
97        let mut y_vals = Vec::with_capacity(steps + 1);
98        for i in 0..=steps {
99            let theta = 2.0 * std::f64::consts::PI * (i as f64) / (steps as f64);
100            let (c, s) = (theta.cos(), theta.sin());
101            let px = mx + a * c * v1x + b * s * v2x;
102            let py = my + a * c * v1y + b * s * v2y;
103            x_vals.push(Value::Float(px));
104            y_vals.push(Value::Float(py));
105        }
106
107        let nrows = x_vals.len();
108        let mut result = DataFrame::new();
109        result.add_column("x".to_string(), x_vals);
110        result.add_column("y".to_string(), y_vals);
111        for col_name in &["color", "fill", "group"] {
112            if let Some(col) = data.column(col_name) {
113                if let Some(first) = col.first() {
114                    result.add_column(col_name.to_string(), vec![first.clone(); nrows]);
115                }
116            }
117        }
118        result
119    }
120
121    fn required_aes(&self) -> Vec<Aesthetic> {
122        vec![Aesthetic::X, Aesthetic::Y]
123    }
124
125    fn name(&self) -> &str {
126        "ellipse"
127    }
128}
129
130#[cfg(test)]
131mod tests {
132    use super::*;
133
134    fn frame(pts: &[(f64, f64)]) -> DataFrame {
135        let mut df = DataFrame::new();
136        df.add_column("x".into(), pts.iter().map(|p| Value::Float(p.0)).collect());
137        df.add_column("y".into(), pts.iter().map(|p| Value::Float(p.1)).collect());
138        df
139    }
140
141    #[test]
142    fn ellipse_of_circular_cloud_is_centered() {
143        // A symmetric ring of points → ellipse centred at the mean.
144        let pts: Vec<(f64, f64)> = (0..40)
145            .map(|i| {
146                let t = 2.0 * std::f64::consts::PI * i as f64 / 40.0;
147                (5.0 + t.cos(), 3.0 + t.sin())
148            })
149            .collect();
150        let out = StatEllipse::default().compute_group(&frame(&pts), &ScaleSet::new());
151        assert_eq!(out.nrows(), StatEllipse::default().segments + 1);
152        let xs: Vec<f64> = out
153            .column("x")
154            .unwrap()
155            .iter()
156            .filter_map(|v| v.as_f64())
157            .collect();
158        let ys: Vec<f64> = out
159            .column("y")
160            .unwrap()
161            .iter()
162            .filter_map(|v| v.as_f64())
163            .collect();
164        let cx = xs.iter().sum::<f64>() / xs.len() as f64;
165        let cy = ys.iter().sum::<f64>() / ys.len() as f64;
166        assert!((cx - 5.0).abs() < 0.2, "center x {cx}");
167        assert!((cy - 3.0).abs() < 0.2, "center y {cy}");
168        // Closed path: first point equals last.
169        assert!((xs[0] - xs[xs.len() - 1]).abs() < 1e-9);
170    }
171
172    #[test]
173    fn too_few_points_returns_empty() {
174        let out = StatEllipse::default()
175            .compute_group(&frame(&[(0.0, 0.0), (1.0, 1.0)]), &ScaleSet::new());
176        assert_eq!(out.nrows(), 0);
177    }
178
179    #[test]
180    fn higher_level_makes_larger_ellipse() {
181        let pts: Vec<(f64, f64)> = (0..30)
182            .map(|i| (i as f64, (i as f64 * 0.7).sin() * 3.0))
183            .collect();
184        let small = StatEllipse::new(0.5).compute_group(&frame(&pts), &ScaleSet::new());
185        let big = StatEllipse::new(0.99).compute_group(&frame(&pts), &ScaleSet::new());
186        let span = |df: &DataFrame| {
187            let xs: Vec<f64> = df
188                .column("x")
189                .unwrap()
190                .iter()
191                .filter_map(|v| v.as_f64())
192                .collect();
193            xs.iter().cloned().fold(f64::MIN, f64::max)
194                - xs.iter().cloned().fold(f64::MAX, f64::min)
195        };
196        assert!(span(&big) > span(&small));
197    }
198}