Skip to main content

ggplot_rs/stat/
smooth.rs

1use crate::aes::Aesthetic;
2use crate::data::{DataFrame, Value};
3use crate::scale::ScaleSet;
4
5use super::Stat;
6
7/// Smoothing method selection.
8#[derive(Clone, Debug, Default)]
9pub enum SmoothMethod {
10    /// Linear regression (y = mx + b).
11    #[default]
12    Lm,
13    /// LOESS with configurable span.
14    Loess { span: f64 },
15    /// Generalized linear model via anofox-regression (Gaussian or Poisson).
16    #[cfg(feature = "regression")]
17    Glm { family: SmoothFamily },
18    /// Robust linear regression (Huber M-estimator) via anofox-regression.
19    #[cfg(feature = "regression")]
20    Rlm,
21    /// Penalized B-spline (P-spline) GAM smoother via anofox-regression —
22    /// ggplot2's `method = "gam"`. λ is chosen by GCV.
23    #[cfg(feature = "regression")]
24    Gam,
25}
26
27/// GLM family for regression-backed smoothing (`SmoothMethod::Glm`).
28#[cfg(feature = "regression")]
29#[derive(Clone, Copy, Debug, Default, PartialEq)]
30pub enum SmoothFamily {
31    /// Ordinary least squares (identity link).
32    #[default]
33    Gaussian,
34    /// Poisson regression (log link) for count responses.
35    Poisson,
36}
37
38/// Smoothing statistic — supports both linear regression and LOESS.
39pub struct StatSmooth {
40    /// Number of points to generate for the fitted line.
41    pub n_points: usize,
42    /// Whether to compute confidence interval.
43    pub se: bool,
44    /// Smoothing method.
45    pub method: SmoothMethod,
46}
47
48impl Default for StatSmooth {
49    fn default() -> Self {
50        StatSmooth {
51            n_points: 80,
52            se: true,
53            method: SmoothMethod::Lm,
54        }
55    }
56}
57
58impl Stat for StatSmooth {
59    fn compute_group(&self, data: &DataFrame, scales: &ScaleSet) -> DataFrame {
60        match &self.method {
61            SmoothMethod::Lm => self.compute_lm(data),
62            SmoothMethod::Loess { span } => {
63                let loess = super::loess::StatLoess {
64                    span: *span,
65                    n_points: self.n_points,
66                    se: self.se,
67                };
68                loess.compute_group(data, scales)
69            }
70            #[cfg(feature = "regression")]
71            SmoothMethod::Glm { family } => self.compute_glm(data, Some(*family)),
72            #[cfg(feature = "regression")]
73            SmoothMethod::Rlm => self.compute_glm(data, None),
74            #[cfg(feature = "regression")]
75            SmoothMethod::Gam => self.compute_gam(data),
76        }
77    }
78
79    fn required_aes(&self) -> Vec<Aesthetic> {
80        vec![Aesthetic::X, Aesthetic::Y]
81    }
82
83    fn name(&self) -> &str {
84        "smooth"
85    }
86}
87
88impl StatSmooth {
89    fn compute_lm(&self, data: &DataFrame) -> DataFrame {
90        let x_col = match data.column("x") {
91            Some(c) => c,
92            None => return DataFrame::new(),
93        };
94        let y_col = match data.column("y") {
95            Some(c) => c,
96            None => return DataFrame::new(),
97        };
98
99        let pairs: Vec<(f64, f64)> = x_col
100            .iter()
101            .zip(y_col.iter())
102            .filter_map(|(x, y)| Some((x.as_f64()?, y.as_f64()?)))
103            .collect();
104
105        if pairs.len() < 2 {
106            return DataFrame::new();
107        }
108
109        let n = pairs.len() as f64;
110        let sum_x: f64 = pairs.iter().map(|(x, _)| x).sum();
111        let sum_y: f64 = pairs.iter().map(|(_, y)| y).sum();
112        let sum_xy: f64 = pairs.iter().map(|(x, y)| x * y).sum();
113        let sum_xx: f64 = pairs.iter().map(|(x, _)| x * x).sum();
114
115        let mean_x = sum_x / n;
116        let mean_y = sum_y / n;
117
118        let denom = sum_xx - sum_x * sum_x / n;
119        let (slope, intercept) = if denom.abs() < f64::EPSILON {
120            (0.0, mean_y)
121        } else {
122            let m = (sum_xy - sum_x * sum_y / n) / denom;
123            let b = mean_y - m * mean_x;
124            (m, b)
125        };
126
127        // Generate fitted values across x range
128        let x_min = pairs.iter().map(|(x, _)| *x).fold(f64::INFINITY, f64::min);
129        let x_max = pairs
130            .iter()
131            .map(|(x, _)| *x)
132            .fold(f64::NEG_INFINITY, f64::max);
133
134        let step = (x_max - x_min) / (self.n_points - 1).max(1) as f64;
135
136        // Compute standard error of prediction if requested
137        let se_values = if self.se && pairs.len() > 2 {
138            let residuals: Vec<f64> = pairs
139                .iter()
140                .map(|(x, y)| y - (slope * x + intercept))
141                .collect();
142            let sse: f64 = residuals.iter().map(|r| r * r).sum();
143            let mse = sse / (n - 2.0);
144            Some((mse, sum_xx, mean_x, n))
145        } else {
146            None
147        };
148
149        let mut x_vals = Vec::with_capacity(self.n_points);
150        let mut y_vals = Vec::with_capacity(self.n_points);
151        let mut ymin_vals = Vec::with_capacity(self.n_points);
152        let mut ymax_vals = Vec::with_capacity(self.n_points);
153
154        for i in 0..self.n_points {
155            let x = x_min + i as f64 * step;
156            let y = slope * x + intercept;
157            x_vals.push(Value::Float(x));
158            y_vals.push(Value::Float(y));
159
160            if let Some((mse, sum_xx, mean_x, n)) = se_values {
161                let se_pred = (mse
162                    * (1.0 / n + (x - mean_x).powi(2) / (sum_xx - n * mean_x * mean_x)))
163                    .sqrt();
164                // 95% confidence interval for the mean response: R's
165                // qt(0.975, n − 2), not the large-sample normal 1.96.
166                let t_val = crate::stat::dist::qt(0.975, (n - 2.0).max(1.0));
167                ymin_vals.push(Value::Float(y - t_val * se_pred));
168                ymax_vals.push(Value::Float(y + t_val * se_pred));
169            }
170        }
171
172        let mut result = DataFrame::new();
173        result.add_column("x".to_string(), x_vals);
174        result.add_column("y".to_string(), y_vals);
175        if !ymin_vals.is_empty() {
176            result.add_column("ymin".to_string(), ymin_vals);
177            result.add_column("ymax".to_string(), ymax_vals);
178        }
179        result
180    }
181
182    /// GLM / robust-linear smoothing backed by anofox-regression. `family = None`
183    /// selects the robust (Huber) fit; `Some(..)` selects a GLM family. A
184    /// confidence-interval ribbon (ymin/ymax) is emitted when `self.se` is set.
185    #[cfg(feature = "regression")]
186    fn compute_glm(&self, data: &DataFrame, family: Option<SmoothFamily>) -> DataFrame {
187        use anofox_regression::solvers::{
188            FittedRegressor, HuberRegressor, OlsRegressor, PoissonRegressor, Regressor,
189        };
190        use anofox_regression::{IntervalType, PoissonFamily, RegressionOptions};
191        use faer::{Col, Mat};
192
193        let (x_col, y_col) = match (data.column("x"), data.column("y")) {
194            (Some(x), Some(y)) => (x, y),
195            _ => return DataFrame::new(),
196        };
197        let pairs: Vec<(f64, f64)> = x_col
198            .iter()
199            .zip(y_col.iter())
200            .filter_map(|(x, y)| Some((x.as_f64()?, y.as_f64()?)))
201            .collect();
202        if pairs.len() < 2 {
203            return DataFrame::new();
204        }
205
206        let n = pairs.len();
207        let x = Mat::from_fn(n, 1, |i, _| pairs[i].0);
208        let y = Col::from_fn(n, |i| pairs[i].1);
209        let x_min = pairs.iter().map(|p| p.0).fold(f64::INFINITY, f64::min);
210        let x_max = pairs.iter().map(|p| p.0).fold(f64::NEG_INFINITY, f64::max);
211        let steps = self.n_points.max(2);
212        let grid = Mat::from_fn(steps, 1, |k, _| {
213            x_min + (x_max - x_min) * k as f64 / (steps - 1) as f64
214        });
215        let interval = if self.se {
216            Some(IntervalType::Confidence)
217        } else {
218            None
219        };
220
221        // Fit the requested model and predict (with interval) over the grid.
222        let pred = match family {
223            None => match HuberRegressor::new().fit(&x, &y) {
224                Ok(f) => f.predict_with_interval(&grid, interval, 0.95),
225                Err(_) => return DataFrame::new(),
226            },
227            Some(SmoothFamily::Gaussian) => {
228                match OlsRegressor::new(RegressionOptions::default()).fit(&x, &y) {
229                    Ok(f) => f.predict_with_interval(&grid, interval, 0.95),
230                    Err(_) => return DataFrame::new(),
231                }
232            }
233            Some(SmoothFamily::Poisson) => {
234                let reg =
235                    PoissonRegressor::new(RegressionOptions::default(), PoissonFamily::default());
236                match reg.fit(&x, &y) {
237                    Ok(f) => f.predict_with_interval(&grid, interval, 0.95),
238                    Err(_) => return DataFrame::new(),
239                }
240            }
241        };
242
243        let mut x_vals = Vec::with_capacity(steps);
244        let mut y_vals = Vec::with_capacity(steps);
245        let mut ymin_vals = Vec::with_capacity(steps);
246        let mut ymax_vals = Vec::with_capacity(steps);
247        for k in 0..steps {
248            x_vals.push(Value::Float(grid[(k, 0)]));
249            y_vals.push(Value::Float(pred.fit[k]));
250            if self.se {
251                ymin_vals.push(Value::Float(pred.lower[k]));
252                ymax_vals.push(Value::Float(pred.upper[k]));
253            }
254        }
255
256        let mut result = DataFrame::new();
257        result.add_column("x".to_string(), x_vals);
258        result.add_column("y".to_string(), y_vals);
259        if self.se {
260            result.add_column("ymin".to_string(), ymin_vals);
261            result.add_column("ymax".to_string(), ymax_vals);
262        }
263        for col_name in &["color", "fill", "group"] {
264            if let Some(col) = data.column(col_name) {
265                if let Some(first) = col.first() {
266                    result.add_column(col_name.to_string(), vec![first.clone(); steps]);
267                }
268            }
269        }
270        result
271    }
272
273    /// GAM smoothing via anofox-regression's penalized B-spline (P-spline) with
274    /// GCV-selected smoothing — ggplot2's `method = "gam"`. Emits a t-based
275    /// confidence ribbon (ymin/ymax) when `self.se` is set.
276    #[cfg(feature = "regression")]
277    fn compute_gam(&self, data: &DataFrame) -> DataFrame {
278        use anofox_regression::solvers::{FittedRegressor, PSplineRegressor, Regressor};
279        use anofox_regression::IntervalType;
280        use faer::{Col, Mat};
281
282        let (x_col, y_col) = match (data.column("x"), data.column("y")) {
283            (Some(x), Some(y)) => (x, y),
284            _ => return DataFrame::new(),
285        };
286        let pairs: Vec<(f64, f64)> = x_col
287            .iter()
288            .zip(y_col.iter())
289            .filter_map(|(x, y)| Some((x.as_f64()?, y.as_f64()?)))
290            .collect();
291        // The P-spline needs enough distinct support to build a basis; fall back
292        // to a straight line for tiny groups.
293        if pairs.len() < 6 {
294            return self.compute_lm(data);
295        }
296
297        let n = pairs.len();
298        let x = Mat::from_fn(n, 1, |i, _| pairs[i].0);
299        let y = Col::from_fn(n, |i| pairs[i].1);
300        let x_min = pairs.iter().map(|p| p.0).fold(f64::INFINITY, f64::min);
301        let x_max = pairs.iter().map(|p| p.0).fold(f64::NEG_INFINITY, f64::max);
302        let steps = self.n_points.max(2);
303        let grid = Mat::from_fn(steps, 1, |k, _| {
304            x_min + (x_max - x_min) * k as f64 / (steps - 1) as f64
305        });
306        let interval = if self.se {
307            Some(IntervalType::Confidence)
308        } else {
309            None
310        };
311
312        let pred = match PSplineRegressor::new().fit(&x, &y) {
313            Ok(f) => f.predict_with_interval(&grid, interval, 0.95),
314            // Degenerate input (e.g. collinear x) — fall back to a line.
315            Err(_) => return self.compute_lm(data),
316        };
317
318        let mut x_vals = Vec::with_capacity(steps);
319        let mut y_vals = Vec::with_capacity(steps);
320        let mut ymin_vals = Vec::with_capacity(steps);
321        let mut ymax_vals = Vec::with_capacity(steps);
322        for k in 0..steps {
323            x_vals.push(Value::Float(grid[(k, 0)]));
324            y_vals.push(Value::Float(pred.fit[k]));
325            if self.se {
326                ymin_vals.push(Value::Float(pred.lower[k]));
327                ymax_vals.push(Value::Float(pred.upper[k]));
328            }
329        }
330
331        let mut result = DataFrame::new();
332        result.add_column("x".to_string(), x_vals);
333        result.add_column("y".to_string(), y_vals);
334        if self.se {
335            result.add_column("ymin".to_string(), ymin_vals);
336            result.add_column("ymax".to_string(), ymax_vals);
337        }
338        for col_name in &["color", "fill", "group"] {
339            if let Some(col) = data.column(col_name) {
340                if let Some(first) = col.first() {
341                    result.add_column(col_name.to_string(), vec![first.clone(); steps]);
342                }
343            }
344        }
345        result
346    }
347}