1use rand::{rngs::StdRng, Rng, SeedableRng};
17use rand_distr::{Distribution, Normal};
18
19use crate::EventStream;
20
21pub fn slice_rng(seed: u64, index: usize) -> StdRng {
29 let mut z = seed.wrapping_add((index as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
30 z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
31 z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
32 StdRng::seed_from_u64(z ^ (z >> 31))
33}
34
35impl EventStream {
36 pub fn random_flip_x(&self, p: f64, seed: u64) -> EventStream {
38 if fires(p, seed) {
39 self.flip_x()
40 } else {
41 self.clone()
42 }
43 }
44
45 pub fn random_flip_y(&self, p: f64, seed: u64) -> EventStream {
47 if fires(p, seed) {
48 self.flip_y()
49 } else {
50 self.clone()
51 }
52 }
53
54 pub fn random_polarity_flip(&self, p: f64, seed: u64) -> EventStream {
60 if fires(p, seed) {
61 self.invert_polarity()
62 } else {
63 self.clone()
64 }
65 }
66
67 pub fn random_crop(&self, width: usize, height: usize, seed: u64) -> EventStream {
70 let (sensor_w, sensor_h) = self.sensor_size();
71 if width >= sensor_w && height >= sensor_h {
72 return self.clone();
73 }
74 let mut rng = slice_rng(seed, 0);
75 let x0 = rng.gen_range(0..=sensor_w.saturating_sub(width)) as i64;
76 let y0 = rng.gen_range(0..=sensor_h.saturating_sub(height)) as i64;
77 self.crop(x0, y0, width, height)
78 }
79
80 pub fn event_drop(&self, p: f64, seed: u64) -> EventStream {
83 if p <= 0.0 {
84 return self.clone();
85 }
86 let mut rng = slice_rng(seed, 0);
87 let (width, height) = self.sensor_size();
88 self.remap(width, height, move |x, y, t, polarity| {
89 (rng.gen::<f64>() >= p).then_some((x, y, t, polarity))
90 })
91 }
92
93 pub fn pixel_dropout(&self, p: f64, seed: u64) -> EventStream {
99 if p <= 0.0 {
100 return self.clone();
101 }
102 let (width, height) = self.sensor_size();
103 let mut rng = slice_rng(seed, 0);
104 let drop: Vec<bool> = (0..width * height).map(|_| rng.gen::<f64>() < p).collect();
107 self.drop_masked_pixels(&drop)
108 }
109
110 pub fn spatial_jitter(&self, sigma: f64, seed: u64) -> EventStream {
113 if sigma <= 0.0 {
114 return self.clone();
115 }
116 let normal = match Normal::new(0.0, sigma) {
117 Ok(normal) => normal,
118 Err(_) => return self.clone(),
119 };
120 let mut rng = slice_rng(seed, 0);
121 let (width, height) = self.sensor_size();
122 self.remap(width, height, move |x, y, t, polarity| {
123 let dx = normal.sample(&mut rng).round() as i64;
124 let dy = normal.sample(&mut rng).round() as i64;
125 Some((x + dx, y + dy, t, polarity))
126 })
127 }
128
129 pub fn time_jitter(&self, sigma: f64, seed: u64) -> EventStream {
135 if sigma <= 0.0 {
136 return self.clone();
137 }
138 let normal = match Normal::new(0.0, sigma) {
139 Ok(normal) => normal,
140 Err(_) => return self.clone(),
141 };
142 let mut rng = slice_rng(seed, 0);
143 let (width, height) = self.sensor_size();
144 let jittered = self.remap(width, height, move |x, y, t, polarity| {
145 Some((x, y, t + normal.sample(&mut rng).round() as i64, polarity))
146 });
147 jittered.sort_by_time()
148 }
149
150 pub fn time_reversal(&self, p: f64, seed: u64) -> EventStream {
156 if !fires(p, seed) || self.is_empty() {
157 return self.clone();
158 }
159 let ts = self.ts();
160 let (&t_min, &t_max) = match (ts.iter().min(), ts.iter().max()) {
161 (Some(min), Some(max)) => (min, max),
162 _ => return self.clone(),
163 };
164 let sum = t_min + t_max;
165 let (width, height) = self.sensor_size();
166 self.remap(width, height, |x, y, t, polarity| {
167 Some((x, y, sum - t, !polarity))
168 })
169 .sort_by_time()
170 }
171}
172
173fn fires(p: f64, seed: u64) -> bool {
176 if p <= 0.0 {
177 return false;
178 }
179 if p >= 1.0 {
180 return true;
181 }
182 slice_rng(seed, 0).gen::<f64>() < p
183}
184
185#[cfg(test)]
186mod tests {
187 use crate::{EventStream, EventStreamBuilder};
188
189 fn sample() -> EventStream {
190 let mut builder = EventStreamBuilder::new(8, 6, 0.001);
191 for i in 0..32u16 {
192 builder.push(i % 8, i % 6, 100 + i64::from(i) * 10, i % 2 == 0);
193 }
194 builder.build()
195 }
196
197 fn coords(stream: &EventStream) -> Vec<(u16, u16)> {
198 stream
199 .xs()
200 .iter()
201 .copied()
202 .zip(stream.ys().iter().copied())
203 .collect()
204 }
205
206 #[test]
207 fn probability_bounds_are_exact() {
208 let s = sample();
209 assert_eq!(coords(&s.random_flip_x(0.0, 7)), coords(&s));
210 assert_eq!(coords(&s.random_flip_x(1.0, 7)), coords(&s.flip_x()));
211 assert_eq!(s.event_drop(0.0, 7).len(), s.len());
212 assert_eq!(s.event_drop(1.0, 7).len(), 0);
213 }
214
215 #[test]
216 fn same_seed_gives_identical_output() {
217 let s = sample();
218 assert_eq!(s.event_drop(0.5, 42).ts(), s.event_drop(0.5, 42).ts());
219 assert_eq!(
220 coords(&s.spatial_jitter(1.5, 42)),
221 coords(&s.spatial_jitter(1.5, 42))
222 );
223 }
224
225 #[test]
226 fn different_seeds_give_different_output() {
227 let s = sample();
228 assert_ne!(s.event_drop(0.5, 1).len(), s.event_drop(0.5, 2).len());
230 }
231
232 #[test]
233 fn slice_rng_decorrelates_adjacent_indices() {
234 use rand::Rng;
235 let draws: Vec<f64> = (0..8)
238 .map(|index| super::slice_rng(0, index).gen::<f64>())
239 .collect();
240 for window in draws.windows(2) {
241 assert!((window[0] - window[1]).abs() > 1e-6);
242 }
243 }
244
245 #[test]
246 fn event_drop_thins_without_moving_events() {
247 let s = sample();
248 let dropped = s.event_drop(0.5, 3);
249 assert!(dropped.len() < s.len() && !dropped.is_empty());
250 let original: Vec<_> = s
252 .ts()
253 .iter()
254 .zip(coords(&s))
255 .map(|(t, xy)| (*t, xy))
256 .collect();
257 for (t, xy) in dropped.ts().iter().zip(coords(&dropped)) {
258 assert!(original.contains(&(*t, xy)));
259 }
260 }
261
262 #[test]
263 fn pixel_dropout_removes_whole_pixels() {
264 let s = sample();
265 let dropped = s.pixel_dropout(0.5, 5);
266 let survivors: std::collections::HashSet<_> = coords(&dropped).into_iter().collect();
267 let removed: std::collections::HashSet<_> = coords(&s)
268 .into_iter()
269 .filter(|xy| !survivors.contains(xy))
270 .collect();
271 assert!(removed.is_disjoint(&survivors));
273 assert!(!removed.is_empty());
274 }
275
276 #[test]
277 fn pixel_dropout_p_is_the_fraction_removed() {
278 let mut builder = EventStreamBuilder::new(40, 40, 0.001);
282 for x in 0..40u16 {
283 for y in 0..40u16 {
284 builder.push(x, y, i64::from(x) * 40 + i64::from(y), true);
285 }
286 }
287 let uniform = builder.build();
290 let kept = uniform.pixel_dropout(0.1, 5).len() as f64 / uniform.len() as f64;
291 assert!(kept > 0.8, "p=0.1 should keep ~90% of pixels, kept {kept}");
292 }
293
294 #[test]
295 fn time_reversal_mirrors_span_and_inverts_polarity() {
296 let s = sample();
297 let reversed = s.time_reversal(1.0, 0);
298 assert_eq!(reversed.len(), s.len());
299 assert_eq!(reversed.ts().first(), s.ts().first());
301 assert_eq!(reversed.ts().last(), s.ts().last());
302 assert!(reversed.ts().windows(2).all(|w| w[0] <= w[1]));
303 assert_eq!(reversed.ps()[0], !s.ps()[s.len() - 1]);
304 }
305
306 #[test]
307 fn time_jitter_leaves_the_stream_sorted() {
308 let jittered = sample().time_jitter(500.0, 11);
309 assert!(jittered.ts().windows(2).all(|w| w[0] <= w[1]));
310 }
311
312 #[test]
313 fn random_crop_larger_than_sensor_is_identity() {
314 let s = sample();
315 assert_eq!(coords(&s.random_crop(64, 64, 9)), coords(&s));
316 }
317
318 #[test]
319 fn random_crop_bounds_the_result() {
320 let cropped = sample().random_crop(3, 2, 9);
321 assert_eq!(cropped.sensor_size(), (3, 2));
322 assert!(cropped.xs().iter().all(|&x| (x as usize) < 3));
323 assert!(cropped.ys().iter().all(|&y| (y as usize) < 2));
324 }
325
326 #[test]
327 fn augmentations_handle_the_empty_stream() {
328 let empty = EventStreamBuilder::new(8, 6, 0.001).build();
329 assert!(empty.random_flip_x(1.0, 0).is_empty());
330 assert!(empty.event_drop(0.5, 0).is_empty());
331 assert!(empty.pixel_dropout(0.5, 0).is_empty());
332 assert!(empty.spatial_jitter(2.0, 0).is_empty());
333 assert!(empty.time_jitter(2.0, 0).is_empty());
334 assert!(empty.time_reversal(1.0, 0).is_empty());
335 assert!(empty.random_crop(3, 2, 0).is_empty());
336 }
337}