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`), R's
28/// `glm(family = …)`.
29///
30/// For every non-Gaussian family the confidence band is computed as ggplot2's
31/// `predictdf.glm` does: the linear predictor and its standard error come from
32/// `predict(type = "link", se.fit = TRUE)`, the interval
33/// `η ± qnorm(0.975)·se(η)` is formed on the link scale and both ends are
34/// mapped through the inverse link — so the band respects the response range
35/// (probabilities stay in `[0, 1]`, means stay positive). Dispersion follows
36/// R: fixed at 1 for binomial / Poisson / negative binomial, Pearson-estimated
37/// for Gamma.
38///
39/// `Gaussian` keeps the ordinary-least-squares fit with its `t`-based interval
40/// (identical to `method = "lm"`).
41#[cfg(feature = "regression")]
42#[derive(Clone, Copy, Debug, Default, PartialEq)]
43pub enum SmoothFamily {
44    /// Ordinary least squares (identity link).
45    #[default]
46    Gaussian,
47    /// Poisson regression (log link) for count responses.
48    Poisson,
49    /// Binomial regression for 0/1 (or proportion) responses —
50    /// `binomial(link)`, R's default link is logit.
51    Binomial(SmoothBinomialLink),
52    /// Gamma regression for positive continuous responses — `Gamma(link)`;
53    /// R's default link is the inverse.
54    Gamma(SmoothGammaLink),
55    /// Negative-binomial regression (log link, θ estimated by maximum
56    /// likelihood) for over-dispersed counts — R's `MASS::glm.nb`.
57    NegativeBinomial,
58}
59
60#[cfg(feature = "regression")]
61impl SmoothFamily {
62    /// `binomial("logit")` — logistic regression.
63    pub fn binomial() -> Self {
64        SmoothFamily::Binomial(SmoothBinomialLink::Logit)
65    }
66
67    /// `Gamma("inverse")` — R's canonical Gamma link.
68    pub fn gamma() -> Self {
69        SmoothFamily::Gamma(SmoothGammaLink::Inverse)
70    }
71}
72
73/// Link function for [`SmoothFamily::Binomial`].
74#[cfg(feature = "regression")]
75#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
76pub enum SmoothBinomialLink {
77    /// `log(μ / (1 − μ))` (canonical).
78    #[default]
79    Logit,
80    /// `Φ⁻¹(μ)`.
81    Probit,
82    /// `log(−log(1 − μ))`.
83    Cloglog,
84}
85
86/// Link function for [`SmoothFamily::Gamma`].
87#[cfg(feature = "regression")]
88#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
89pub enum SmoothGammaLink {
90    /// `1 / μ` (R's canonical Gamma link).
91    #[default]
92    Inverse,
93    /// `log(μ)`.
94    Log,
95}
96
97/// Smoothing statistic — supports both linear regression and LOESS.
98pub struct StatSmooth {
99    /// Number of points to generate for the fitted line.
100    pub n_points: usize,
101    /// Whether to compute confidence interval.
102    pub se: bool,
103    /// Smoothing method.
104    pub method: SmoothMethod,
105}
106
107impl Default for StatSmooth {
108    fn default() -> Self {
109        StatSmooth {
110            n_points: 80,
111            se: true,
112            method: SmoothMethod::Lm,
113        }
114    }
115}
116
117impl Stat for StatSmooth {
118    fn compute_group(&self, data: &DataFrame, scales: &ScaleSet) -> DataFrame {
119        match &self.method {
120            SmoothMethod::Lm => self.compute_lm(data),
121            SmoothMethod::Loess { span } => {
122                let loess = super::loess::StatLoess {
123                    span: *span,
124                    n_points: self.n_points,
125                    se: self.se,
126                };
127                loess.compute_group(data, scales)
128            }
129            #[cfg(feature = "regression")]
130            SmoothMethod::Glm { family } => self.compute_glm(data, Some(*family)),
131            #[cfg(feature = "regression")]
132            SmoothMethod::Rlm => self.compute_glm(data, None),
133            #[cfg(feature = "regression")]
134            SmoothMethod::Gam => self.compute_gam(data),
135        }
136    }
137
138    fn required_aes(&self) -> Vec<Aesthetic> {
139        vec![Aesthetic::X, Aesthetic::Y]
140    }
141
142    fn name(&self) -> &str {
143        "smooth"
144    }
145}
146
147impl StatSmooth {
148    fn compute_lm(&self, data: &DataFrame) -> DataFrame {
149        let x_col = match data.column("x") {
150            Some(c) => c,
151            None => return DataFrame::new(),
152        };
153        let y_col = match data.column("y") {
154            Some(c) => c,
155            None => return DataFrame::new(),
156        };
157
158        let pairs: Vec<(f64, f64)> = x_col
159            .iter()
160            .zip(y_col.iter())
161            .filter_map(|(x, y)| Some((x.as_f64()?, y.as_f64()?)))
162            .collect();
163
164        if pairs.len() < 2 {
165            return DataFrame::new();
166        }
167
168        let n = pairs.len() as f64;
169        let sum_x: f64 = pairs.iter().map(|(x, _)| x).sum();
170        let sum_y: f64 = pairs.iter().map(|(_, y)| y).sum();
171        let sum_xy: f64 = pairs.iter().map(|(x, y)| x * y).sum();
172        let sum_xx: f64 = pairs.iter().map(|(x, _)| x * x).sum();
173
174        let mean_x = sum_x / n;
175        let mean_y = sum_y / n;
176
177        let denom = sum_xx - sum_x * sum_x / n;
178        let (slope, intercept) = if denom.abs() < f64::EPSILON {
179            (0.0, mean_y)
180        } else {
181            let m = (sum_xy - sum_x * sum_y / n) / denom;
182            let b = mean_y - m * mean_x;
183            (m, b)
184        };
185
186        // Generate fitted values across x range
187        let x_min = pairs.iter().map(|(x, _)| *x).fold(f64::INFINITY, f64::min);
188        let x_max = pairs
189            .iter()
190            .map(|(x, _)| *x)
191            .fold(f64::NEG_INFINITY, f64::max);
192
193        let step = (x_max - x_min) / (self.n_points - 1).max(1) as f64;
194
195        // Compute standard error of prediction if requested
196        let se_values = if self.se && pairs.len() > 2 {
197            let residuals: Vec<f64> = pairs
198                .iter()
199                .map(|(x, y)| y - (slope * x + intercept))
200                .collect();
201            let sse: f64 = residuals.iter().map(|r| r * r).sum();
202            let mse = sse / (n - 2.0);
203            Some((mse, sum_xx, mean_x, n))
204        } else {
205            None
206        };
207
208        let mut x_vals = Vec::with_capacity(self.n_points);
209        let mut y_vals = Vec::with_capacity(self.n_points);
210        let mut ymin_vals = Vec::with_capacity(self.n_points);
211        let mut ymax_vals = Vec::with_capacity(self.n_points);
212
213        for i in 0..self.n_points {
214            let x = x_min + i as f64 * step;
215            let y = slope * x + intercept;
216            x_vals.push(Value::Float(x));
217            y_vals.push(Value::Float(y));
218
219            if let Some((mse, sum_xx, mean_x, n)) = se_values {
220                let se_pred = (mse
221                    * (1.0 / n + (x - mean_x).powi(2) / (sum_xx - n * mean_x * mean_x)))
222                    .sqrt();
223                // 95% confidence interval for the mean response: R's
224                // qt(0.975, n − 2), not the large-sample normal 1.96.
225                let t_val = crate::stat::dist::qt(0.975, (n - 2.0).max(1.0));
226                ymin_vals.push(Value::Float(y - t_val * se_pred));
227                ymax_vals.push(Value::Float(y + t_val * se_pred));
228            }
229        }
230
231        let mut result = DataFrame::new();
232        result.add_column("x".to_string(), x_vals);
233        result.add_column("y".to_string(), y_vals);
234        if !ymin_vals.is_empty() {
235            result.add_column("ymin".to_string(), ymin_vals);
236            result.add_column("ymax".to_string(), ymax_vals);
237        }
238        result
239    }
240
241    /// GLM / robust-linear smoothing backed by anofox-regression. `family = None`
242    /// selects the robust (Huber) fit; `Some(..)` selects a GLM family. A
243    /// confidence-interval ribbon (ymin/ymax) is emitted when `self.se` is set.
244    #[cfg(feature = "regression")]
245    fn compute_glm(&self, data: &DataFrame, family: Option<SmoothFamily>) -> DataFrame {
246        use anofox_regression::solvers::{
247            FittedRegressor, HuberRegressor, OlsRegressor, Regressor,
248        };
249        use anofox_regression::{IntervalType, RegressionOptions};
250        use faer::{Col, Mat};
251
252        let (x_col, y_col) = match (data.column("x"), data.column("y")) {
253            (Some(x), Some(y)) => (x, y),
254            _ => return DataFrame::new(),
255        };
256        let pairs: Vec<(f64, f64)> = x_col
257            .iter()
258            .zip(y_col.iter())
259            .filter_map(|(x, y)| Some((x.as_f64()?, y.as_f64()?)))
260            .collect();
261        if pairs.len() < 2 {
262            return DataFrame::new();
263        }
264
265        let n = pairs.len();
266        let x = Mat::from_fn(n, 1, |i, _| pairs[i].0);
267        let y = Col::from_fn(n, |i| pairs[i].1);
268        let x_min = pairs.iter().map(|p| p.0).fold(f64::INFINITY, f64::min);
269        let x_max = pairs.iter().map(|p| p.0).fold(f64::NEG_INFINITY, f64::max);
270        let steps = self.n_points.max(2);
271        let grid = Mat::from_fn(steps, 1, |k, _| {
272            x_min + (x_max - x_min) * k as f64 / (steps - 1) as f64
273        });
274        let interval = if self.se {
275            Some(IntervalType::Confidence)
276        } else {
277            None
278        };
279
280        // Fit the requested model and predict (with interval) over the grid.
281        let pred = match family {
282            None => match HuberRegressor::new().fit(&x, &y) {
283                Ok(f) => f.predict_with_interval(&grid, interval, 0.95),
284                Err(_) => return DataFrame::new(),
285            },
286            Some(SmoothFamily::Gaussian) => {
287                match OlsRegressor::new(RegressionOptions::default()).fit(&x, &y) {
288                    Ok(f) => f.predict_with_interval(&grid, interval, 0.95),
289                    Err(_) => return DataFrame::new(),
290                }
291            }
292            Some(family) => match glm_link_prediction(family, &x, &y, &grid, self.se) {
293                Some(p) => p,
294                None => return DataFrame::new(),
295            },
296        };
297
298        let mut x_vals = Vec::with_capacity(steps);
299        let mut y_vals = Vec::with_capacity(steps);
300        let mut ymin_vals = Vec::with_capacity(steps);
301        let mut ymax_vals = Vec::with_capacity(steps);
302        for k in 0..steps {
303            x_vals.push(Value::Float(grid[(k, 0)]));
304            y_vals.push(Value::Float(pred.fit[k]));
305            if self.se {
306                // An inverse link can swap the ends (e.g. Gamma's 1/η).
307                let (a, b) = (pred.lower[k], pred.upper[k]);
308                ymin_vals.push(Value::Float(a.min(b)));
309                ymax_vals.push(Value::Float(a.max(b)));
310            }
311        }
312
313        let mut result = DataFrame::new();
314        result.add_column("x".to_string(), x_vals);
315        result.add_column("y".to_string(), y_vals);
316        if self.se {
317            result.add_column("ymin".to_string(), ymin_vals);
318            result.add_column("ymax".to_string(), ymax_vals);
319        }
320        for col_name in &["color", "fill", "group"] {
321            if let Some(col) = data.column(col_name) {
322                if let Some(first) = col.first() {
323                    result.add_column(col_name.to_string(), vec![first.clone(); steps]);
324                }
325            }
326        }
327        result
328    }
329
330    /// GAM smoothing via anofox-regression's penalized B-spline (P-spline) with
331    /// GCV-selected smoothing — ggplot2's `method = "gam"`. Emits a t-based
332    /// confidence ribbon (ymin/ymax) when `self.se` is set.
333    #[cfg(feature = "regression")]
334    fn compute_gam(&self, data: &DataFrame) -> DataFrame {
335        use anofox_regression::solvers::{FittedRegressor, PSplineRegressor, Regressor};
336        use anofox_regression::IntervalType;
337        use faer::{Col, Mat};
338
339        let (x_col, y_col) = match (data.column("x"), data.column("y")) {
340            (Some(x), Some(y)) => (x, y),
341            _ => return DataFrame::new(),
342        };
343        let pairs: Vec<(f64, f64)> = x_col
344            .iter()
345            .zip(y_col.iter())
346            .filter_map(|(x, y)| Some((x.as_f64()?, y.as_f64()?)))
347            .collect();
348        // The P-spline needs enough distinct support to build a basis; fall back
349        // to a straight line for tiny groups.
350        if pairs.len() < 6 {
351            return self.compute_lm(data);
352        }
353
354        let n = pairs.len();
355        let x = Mat::from_fn(n, 1, |i, _| pairs[i].0);
356        let y = Col::from_fn(n, |i| pairs[i].1);
357        let x_min = pairs.iter().map(|p| p.0).fold(f64::INFINITY, f64::min);
358        let x_max = pairs.iter().map(|p| p.0).fold(f64::NEG_INFINITY, f64::max);
359        let steps = self.n_points.max(2);
360        let grid = Mat::from_fn(steps, 1, |k, _| {
361            x_min + (x_max - x_min) * k as f64 / (steps - 1) as f64
362        });
363        let interval = if self.se {
364            Some(IntervalType::Confidence)
365        } else {
366            None
367        };
368
369        let pred = match PSplineRegressor::new().fit(&x, &y) {
370            Ok(f) => f.predict_with_interval(&grid, interval, 0.95),
371            // Degenerate input (e.g. collinear x) — fall back to a line.
372            Err(_) => return self.compute_lm(data),
373        };
374
375        let mut x_vals = Vec::with_capacity(steps);
376        let mut y_vals = Vec::with_capacity(steps);
377        let mut ymin_vals = Vec::with_capacity(steps);
378        let mut ymax_vals = Vec::with_capacity(steps);
379        for k in 0..steps {
380            x_vals.push(Value::Float(grid[(k, 0)]));
381            y_vals.push(Value::Float(pred.fit[k]));
382            if self.se {
383                ymin_vals.push(Value::Float(pred.lower[k]));
384                ymax_vals.push(Value::Float(pred.upper[k]));
385            }
386        }
387
388        let mut result = DataFrame::new();
389        result.add_column("x".to_string(), x_vals);
390        result.add_column("y".to_string(), y_vals);
391        if self.se {
392            result.add_column("ymin".to_string(), ymin_vals);
393            result.add_column("ymax".to_string(), ymax_vals);
394        }
395        for col_name in &["color", "fill", "group"] {
396            if let Some(col) = data.column(col_name) {
397                if let Some(first) = col.first() {
398                    result.add_column(col_name.to_string(), vec![first.clone(); steps]);
399                }
400            }
401        }
402        result
403    }
404}
405
406/// IRLS convergence settings shared with the anofox-statistics DuckDB
407/// extension's GLM aggregates (R's `glm.control(epsilon = 1e-8)`), so a fit in
408/// SQL and the plotted smooth agree numerically.
409#[cfg(feature = "regression")]
410const GLM_TOL: f64 = 1e-8;
411#[cfg(feature = "regression")]
412const GLM_MAX_ITER: usize = 100;
413
414/// Fitted mean and (optionally) a 95% confidence band for a non-Gaussian GLM
415/// family over `grid`, computed like ggplot2's `predictdf.glm`: the band is
416/// `linkinv(η ± qnorm(0.975)·se(η))`. `None` when the fit fails (degenerate
417/// or out-of-range responses, e.g. a negative count).
418#[cfg(feature = "regression")]
419fn glm_link_prediction(
420    family: SmoothFamily,
421    x: &faer::Mat<f64>,
422    y: &faer::Col<f64>,
423    grid: &faer::Mat<f64>,
424    se: bool,
425) -> Option<anofox_regression::PredictionResult> {
426    use anofox_regression::core::PredictionType;
427    use anofox_regression::solvers::{
428        BinomialRegressor, GammaRegressor, NegativeBinomialRegressor, PoissonRegressor, Regressor,
429    };
430    use anofox_regression::{BinomialLink, PredictionResult};
431    use faer::Col;
432
433    // Link-scale prediction (η and se(η)) plus the inverse link.
434    let (link_pred, linkinv): (PredictionResult, Box<dyn Fn(f64) -> f64>) = match family {
435        SmoothFamily::Gaussian => return None,
436        SmoothFamily::Poisson => {
437            let f = PoissonRegressor::log()
438                .tolerance(GLM_TOL)
439                .max_iterations(GLM_MAX_ITER)
440                .build()
441                .fit(x, y)
442                .ok()?;
443            (
444                f.predict_with_se(grid, PredictionType::Link, None, 0.95),
445                Box::new(f64::exp),
446            )
447        }
448        SmoothFamily::NegativeBinomial => {
449            let f = NegativeBinomialRegressor::builder()
450                .tolerance(GLM_TOL)
451                .max_iterations(GLM_MAX_ITER)
452                .build()
453                .fit(x, y)
454                .ok()?;
455            (
456                f.predict_with_se(grid, PredictionType::Link, None, 0.95),
457                Box::new(f64::exp),
458            )
459        }
460        SmoothFamily::Binomial(link) => {
461            let link = match link {
462                SmoothBinomialLink::Logit => BinomialLink::Logit,
463                SmoothBinomialLink::Probit => BinomialLink::Probit,
464                SmoothBinomialLink::Cloglog => BinomialLink::Cloglog,
465            };
466            let f = BinomialRegressor::builder()
467                .link(link)
468                .tolerance(GLM_TOL)
469                .max_iterations(GLM_MAX_ITER)
470                .build()
471                .fit(x, y)
472                .ok()?;
473            (
474                f.predict_with_se(grid, PredictionType::Link, None, 0.95),
475                Box::new(move |eta| link.link_inverse(eta)),
476            )
477        }
478        SmoothFamily::Gamma(link) => {
479            let (power, inv): (f64, Box<dyn Fn(f64) -> f64>) = match link {
480                SmoothGammaLink::Inverse => (-1.0, Box::new(|eta: f64| 1.0 / eta)),
481                SmoothGammaLink::Log => (0.0, Box::new(f64::exp)),
482            };
483            let f = GammaRegressor::builder()
484                .link_power(power)
485                .tolerance(GLM_TOL)
486                .max_iterations(GLM_MAX_ITER)
487                .build()
488                .fit(x, y)
489                .ok()?;
490            (
491                f.inner()
492                    .predict_with_se(grid, PredictionType::Link, None, 0.95),
493                inv,
494            )
495        }
496    };
497
498    let n = grid.nrows();
499    let eta = &link_pred.fit;
500    let fit = Col::from_fn(n, |k| linkinv(eta[k]));
501    if !se {
502        return Some(PredictionResult::with_intervals(
503            fit,
504            Col::zeros(n),
505            Col::zeros(n),
506            Col::zeros(n),
507        ));
508    }
509    let z = crate::stat::dist::qnorm(0.975);
510    let lower = Col::from_fn(n, |k| linkinv(eta[k] - z * link_pred.se[k]));
511    let upper = Col::from_fn(n, |k| linkinv(eta[k] + z * link_pred.se[k]));
512    Some(PredictionResult::with_intervals(
513        fit,
514        lower,
515        upper,
516        link_pred.se.clone(),
517    ))
518}