#[derive(Clone, Debug)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct PoissonEncoder {
pub num_steps: usize,
}
fn rate_dt_produces_spikes(rate_hz: f32, dt_seconds: f32) -> bool {
rate_hz.is_finite() && rate_hz > 0.0 && dt_seconds.is_finite() && dt_seconds > 0.0
}
pub fn probability_from_rate_hz(rate_hz: f32, dt_seconds: f32) -> f32 {
if !rate_dt_produces_spikes(rate_hz, dt_seconds) {
return 0.0;
}
let x = rate_hz * dt_seconds;
(-(-x).exp_m1()).clamp(0.0, 1.0)
}
impl PoissonEncoder {
pub fn new(steps: usize) -> Self {
Self { num_steps: steps }
}
pub fn encode_rate_hz(&self, rate_hz: f32, dt_seconds: f32) -> Vec<u8> {
self.encode(probability_from_rate_hz(rate_hz, dt_seconds))
}
pub fn encode_rate_hz_step(&self, rate_hz: f32, dt_seconds: f32) -> u8 {
self.encode_step(probability_from_rate_hz(rate_hz, dt_seconds))
}
pub fn encode(&self, input: f32) -> Vec<u8> {
let probability = input.clamp(0.0, 1.0);
let mut rng = rand::rng();
(0..self.num_steps)
.map(|_| {
if crate::rng::gen_unit_f32_with_rng(&mut rng) < probability {
1
} else {
0
}
})
.collect()
}
pub fn encode_step(&self, input: f32) -> u8 {
let probability = input.clamp(0.0, 1.0);
let mut rng = rand::rng();
if crate::rng::gen_unit_f32_with_rng(&mut rng) < probability {
1
} else {
0
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn length_matches_num_steps() {
let enc = PoissonEncoder::new(50);
let spikes = enc.encode(0.5);
assert_eq!(spikes.len(), 50);
}
#[test]
fn zero_input_produces_no_spikes() {
let enc = PoissonEncoder::new(100);
let spikes = enc.encode(0.0);
assert!(spikes.iter().all(|&s| s == 0));
}
#[test]
fn full_input_produces_all_spikes() {
let enc = PoissonEncoder::new(100);
let spikes = enc.encode(1.0);
assert!(spikes.iter().all(|&s| s == 1));
}
#[test]
fn values_are_binary() {
let enc = PoissonEncoder::new(200);
let spikes = enc.encode(0.4);
assert!(spikes.iter().all(|&s| s == 0 || s == 1));
}
#[test]
fn empty_steps_produces_empty() {
let enc = PoissonEncoder::new(0);
let spikes = enc.encode(0.5);
assert_eq!(spikes.len(), 0);
}
#[test]
fn negative_input_clamped_to_zero() {
let enc = PoissonEncoder::new(50);
let spikes = enc.encode(-0.5);
assert!(spikes.iter().all(|&s| s == 0));
}
#[test]
fn above_one_input_clamped_to_one() {
let enc = PoissonEncoder::new(100);
let spikes = enc.encode(1.5);
assert!(spikes.iter().all(|&s| s == 1));
}
#[test]
fn spike_count_produces_mixed_output() {
let enc = PoissonEncoder::new(100);
let spikes = enc.encode(0.5);
let count = spikes.iter().filter(|&&s| s == 1).count();
assert!(
count > 0 && count < 100,
"p=0.5 should produce mixed output, got {} spikes",
count
);
}
#[test]
fn test_poisson_encode_step() {
let enc = PoissonEncoder::new(1);
let mut ones = 0;
let mut zeros = 0;
for _ in 0..100 {
let s = enc.encode_step(0.5);
if s == 1 {
ones += 1;
} else {
zeros += 1;
}
}
assert!(ones > 0 && zeros > 0);
}
#[test]
fn never_panics() {
let enc = PoissonEncoder::new(50);
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| enc.encode(0.5)));
assert!(result.is_ok());
let result =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| enc.encode(f32::NAN)));
assert!(result.is_ok());
let result =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| enc.encode(f32::INFINITY)));
assert!(result.is_ok());
}
#[test]
fn rate_probability_uses_explicit_dt_seconds() {
let probability = probability_from_rate_hz(10.0, 0.01);
let expected = -(-0.1_f32).exp_m1();
assert!((probability - expected).abs() < f32::EPSILON);
}
#[test]
fn tiny_rate_dt_product_stays_positive() {
let probability = probability_from_rate_hz(1.0, 1e-8);
assert!(
probability > 0.0,
"tiny rate*dt must remain a positive Poisson probability, got {probability}"
);
assert!(probability < 1e-6);
assert!(probability < probability_from_rate_hz(1.0, 0.1));
}
#[test]
fn rate_probability_invalid_inputs_are_silent() {
for (rate_hz, dt_seconds) in [
(0.0, 0.01),
(-1.0, 0.01),
(f32::NAN, 0.01),
(10.0, 0.0),
(10.0, f32::NAN),
(10.0, -0.01),
(f32::INFINITY, 0.01),
(10.0, f32::INFINITY),
] {
assert_eq!(probability_from_rate_hz(rate_hz, dt_seconds), 0.0);
}
}
#[test]
fn encode_rate_hz_uses_probability_from_rate() {
let enc = PoissonEncoder::new(200);
let spikes = enc.encode_rate_hz(1_000.0, 1.0);
assert_eq!(spikes.len(), 200);
assert!(spikes.iter().all(|&s| s == 1));
let silent = enc.encode_rate_hz(0.0, 0.01);
assert!(silent.iter().all(|&s| s == 0));
}
#[test]
fn encode_rate_hz_step_returns_binary() {
let enc = PoissonEncoder::new(1);
assert_eq!(enc.encode_rate_hz_step(0.0, 0.01), 0);
let mut ones = 0;
for _ in 0..50 {
ones += enc.encode_rate_hz_step(1_000.0, 1.0) as usize;
}
assert_eq!(ones, 50);
}
}