Skip to main content

ggplot_rs/stat/
ydensity.rs

1use crate::aes::Aesthetic;
2use crate::data::{DataFrame, Value};
3use crate::scale::ScaleSet;
4
5use super::Stat;
6
7/// Kernel density estimation on Y per group (for violin plots).
8/// Outputs: x (group value), y (eval points), violinwidth (density normalized to
9/// [0, 1]). The geom mirrors `violinwidth` around the group's x slot.
10pub struct StatYDensity {
11    pub n_points: usize,
12}
13
14impl Default for StatYDensity {
15    fn default() -> Self {
16        StatYDensity { n_points: 512 }
17    }
18}
19
20impl Stat for StatYDensity {
21    fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
22        let x_col = data.column("x");
23        let y_col = match data.column("y") {
24            Some(c) => c,
25            None => return DataFrame::new(),
26        };
27
28        let values: Vec<f64> = y_col.iter().filter_map(|v| v.as_f64()).collect();
29        if values.len() < 2 {
30            return DataFrame::new();
31        }
32
33        // Keep the group's x *value* as-is (e.g. the discrete label "A"). The geom
34        // maps it through the X scale, exactly like boxplot — converting to f64 here
35        // would collapse every discrete group to 0.0.
36        let group_x = x_col
37            .and_then(|c| c.first())
38            .cloned()
39            .unwrap_or(Value::Float(0.0));
40
41        let n = values.len() as f64;
42        // R's bw.nrd0 (positive even for zero-spread data).
43        let bandwidth = super::bw_nrd0(&values);
44
45        // ggplot2's geom_violin defaults to trim = TRUE: evaluate the density
46        // over the observed data range, not extended by ±3 bandwidths.
47        let y_min = values.iter().cloned().fold(f64::INFINITY, f64::min);
48        let y_max = values.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
49        let step = (y_max - y_min) / (self.n_points - 1) as f64;
50
51        let mut x_vals = Vec::with_capacity(self.n_points);
52        let mut y_vals = Vec::with_capacity(self.n_points);
53
54        // Compute density at each evaluation point
55        let mut densities = Vec::with_capacity(self.n_points);
56        let mut max_density: f64 = 0.0;
57        for i in 0..self.n_points {
58            let y = y_min + i as f64 * step;
59            let density: f64 = values
60                .iter()
61                .map(|yi| gaussian_kernel((y - yi) / bandwidth))
62                .sum::<f64>()
63                / (n * bandwidth);
64            densities.push((y, density));
65            if density > max_density {
66                max_density = density;
67            }
68        }
69
70        // Normalize density to [0, 1] (peak = 1). The geom scales this by the
71        // per-group slot half-width, so the widest point fills the group's slot.
72        let scale = if max_density > 0.0 {
73            1.0 / max_density
74        } else {
75            1.0
76        };
77
78        let mut width_vals = Vec::with_capacity(self.n_points);
79        for (y, density) in &densities {
80            x_vals.push(group_x.clone());
81            y_vals.push(Value::Float(*y));
82            width_vals.push(Value::Float(density * scale));
83        }
84
85        let mut result = DataFrame::new();
86        result.add_column("x".to_string(), x_vals);
87        result.add_column("y".to_string(), y_vals);
88        result.add_column("violinwidth".to_string(), width_vals);
89
90        // Carry over grouping columns
91        for col_name in &["color", "fill", "group"] {
92            if let Some(col) = data.column(col_name) {
93                if let Some(first) = col.first() {
94                    result.add_column(col_name.to_string(), vec![first.clone(); self.n_points]);
95                }
96            }
97        }
98
99        result
100    }
101
102    fn required_aes(&self) -> Vec<Aesthetic> {
103        vec![Aesthetic::X, Aesthetic::Y]
104    }
105
106    fn name(&self) -> &str {
107        "ydensity"
108    }
109}
110
111fn gaussian_kernel(x: f64) -> f64 {
112    (-(x * x) / 2.0).exp() / (2.0 * std::f64::consts::PI).sqrt()
113}
114
115#[cfg(test)]
116mod tests {
117    use super::*;
118
119    #[test]
120    fn test_ydensity_basic() {
121        let mut data = DataFrame::new();
122        data.add_column("x".to_string(), vec![Value::Float(1.0); 50]);
123        let y_vals: Vec<Value> = (0..50).map(|i| Value::Float(i as f64)).collect();
124        data.add_column("y".to_string(), y_vals);
125
126        let stat = StatYDensity::default();
127        let scales = ScaleSet::new();
128        let result = stat.compute_group(&data, &scales);
129
130        assert!(result.nrows() > 0);
131        assert!(result.column("x").is_some());
132        assert!(result.column("y").is_some());
133        assert!(result.column("violinwidth").is_some());
134        // Normalized width peaks at 1.0.
135        let max_w = result
136            .column("violinwidth")
137            .unwrap()
138            .iter()
139            .filter_map(|v| v.as_f64())
140            .fold(0.0_f64, f64::max);
141        assert!(
142            (max_w - 1.0).abs() < 1e-9,
143            "peak width should be 1.0, got {max_w}"
144        );
145    }
146}