Skip to main content

ggplot_rs/geom/
smooth.rs

1use crate::aes::Aesthetic;
2use crate::coord::Coord;
3use crate::data::{DataFrame, Value};
4use crate::position::identity::PositionIdentity;
5use crate::position::Position;
6use crate::render::backend::{DrawBackend, LineStyle, Linetype, PointShape, PointStyle, RectStyle};
7use crate::render::RenderError;
8use crate::scale::ScaleSet;
9use crate::stat::smooth::{SmoothMethod, StatSmooth};
10use crate::stat::Stat;
11use crate::theme::Theme;
12
13use super::{Geom, GeomParams};
14
15/// Smooth line with optional confidence ribbon.
16pub struct GeomSmooth {
17    pub color: (u8, u8, u8),
18    pub fill: (u8, u8, u8),
19    pub line_width: f64,
20    pub alpha: f64,
21    pub se: bool,
22    pub n_points: usize,
23    pub method: SmoothMethod,
24}
25
26impl Default for GeomSmooth {
27    fn default() -> Self {
28        GeomSmooth {
29            color: (51, 102, 204),
30            fill: (51, 102, 204),
31            line_width: 1.5,
32            alpha: 0.2,
33            se: true,
34            n_points: 80,
35            method: SmoothMethod::Lm,
36        }
37    }
38}
39
40impl GeomSmooth {
41    /// Use LOESS smoothing with the given span.
42    pub fn loess(mut self, span: f64) -> Self {
43        self.method = SmoothMethod::Loess { span };
44        self
45    }
46
47    /// Use penalized B-spline (P-spline) GAM smoothing — ggplot2's
48    /// `method = "gam"`, backed by anofox-regression with GCV-selected λ.
49    #[cfg(feature = "regression")]
50    pub fn gam(mut self) -> Self {
51        self.method = SmoothMethod::Gam;
52        self
53    }
54
55    /// Use a generalized linear model — ggplot2's
56    /// `geom_smooth(method = "glm", method.args = list(family = …))` — e.g.
57    /// `GeomSmooth::default().glm(SmoothFamily::binomial())`. The band is the
58    /// link-scale confidence interval mapped through the inverse link (see
59    /// [`SmoothFamily`](crate::stat::smooth::SmoothFamily)).
60    #[cfg(feature = "regression")]
61    pub fn glm(mut self, family: crate::stat::smooth::SmoothFamily) -> Self {
62        self.method = SmoothMethod::Glm { family };
63        self
64    }
65}
66
67impl Geom for GeomSmooth {
68    fn draw(
69        &self,
70        data: &DataFrame,
71        coord: &dyn Coord,
72        scales: &ScaleSet,
73        _theme: &Theme,
74        backend: &mut dyn DrawBackend,
75    ) -> Result<(), RenderError> {
76        let x_col = data
77            .column("x")
78            .ok_or(RenderError::MissingAesthetic("x".into()))?;
79        let y_col = data
80            .column("y")
81            .ok_or(RenderError::MissingAesthetic("y".into()))?;
82        let ymin_col = data.column("ymin");
83        let ymax_col = data.column("ymax");
84        let color_col = data.column("color");
85        let fill_col = data.column("fill");
86
87        let plot_area = backend.plot_area();
88        let x_scale = scales.get(&Aesthetic::X);
89        let y_scale = scales.get(&Aesthetic::Y);
90
91        // If there's a color/fill aesthetic, draw separate smooths per group
92        if let Some(cc) = color_col.or(fill_col) {
93            let mut groups: Vec<(String, Vec<usize>)> = Vec::new();
94            for (i, v) in cc.iter().enumerate() {
95                let key = v.to_group_key();
96                if let Some(entry) = groups.iter_mut().find(|(k, _)| k == &key) {
97                    entry.1.push(i);
98                } else {
99                    groups.push((key, vec![i]));
100                }
101            }
102
103            for (_, indices) in &groups {
104                let first_idx = indices[0];
105
106                // Determine colors from mapped aesthetics
107                let line_color = color_col
108                    .and_then(|c| scales.map_color(&Aesthetic::Color, &c[first_idx]))
109                    .unwrap_or(self.color);
110                let ribbon_fill = fill_col
111                    .and_then(|f| scales.map_color(&Aesthetic::Fill, &f[first_idx]))
112                    .or_else(|| {
113                        color_col.and_then(|c| scales.map_color(&Aesthetic::Color, &c[first_idx]))
114                    })
115                    .unwrap_or(self.fill);
116
117                // Draw confidence ribbon
118                if self.se {
119                    if let (Some(ymin), Some(ymax)) = (ymin_col, ymax_col) {
120                        let mut upper_points: Vec<(f64, f64)> = Vec::new();
121                        let mut lower_points: Vec<(f64, f64)> = Vec::new();
122
123                        for &i in indices {
124                            let nx = x_scale.map(|s| s.map(&x_col[i])).unwrap_or(0.0);
125                            let ny_max = y_scale.map(|s| s.map(&ymax[i])).unwrap_or(0.0);
126                            let ny_min = y_scale.map(|s| s.map(&ymin[i])).unwrap_or(0.0);
127
128                            upper_points.push(coord.transform((nx, ny_max), &plot_area));
129                            lower_points.push(coord.transform((nx, ny_min), &plot_area));
130                        }
131
132                        let mut polygon = upper_points;
133                        lower_points.reverse();
134                        polygon.extend(lower_points);
135
136                        if polygon.len() >= 3 {
137                            backend.draw_polygon(
138                                &polygon,
139                                &RectStyle {
140                                    fill: Some(ribbon_fill),
141                                    stroke: None,
142                                    stroke_width: 0.0,
143                                    alpha: self.alpha,
144                                    clip: true,
145                                },
146                            )?;
147                        }
148                    }
149                }
150
151                // Draw fitted line
152                let points: Vec<(f64, f64)> = indices
153                    .iter()
154                    .map(|&i| {
155                        let nx = x_scale.map(|s| s.map(&x_col[i])).unwrap_or(0.0);
156                        let ny = y_scale.map(|s| s.map(&y_col[i])).unwrap_or(0.0);
157                        coord.transform((nx, ny), &plot_area)
158                    })
159                    .collect();
160
161                if points.len() >= 2 {
162                    backend.draw_line(
163                        &points,
164                        &LineStyle {
165                            color: line_color,
166                            alpha: 1.0,
167                            width: self.line_width,
168                            linetype: Linetype::Solid,
169                        },
170                    )?;
171                    draw_hover_marks(
172                        backend, &points, indices, x_col, y_col, ymin_col, ymax_col, line_color,
173                    )?;
174                }
175            }
176        } else {
177            // No grouping — original behavior with fixed colors
178
179            // Draw confidence ribbon first (behind line)
180            if self.se {
181                if let (Some(ymin), Some(ymax)) = (ymin_col, ymax_col) {
182                    let mut upper_points: Vec<(f64, f64)> = Vec::new();
183                    let mut lower_points: Vec<(f64, f64)> = Vec::new();
184
185                    for i in 0..data.nrows() {
186                        let nx = x_scale.map(|s| s.map(&x_col[i])).unwrap_or(0.0);
187                        let ny_max = y_scale.map(|s| s.map(&ymax[i])).unwrap_or(0.0);
188                        let ny_min = y_scale.map(|s| s.map(&ymin[i])).unwrap_or(0.0);
189
190                        upper_points.push(coord.transform((nx, ny_max), &plot_area));
191                        lower_points.push(coord.transform((nx, ny_min), &plot_area));
192                    }
193
194                    // Build polygon: upper left-to-right, then lower right-to-left
195                    let mut polygon = upper_points;
196                    lower_points.reverse();
197                    polygon.extend(lower_points);
198
199                    if polygon.len() >= 3 {
200                        backend.draw_polygon(
201                            &polygon,
202                            &RectStyle {
203                                fill: Some(self.fill),
204                                stroke: None,
205                                stroke_width: 0.0,
206                                alpha: self.alpha,
207                                clip: true,
208                            },
209                        )?;
210                    }
211                }
212            }
213
214            // Draw fitted line
215            let points: Vec<(f64, f64)> = (0..data.nrows())
216                .map(|i| {
217                    let nx = x_scale.map(|s| s.map(&x_col[i])).unwrap_or(0.0);
218                    let ny = y_scale.map(|s| s.map(&y_col[i])).unwrap_or(0.0);
219                    coord.transform((nx, ny), &plot_area)
220                })
221                .collect();
222
223            if points.len() >= 2 {
224                backend.draw_line(
225                    &points,
226                    &LineStyle {
227                        color: self.color,
228                        alpha: 1.0,
229                        width: self.line_width,
230                        linetype: Linetype::Solid,
231                    },
232                )?;
233                let rows: Vec<usize> = (0..data.nrows()).collect();
234                draw_hover_marks(
235                    backend, &points, &rows, x_col, y_col, ymin_col, ymax_col, self.color,
236                )?;
237            }
238        }
239
240        Ok(())
241    }
242
243    fn required_aes(&self) -> Vec<Aesthetic> {
244        vec![Aesthetic::X, Aesthetic::Y]
245    }
246
247    fn default_stat(&self) -> Box<dyn Stat> {
248        Box::new(StatSmooth {
249            n_points: self.n_points,
250            se: self.se,
251            method: self.method.clone(),
252        })
253    }
254
255    fn default_position(&self) -> Box<dyn Position> {
256        Box::new(PositionIdentity)
257    }
258
259    fn default_params(&self) -> GeomParams {
260        GeomParams::default()
261    }
262
263    fn name(&self) -> &str {
264        "smooth"
265    }
266
267    fn set_series_color(&mut self, color: (u8, u8, u8)) {
268        self.color = color;
269        self.fill = color;
270    }
271}
272
273/// The fitted curve is drawn as a single path and so carries no per-point
274/// marks. Emit a sparse set of transparent hover points along it — each tagged
275/// with the fitted `ŷ` (and the CI when present) — so the axis pointer can read
276/// the smoother on hover, mirroring `geom_density`.
277#[allow(clippy::too_many_arguments)]
278fn draw_hover_marks(
279    backend: &mut dyn DrawBackend,
280    points: &[(f64, f64)],
281    rows: &[usize],
282    x_col: &[Value],
283    y_col: &[Value],
284    ymin_col: Option<&[Value]>,
285    ymax_col: Option<&[Value]>,
286    color: (u8, u8, u8),
287) -> Result<(), RenderError> {
288    let step = (rows.len() / 40).max(1);
289    for (k, &i) in rows.iter().enumerate() {
290        if k % step != 0 {
291            continue;
292        }
293        super::set_mark(
294            backend,
295            Some(smooth_tip(y_col, ymin_col, ymax_col, i)),
296            Some(super::tip_value(&x_col[i])),
297            None,
298            super::raw_value(&y_col[i]),
299        );
300        backend.draw_shape(
301            points[k],
302            0.6,
303            &PointStyle {
304                color,
305                alpha: 0.0,
306                filled: true,
307                shape: PointShape::Circle,
308            },
309        )?;
310    }
311    super::clear_mark(backend);
312    Ok(())
313}
314
315/// `ŷ = <fitted>` plus ` [lo, hi]` when a confidence band is present.
316fn smooth_tip(
317    y_col: &[Value],
318    ymin_col: Option<&[Value]>,
319    ymax_col: Option<&[Value]>,
320    i: usize,
321) -> String {
322    let yv = y_col[i]
323        .as_f64()
324        .map(|f| format!("{f:.3}"))
325        .unwrap_or_default();
326    let ci = match (ymin_col, ymax_col) {
327        (Some(lo), Some(hi)) => match (lo[i].as_f64(), hi[i].as_f64()) {
328            (Some(a), Some(b)) => format!(" [{a:.3}, {b:.3}]"),
329            _ => String::new(),
330        },
331        _ => String::new(),
332    };
333    format!("ŷ = {yv}{ci}")
334}