Skip to main content

axon_encoder/
poisson.rs

1/// Poisson spike train encoder.
2///
3/// Generates spike trains with Poisson-distributed timing based on either a per-step
4/// probability or an explicit firing rate and time step.
5///
6/// # Mathematical Model
7///
8/// ```text
9/// // Dimensionless probability input:
10/// probability = clamp(input, 0.0, 1.0)
11/// spike[i] = 1 if random() < probability else 0
12///
13/// // Physical rate input:
14/// probability = 1 - exp(-rate_hz * dt_seconds)  // via -exp_m1(-x) in f32
15/// ```
16///
17/// # When to Use
18///
19/// - Generating baseline spike trains with controllable average rates
20/// - Poisson-like random spike generation for stochastic encoders
21/// - Creating temporal patterns with controllable firing rates
22///
23/// # Note
24///
25/// This encoder is NOT part of the `Encoder` trait because its output type (`Vec<u8>`)
26/// differs from other encoders (`EncodedOutput`). It operates in a different mode:
27/// the input is a single probability (0.0 to 1.0) and the output is a spike train
28/// over multiple time steps.
29///
30/// # Examples
31///
32/// ```rust
33/// use axon_encoder::prelude::*;
34///
35/// let enc = PoissonEncoder::new(32);
36/// let train = enc.encode(0.0); // never spikes
37/// assert_eq!(train.len(), 32);
38/// assert!(train.iter().all(|&b| b == 0));
39///
40/// let p = probability_from_rate_hz(50.0, 0.001);
41/// assert!((0.0..=1.0).contains(&p));
42/// ```
43#[derive(Clone, Debug)]
44#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
45pub struct PoissonEncoder {
46    pub num_steps: usize,
47}
48
49/// Converts a firing rate in hertz and a time-bin width in seconds into the
50/// per-bin spike probability for a homogeneous Poisson process.
51///
52/// Mathematically this is `1 - exp(-rate_hz * dt_seconds)`. The implementation
53/// uses `exp_m1` so tiny products (high sample rate, low Hz) stay nonzero in
54/// `f32` instead of rounding to `0.0`. Non-finite or non-positive rates produce
55/// `0.0`; invalid `dt_seconds` is treated as silent here so stochastic paths
56/// never emit NaN probabilities (callers should still validate `dt_seconds`
57/// when constructing encoders).
58fn rate_dt_produces_spikes(rate_hz: f32, dt_seconds: f32) -> bool {
59    rate_hz.is_finite() && rate_hz > 0.0 && dt_seconds.is_finite() && dt_seconds > 0.0
60}
61
62pub fn probability_from_rate_hz(rate_hz: f32, dt_seconds: f32) -> f32 {
63    if !rate_dt_produces_spikes(rate_hz, dt_seconds) {
64        return 0.0;
65    }
66    // 1 - exp(-x) == -exp_m1(-x); exp_m1 stays accurate for tiny x in f32.
67    let x = rate_hz * dt_seconds;
68    (-(-x).exp_m1()).clamp(0.0, 1.0)
69}
70
71impl PoissonEncoder {
72    pub fn new(steps: usize) -> Self {
73        Self { num_steps: steps }
74    }
75
76    /// Encodes a firing rate in hertz into a spike train using an explicit time
77    /// step in seconds for each bin.
78    pub fn encode_rate_hz(&self, rate_hz: f32, dt_seconds: f32) -> Vec<u8> {
79        self.encode(probability_from_rate_hz(rate_hz, dt_seconds))
80    }
81
82    /// Encodes a single rate-based step using an explicit time step in seconds.
83    pub fn encode_rate_hz_step(&self, rate_hz: f32, dt_seconds: f32) -> u8 {
84        self.encode_step(probability_from_rate_hz(rate_hz, dt_seconds))
85    }
86
87    /// Encodes a single probability value into a spike train.
88    ///
89    /// Each of the `num_steps` represents an independent time step where
90    /// a spike occurs with the given probability.
91    pub fn encode(&self, input: f32) -> Vec<u8> {
92        let probability = input.clamp(0.0, 1.0);
93        let mut rng = rand::rng();
94        (0..self.num_steps)
95            .map(|_| {
96                if crate::rng::gen_unit_f32_with_rng(&mut rng) < probability {
97                    1
98                } else {
99                    0
100                }
101            })
102            .collect()
103    }
104
105    /// Encodes a single step - returns 1 or 0 based on input probability.
106    ///
107    /// Useful for streaming mode where you want one spike decision at a time.
108    pub fn encode_step(&self, input: f32) -> u8 {
109        let probability = input.clamp(0.0, 1.0);
110        let mut rng = rand::rng();
111        if crate::rng::gen_unit_f32_with_rng(&mut rng) < probability {
112            1
113        } else {
114            0
115        }
116    }
117}
118
119#[cfg(test)]
120mod tests {
121    use super::*;
122
123    #[test]
124    fn length_matches_num_steps() {
125        let enc = PoissonEncoder::new(50);
126        let spikes = enc.encode(0.5);
127        assert_eq!(spikes.len(), 50);
128    }
129
130    #[test]
131    fn zero_input_produces_no_spikes() {
132        let enc = PoissonEncoder::new(100);
133        let spikes = enc.encode(0.0);
134        assert!(spikes.iter().all(|&s| s == 0));
135    }
136
137    #[test]
138    fn full_input_produces_all_spikes() {
139        let enc = PoissonEncoder::new(100);
140        let spikes = enc.encode(1.0);
141        assert!(spikes.iter().all(|&s| s == 1));
142    }
143
144    #[test]
145    fn values_are_binary() {
146        let enc = PoissonEncoder::new(200);
147        let spikes = enc.encode(0.4);
148        assert!(spikes.iter().all(|&s| s == 0 || s == 1));
149    }
150
151    #[test]
152    fn empty_steps_produces_empty() {
153        let enc = PoissonEncoder::new(0);
154        let spikes = enc.encode(0.5);
155        assert_eq!(spikes.len(), 0);
156    }
157
158    #[test]
159    fn negative_input_clamped_to_zero() {
160        let enc = PoissonEncoder::new(50);
161        let spikes = enc.encode(-0.5);
162        assert!(spikes.iter().all(|&s| s == 0));
163    }
164
165    #[test]
166    fn above_one_input_clamped_to_one() {
167        let enc = PoissonEncoder::new(100);
168        let spikes = enc.encode(1.5);
169        assert!(spikes.iter().all(|&s| s == 1));
170    }
171
172    #[test]
173    fn spike_count_produces_mixed_output() {
174        let enc = PoissonEncoder::new(100);
175        let spikes = enc.encode(0.5);
176        let count = spikes.iter().filter(|&&s| s == 1).count();
177        assert!(
178            count > 0 && count < 100,
179            "p=0.5 should produce mixed output, got {} spikes",
180            count
181        );
182    }
183
184    #[test]
185    fn test_poisson_encode_step() {
186        let enc = PoissonEncoder::new(1);
187        let mut ones = 0;
188        let mut zeros = 0;
189        for _ in 0..100 {
190            let s = enc.encode_step(0.5);
191            if s == 1 {
192                ones += 1;
193            } else {
194                zeros += 1;
195            }
196        }
197        assert!(ones > 0 && zeros > 0);
198    }
199
200    #[test]
201    fn never_panics() {
202        let enc = PoissonEncoder::new(50);
203        let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| enc.encode(0.5)));
204        assert!(result.is_ok());
205        let result =
206            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| enc.encode(f32::NAN)));
207        assert!(result.is_ok());
208        let result =
209            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| enc.encode(f32::INFINITY)));
210        assert!(result.is_ok());
211    }
212
213    #[test]
214    fn rate_probability_uses_explicit_dt_seconds() {
215        let probability = probability_from_rate_hz(10.0, 0.01);
216        // Match the exp_m1 implementation (equivalent to 1 - exp(-x) for this x).
217        let expected = -(-0.1_f32).exp_m1();
218        assert!((probability - expected).abs() < f32::EPSILON);
219    }
220
221    #[test]
222    fn tiny_rate_dt_product_stays_positive() {
223        // 1 Hz at 10 ns: naive 1 - exp(-x) rounds to 0 in f32; exp_m1 keeps it > 0.
224        let probability = probability_from_rate_hz(1.0, 1e-8);
225        assert!(
226            probability > 0.0,
227            "tiny rate*dt must remain a positive Poisson probability, got {probability}"
228        );
229        assert!(probability < 1e-6);
230        // Sanity: also smaller than the large-x path.
231        assert!(probability < probability_from_rate_hz(1.0, 0.1));
232    }
233
234    #[test]
235    fn rate_probability_invalid_inputs_are_silent() {
236        for (rate_hz, dt_seconds) in [
237            (0.0, 0.01),
238            (-1.0, 0.01),
239            (f32::NAN, 0.01),
240            (10.0, 0.0),
241            (10.0, f32::NAN),
242            (10.0, -0.01),
243            (f32::INFINITY, 0.01),
244            (10.0, f32::INFINITY),
245        ] {
246            assert_eq!(probability_from_rate_hz(rate_hz, dt_seconds), 0.0);
247        }
248    }
249
250    #[test]
251    fn encode_rate_hz_uses_probability_from_rate() {
252        let enc = PoissonEncoder::new(200);
253        // High rate * dt saturates probability to ~1.0 so every bin spikes.
254        let spikes = enc.encode_rate_hz(1_000.0, 1.0);
255        assert_eq!(spikes.len(), 200);
256        assert!(spikes.iter().all(|&s| s == 1));
257
258        // Zero / invalid rates produce a silent train.
259        let silent = enc.encode_rate_hz(0.0, 0.01);
260        assert!(silent.iter().all(|&s| s == 0));
261    }
262
263    #[test]
264    fn encode_rate_hz_step_returns_binary() {
265        let enc = PoissonEncoder::new(1);
266        assert_eq!(enc.encode_rate_hz_step(0.0, 0.01), 0);
267        // Saturated rate almost always spikes; sample enough to be robust.
268        let mut ones = 0;
269        for _ in 0..50 {
270            ones += enc.encode_rate_hz_step(1_000.0, 1.0) as usize;
271        }
272        assert_eq!(ones, 50);
273    }
274}