use rand::{rngs::StdRng, Rng, SeedableRng};
use rand_distr::{Distribution, Normal};
use crate::EventStream;
pub fn slice_rng(seed: u64, index: usize) -> StdRng {
let mut z = seed.wrapping_add((index as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
StdRng::seed_from_u64(z ^ (z >> 31))
}
impl EventStream {
pub fn random_flip_x(&self, p: f64, seed: u64) -> EventStream {
if fires(p, seed) {
self.flip_x()
} else {
self.clone()
}
}
pub fn random_flip_y(&self, p: f64, seed: u64) -> EventStream {
if fires(p, seed) {
self.flip_y()
} else {
self.clone()
}
}
pub fn random_polarity_flip(&self, p: f64, seed: u64) -> EventStream {
if fires(p, seed) {
self.invert_polarity()
} else {
self.clone()
}
}
pub fn random_crop(&self, width: usize, height: usize, seed: u64) -> EventStream {
let (sensor_w, sensor_h) = self.sensor_size();
if width >= sensor_w && height >= sensor_h {
return self.clone();
}
let mut rng = slice_rng(seed, 0);
let x0 = rng.gen_range(0..=sensor_w.saturating_sub(width)) as i64;
let y0 = rng.gen_range(0..=sensor_h.saturating_sub(height)) as i64;
self.crop(x0, y0, width, height)
}
pub fn event_drop(&self, p: f64, seed: u64) -> EventStream {
if p <= 0.0 {
return self.clone();
}
let mut rng = slice_rng(seed, 0);
let (width, height) = self.sensor_size();
self.remap(width, height, move |x, y, t, polarity| {
(rng.gen::<f64>() >= p).then_some((x, y, t, polarity))
})
}
pub fn pixel_dropout(&self, p: f64, seed: u64) -> EventStream {
if p <= 0.0 {
return self.clone();
}
let (width, height) = self.sensor_size();
let mut rng = slice_rng(seed, 0);
let drop: Vec<bool> = (0..width * height).map(|_| rng.gen::<f64>() < p).collect();
self.drop_masked_pixels(&drop)
}
pub fn spatial_jitter(&self, sigma: f64, seed: u64) -> EventStream {
if sigma <= 0.0 {
return self.clone();
}
let normal = match Normal::new(0.0, sigma) {
Ok(normal) => normal,
Err(_) => return self.clone(),
};
let mut rng = slice_rng(seed, 0);
let (width, height) = self.sensor_size();
self.remap(width, height, move |x, y, t, polarity| {
let dx = normal.sample(&mut rng).round() as i64;
let dy = normal.sample(&mut rng).round() as i64;
Some((x + dx, y + dy, t, polarity))
})
}
pub fn time_jitter(&self, sigma: f64, seed: u64) -> EventStream {
if sigma <= 0.0 {
return self.clone();
}
let normal = match Normal::new(0.0, sigma) {
Ok(normal) => normal,
Err(_) => return self.clone(),
};
let mut rng = slice_rng(seed, 0);
let (width, height) = self.sensor_size();
let jittered = self.remap(width, height, move |x, y, t, polarity| {
Some((x, y, t + normal.sample(&mut rng).round() as i64, polarity))
});
jittered.sort_by_time()
}
pub fn time_reversal(&self, p: f64, seed: u64) -> EventStream {
if !fires(p, seed) || self.is_empty() {
return self.clone();
}
let ts = self.ts();
let (&t_min, &t_max) = match (ts.iter().min(), ts.iter().max()) {
(Some(min), Some(max)) => (min, max),
_ => return self.clone(),
};
let sum = t_min + t_max;
let (width, height) = self.sensor_size();
self.remap(width, height, |x, y, t, polarity| {
Some((x, y, sum - t, !polarity))
})
.sort_by_time()
}
}
fn fires(p: f64, seed: u64) -> bool {
if p <= 0.0 {
return false;
}
if p >= 1.0 {
return true;
}
slice_rng(seed, 0).gen::<f64>() < p
}
#[cfg(test)]
mod tests {
use crate::{EventStream, EventStreamBuilder};
fn sample() -> EventStream {
let mut builder = EventStreamBuilder::new(8, 6, 0.001);
for i in 0..32u16 {
builder.push(i % 8, i % 6, 100 + i64::from(i) * 10, i % 2 == 0);
}
builder.build()
}
fn coords(stream: &EventStream) -> Vec<(u16, u16)> {
stream
.xs()
.iter()
.copied()
.zip(stream.ys().iter().copied())
.collect()
}
#[test]
fn probability_bounds_are_exact() {
let s = sample();
assert_eq!(coords(&s.random_flip_x(0.0, 7)), coords(&s));
assert_eq!(coords(&s.random_flip_x(1.0, 7)), coords(&s.flip_x()));
assert_eq!(s.event_drop(0.0, 7).len(), s.len());
assert_eq!(s.event_drop(1.0, 7).len(), 0);
}
#[test]
fn same_seed_gives_identical_output() {
let s = sample();
assert_eq!(s.event_drop(0.5, 42).ts(), s.event_drop(0.5, 42).ts());
assert_eq!(
coords(&s.spatial_jitter(1.5, 42)),
coords(&s.spatial_jitter(1.5, 42))
);
}
#[test]
fn different_seeds_give_different_output() {
let s = sample();
assert_ne!(s.event_drop(0.5, 1).len(), s.event_drop(0.5, 2).len());
}
#[test]
fn slice_rng_decorrelates_adjacent_indices() {
use rand::Rng;
let draws: Vec<f64> = (0..8)
.map(|index| super::slice_rng(0, index).gen::<f64>())
.collect();
for window in draws.windows(2) {
assert!((window[0] - window[1]).abs() > 1e-6);
}
}
#[test]
fn event_drop_thins_without_moving_events() {
let s = sample();
let dropped = s.event_drop(0.5, 3);
assert!(dropped.len() < s.len() && !dropped.is_empty());
let original: Vec<_> = s
.ts()
.iter()
.zip(coords(&s))
.map(|(t, xy)| (*t, xy))
.collect();
for (t, xy) in dropped.ts().iter().zip(coords(&dropped)) {
assert!(original.contains(&(*t, xy)));
}
}
#[test]
fn pixel_dropout_removes_whole_pixels() {
let s = sample();
let dropped = s.pixel_dropout(0.5, 5);
let survivors: std::collections::HashSet<_> = coords(&dropped).into_iter().collect();
let removed: std::collections::HashSet<_> = coords(&s)
.into_iter()
.filter(|xy| !survivors.contains(xy))
.collect();
assert!(removed.is_disjoint(&survivors));
assert!(!removed.is_empty());
}
#[test]
fn pixel_dropout_p_is_the_fraction_removed() {
let mut builder = EventStreamBuilder::new(40, 40, 0.001);
for x in 0..40u16 {
for y in 0..40u16 {
builder.push(x, y, i64::from(x) * 40 + i64::from(y), true);
}
}
let uniform = builder.build();
let kept = uniform.pixel_dropout(0.1, 5).len() as f64 / uniform.len() as f64;
assert!(kept > 0.8, "p=0.1 should keep ~90% of pixels, kept {kept}");
}
#[test]
fn time_reversal_mirrors_span_and_inverts_polarity() {
let s = sample();
let reversed = s.time_reversal(1.0, 0);
assert_eq!(reversed.len(), s.len());
assert_eq!(reversed.ts().first(), s.ts().first());
assert_eq!(reversed.ts().last(), s.ts().last());
assert!(reversed.ts().windows(2).all(|w| w[0] <= w[1]));
assert_eq!(reversed.ps()[0], !s.ps()[s.len() - 1]);
}
#[test]
fn time_jitter_leaves_the_stream_sorted() {
let jittered = sample().time_jitter(500.0, 11);
assert!(jittered.ts().windows(2).all(|w| w[0] <= w[1]));
}
#[test]
fn random_crop_larger_than_sensor_is_identity() {
let s = sample();
assert_eq!(coords(&s.random_crop(64, 64, 9)), coords(&s));
}
#[test]
fn random_crop_bounds_the_result() {
let cropped = sample().random_crop(3, 2, 9);
assert_eq!(cropped.sensor_size(), (3, 2));
assert!(cropped.xs().iter().all(|&x| (x as usize) < 3));
assert!(cropped.ys().iter().all(|&y| (y as usize) < 2));
}
#[test]
fn augmentations_handle_the_empty_stream() {
let empty = EventStreamBuilder::new(8, 6, 0.001).build();
assert!(empty.random_flip_x(1.0, 0).is_empty());
assert!(empty.event_drop(0.5, 0).is_empty());
assert!(empty.pixel_dropout(0.5, 0).is_empty());
assert!(empty.spatial_jitter(2.0, 0).is_empty());
assert!(empty.time_jitter(2.0, 0).is_empty());
assert!(empty.time_reversal(1.0, 0).is_empty());
assert!(empty.random_crop(3, 2, 0).is_empty());
}
}