Skip to main content

ggplot_rs/position/
jitter.rs

1use crate::data::{DataFrame, Value};
2use crate::rng::SplitMix64;
3
4use super::{Position, PositionParams};
5
6/// Default seed: the same data must render the same SVG (stable snapshots,
7/// cacheable output, no flicker when a host re-renders), as with ggplot2's
8/// `position_jitter(seed = …)`.
9pub const JITTER_SEED: u64 = 0x6A17_7E2D;
10
11/// Add random noise to x and y positions to reduce overplotting.
12///
13/// Each value moves by a uniform amount in `[-width, width)` (x) and
14/// `[-height, height)` (y). `None` (the default) means "40% of the data's
15/// resolution", as in ggplot2: the smallest gap between distinct values, or 1
16/// for integer/constant data — so integer-coded categories (1, 2, 3, …) get
17/// ±0.4 and finely-spaced continuous data gets proportionally small noise.
18/// `Some(0.0)` disables jitter along that axis.
19#[derive(Clone, Debug)]
20pub struct PositionJitter {
21    /// Horizontal jitter amount in data units; `None` = 0.4 × resolution(x).
22    pub width: Option<f64>,
23    /// Vertical jitter amount in data units; `None` = 0.4 × resolution(y).
24    pub height: Option<f64>,
25    /// Seed of the deterministic noise stream.
26    pub seed: u64,
27}
28
29impl Default for PositionJitter {
30    fn default() -> Self {
31        PositionJitter {
32            width: None,
33            height: None,
34            seed: JITTER_SEED,
35        }
36    }
37}
38
39impl PositionJitter {
40    /// Explicit jitter amounts (data units), default seed.
41    pub fn new(width: f64, height: f64) -> Self {
42        PositionJitter {
43            width: Some(width),
44            height: Some(height),
45            seed: JITTER_SEED,
46        }
47    }
48
49    /// Set the horizontal jitter amount (data units).
50    pub fn with_width(mut self, width: f64) -> Self {
51        self.width = Some(width);
52        self
53    }
54
55    /// Set the vertical jitter amount (data units).
56    pub fn with_height(mut self, height: f64) -> Self {
57        self.height = Some(height);
58        self
59    }
60
61    /// Use a different seed (ggplot2's `position_jitter(seed = …)`).
62    pub fn with_seed(mut self, seed: u64) -> Self {
63        self.seed = seed;
64        self
65    }
66}
67
68/// ggplot2's `resolution(x, zero = FALSE)`: the smallest positive gap between
69/// distinct finite values; 1 for integer-typed data, a single distinct value, or
70/// no numeric values at all.
71pub fn resolution(values: &[Value]) -> f64 {
72    let mut all_int = true;
73    let mut xs: Vec<f64> = Vec::with_capacity(values.len());
74    for v in values {
75        match v {
76            Value::Integer(i) => xs.push(*i as f64),
77            Value::Float(f) if f.is_finite() => {
78                all_int = false;
79                xs.push(*f);
80            }
81            Value::DateTime(s) => {
82                all_int = false;
83                xs.push(*s as f64);
84            }
85            _ => {}
86        }
87    }
88    if all_int || xs.len() < 2 {
89        return 1.0;
90    }
91    xs.sort_by(|a, b| a.total_cmp(b));
92    let span = xs[xs.len() - 1] - xs[0];
93    // Ignore floating-point dust between "equal" values (ggplot2 uses a
94    // tolerance of sqrt(.Machine$double.eps) relative to the range).
95    let tol = span.abs() * 1e-8;
96    let min_gap = xs
97        .windows(2)
98        .map(|w| w[1] - w[0])
99        .filter(|d| *d > tol)
100        .fold(f64::INFINITY, f64::min);
101    if min_gap.is_finite() && min_gap > 0.0 {
102        min_gap
103    } else {
104        1.0
105    }
106}
107
108fn jitter_column(col: &mut [Value], amount: f64, rng: &mut SplitMix64) {
109    for v in col.iter_mut() {
110        if let Some(x) = v.as_f64() {
111            *v = Value::Float(x + rng.range_f64(-amount, amount));
112        }
113    }
114}
115
116impl Position for PositionJitter {
117    fn compute(&self, data: &mut DataFrame, _params: &PositionParams) {
118        let mut rng = SplitMix64::new(self.seed);
119
120        if let Some(x_col) = data.column_mut("x") {
121            let w = self.width.unwrap_or_else(|| 0.4 * resolution(x_col));
122            if w > 0.0 && w.is_finite() {
123                jitter_column(x_col, w, &mut rng);
124            }
125        }
126
127        if let Some(y_col) = data.column_mut("y") {
128            let h = self.height.unwrap_or_else(|| 0.4 * resolution(y_col));
129            if h > 0.0 && h.is_finite() {
130                jitter_column(y_col, h, &mut rng);
131            }
132        }
133    }
134
135    fn name(&self) -> &str {
136        "jitter"
137    }
138}
139
140#[cfg(test)]
141mod tests {
142    use super::*;
143
144    fn frame(x: Vec<Value>, y: Vec<Value>) -> DataFrame {
145        let mut df = DataFrame::new();
146        df.add_column("x".into(), x);
147        df.add_column("y".into(), y);
148        df
149    }
150
151    fn floats(v: &[f64]) -> Vec<Value> {
152        v.iter().map(|f| Value::Float(*f)).collect()
153    }
154
155    #[test]
156    fn resolution_matches_ggplot2() {
157        assert_eq!(resolution(&floats(&[1.0, 2.0, 3.0])), 1.0);
158        assert_eq!(resolution(&floats(&[0.1, 0.3, 0.2])), 0.09999999999999998);
159        assert_eq!(resolution(&floats(&[5.0])), 1.0);
160        assert_eq!(resolution(&floats(&[2.0, 2.0])), 1.0);
161        assert_eq!(
162            resolution(&[Value::Integer(10), Value::Integer(20)]),
163            1.0,
164            "integer data has resolution 1"
165        );
166        assert_eq!(resolution(&[Value::Str("a".into())]), 1.0);
167    }
168
169    #[test]
170    fn default_width_scales_with_resolution() {
171        // Values spaced 0.01 apart: jitter must stay within ±0.004.
172        let xs: Vec<f64> = (0..50).map(|i| i as f64 * 0.01).collect();
173        let mut df = frame(floats(&xs), floats(&xs));
174        PositionJitter::default().compute(&mut df, &PositionParams::default());
175        for (orig, v) in xs.iter().zip(df.column("x").unwrap()) {
176            let d = (v.as_f64().unwrap() - orig).abs();
177            assert!(d <= 0.004 + 1e-12, "moved {d}");
178        }
179        // Integer-coded positions keep the classic ±0.4.
180        let mut df = frame(floats(&[1.0, 2.0, 3.0]), floats(&[1.0, 2.0, 3.0]));
181        PositionJitter::default().compute(&mut df, &PositionParams::default());
182        let moved: Vec<f64> = df
183            .column("x")
184            .unwrap()
185            .iter()
186            .zip([1.0, 2.0, 3.0])
187            .map(|(v, o)| (v.as_f64().unwrap() - o).abs())
188            .collect();
189        assert!(moved.iter().all(|d| *d <= 0.4));
190        assert!(moved.iter().any(|d| *d > 0.01));
191    }
192
193    #[test]
194    fn seed_controls_output() {
195        let run = |p: PositionJitter| {
196            let mut df = frame(floats(&[1.0, 2.0, 3.0]), floats(&[1.0, 2.0, 3.0]));
197            p.compute(&mut df, &PositionParams::default());
198            df.column("x").unwrap().to_vec()
199        };
200        assert_eq!(
201            run(PositionJitter::default()),
202            run(PositionJitter::default())
203        );
204        assert_ne!(
205            run(PositionJitter::default()),
206            run(PositionJitter::default().with_seed(1))
207        );
208        assert_eq!(
209            run(PositionJitter::default().with_seed(1)),
210            run(PositionJitter::default().with_seed(1))
211        );
212    }
213
214    #[test]
215    fn zero_amount_disables_axis() {
216        let mut df = frame(floats(&[1.0, 2.0]), floats(&[5.0, 6.0]));
217        PositionJitter::new(0.3, 0.0).compute(&mut df, &PositionParams::default());
218        assert_eq!(df.column("y").unwrap(), &floats(&[5.0, 6.0])[..]);
219        assert_ne!(df.column("x").unwrap(), &floats(&[1.0, 2.0])[..]);
220    }
221}