1#[derive(Clone, Debug)]
44#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
45pub struct PoissonEncoder {
46 pub num_steps: usize,
47}
48
49fn 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 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 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 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 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 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 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 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 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 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 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 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}