ggplot_rs/position/
jitter.rs1use crate::data::{DataFrame, Value};
2use crate::rng::SplitMix64;
3
4use super::{Position, PositionParams};
5
6pub const JITTER_SEED: u64 = 0x6A17_7E2D;
10
11#[derive(Clone, Debug)]
20pub struct PositionJitter {
21 pub width: Option<f64>,
23 pub height: Option<f64>,
25 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 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 pub fn with_width(mut self, width: f64) -> Self {
51 self.width = Some(width);
52 self
53 }
54
55 pub fn with_height(mut self, height: f64) -> Self {
57 self.height = Some(height);
58 self
59 }
60
61 pub fn with_seed(mut self, seed: u64) -> Self {
63 self.seed = seed;
64 self
65 }
66}
67
68pub 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 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 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 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}