axon_encoder/encoders/
derivative.rs1use crate::prelude::*;
2
3#[derive(Clone, Debug, PartialEq)]
22#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
23#[cfg_attr(feature = "serde", serde(try_from = "DerivativeEncoderRepr"))]
24pub struct DerivativeEncoder {
25 last_values: Vec<f32>,
26 thresholds: Vec<f32>,
27}
28
29#[cfg(feature = "serde")]
30#[derive(serde::Deserialize)]
31struct DerivativeEncoderRepr {
32 last_values: Vec<f32>,
33 thresholds: Vec<f32>,
34}
35
36#[cfg(feature = "serde")]
37impl TryFrom<DerivativeEncoderRepr> for DerivativeEncoder {
38 type Error = String;
39
40 fn try_from(r: DerivativeEncoderRepr) -> Result<Self, String> {
41 if r.last_values.len() != r.thresholds.len() {
42 return Err(format!(
43 "mismatched last_values length ({}) and thresholds length ({})",
44 r.last_values.len(),
45 r.thresholds.len()
46 ));
47 }
48 if r.last_values.iter().any(|v| !v.is_finite()) {
49 return Err("last_values must be finite".into());
50 }
51 let mut encoder = Self::try_new(r.thresholds).map_err(|error| error.to_string())?;
52 encoder.last_values = r.last_values;
53 Ok(encoder)
54 }
55}
56
57impl DerivativeEncoder {
58 pub fn new(thresholds: Vec<f32>) -> Self {
66 Self::try_new(thresholds).expect("invalid DerivativeEncoder configuration")
67 }
68
69 pub fn try_new(thresholds: Vec<f32>) -> Result<Self, EncoderError> {
73 crate::error::validate_channel_count(thresholds.len())?;
74 for &threshold in &thresholds {
75 crate::error::validate_non_negative_finite("threshold", threshold)?;
76 }
77 let num_channels = thresholds.len();
78 Ok(Self {
79 last_values: vec![0.0; num_channels],
80 thresholds,
81 })
82 }
83}
84
85impl Encoder for DerivativeEncoder {
86 fn encode(&mut self, input: &[f32]) -> EncodedOutput {
87 self.encode_step(input)
88 }
89
90 fn encode_step(&mut self, current_values: &[f32]) -> EncodedOutput {
91 let mut output = EncodedOutput::new();
92
93 for (i, ¤t_val) in current_values.iter().enumerate() {
94 if i >= self.thresholds.len() {
95 break;
96 }
97
98 let delta = current_val - self.last_values[i];
99
100 if delta > self.thresholds[i] {
102 output.spikes.push(SpikeEvent {
103 channel: u16::try_from(i).expect("channel index exceeds u16::MAX"),
104 timestamp: 0,
105 polarity: true,
106 });
107 }
108 else if delta < -self.thresholds[i] {
110 output.spikes.push(SpikeEvent {
111 channel: u16::try_from(i).expect("channel index exceeds u16::MAX"),
112 timestamp: 0,
113 polarity: false,
114 });
115 }
116
117 self.last_values[i] = current_val;
118 }
119 output
120 }
121
122 fn reset(&mut self) {
123 for val in self.last_values.iter_mut() {
124 *val = 0.0;
125 }
126 }
127}
128
129#[cfg(test)]
130mod tests {
131 use super::*;
132
133 #[test]
134 fn test_derivative_encoder_basic() {
135 let mut encoder = DerivativeEncoder::new(vec![1.0, 2.0]);
136
137 let output = encoder.encode(&[1.5, 1.5]);
139 assert_eq!(output.spikes.len(), 1);
140 assert_eq!(output.spikes[0].channel, 0);
141 assert!(output.spikes[0].polarity);
142
143 let output = encoder.encode(&[1.5, 1.5]);
145 assert!(output.spikes.is_empty());
146
147 let output = encoder.encode(&[0.0, 1.5]);
149 assert_eq!(output.spikes.len(), 1);
150 assert_eq!(output.spikes[0].channel, 0);
151 assert!(!output.spikes[0].polarity);
152
153 let output = encoder.encode(&[0.0, 4.0]);
155 assert_eq!(output.spikes.len(), 1);
156 assert_eq!(output.spikes[0].channel, 1);
157 assert!(output.spikes[0].polarity);
158 }
159
160 #[test]
161 fn test_derivative_encoder_reset() {
162 let mut encoder = DerivativeEncoder::new(vec![1.0]);
163 encoder.encode(&[5.0]);
164 encoder.reset();
165 assert_eq!(encoder.last_values[0], 0.0);
166 }
167
168 #[test]
169 fn test_derivative_encoder_empty_and_mismatched() {
170 let mut encoder = DerivativeEncoder::new(vec![1.0]);
171 let output = encoder.encode(&[]);
172 assert!(output.spikes.is_empty());
173
174 let output = encoder.encode(&[2.0, 3.0]);
175 assert_eq!(output.spikes.len(), 1); }
177
178 #[cfg(feature = "serde")]
179 #[test]
180 fn test_derivative_serde_rejects_too_many_channels() {
181 let values: Vec<f32> = vec![0.0; (u16::MAX as usize) + 2];
182 let value = serde_json::json!({
183 "last_values": values.clone(),
184 "thresholds": values,
185 });
186 let res: Result<DerivativeEncoder, _> = serde_json::from_value(value);
187 assert!(res.is_err());
188 }
189
190 #[test]
191 fn test_derivative_encoder_try_new_validation() {
192 assert!(DerivativeEncoder::try_new(vec![0.0, 1.0]).is_ok());
193 assert_eq!(
194 DerivativeEncoder::try_new(vec![f32::NAN]).err(),
195 Some(EncoderError::NonNegativeFinite {
196 parameter: "threshold"
197 })
198 );
199 assert_eq!(
200 DerivativeEncoder::try_new(vec![-1.0]).err(),
201 Some(EncoderError::NonNegativeFinite {
202 parameter: "threshold"
203 })
204 );
205 assert_eq!(
206 DerivativeEncoder::try_new(vec![1.0; u16::MAX as usize + 2]).err(),
207 Some(EncoderError::NumChannelsTooLarge)
208 );
209 }
210}
211
212#[cfg(test)]
213mod branch_coverage_tests {
214 use super::*;
215
216 #[test]
217 fn derivative_encoder_initializes_channel_state() {
218 let encoder = DerivativeEncoder::new(vec![1.0, 2.0, 3.0]);
219 assert_eq!(encoder.thresholds, vec![1.0, 2.0, 3.0]);
220 assert_eq!(encoder.last_values, vec![0.0, 0.0, 0.0]);
221 }
222
223 #[test]
224 fn derivative_encoder_tracks_positive_and_negative_steps() {
225 let mut encoder = DerivativeEncoder::new(vec![1.0, 2.0]);
226
227 let output = encoder.encode_step(&[1.5, -2.5]);
228 assert_eq!(output.spikes.len(), 2);
229 assert_eq!(output.spikes[0].channel, 0);
230 assert!(output.spikes[0].polarity);
231 assert_eq!(output.spikes[1].channel, 1);
232 assert!(!output.spikes[1].polarity);
233 assert_eq!(encoder.last_values, vec![1.5, -2.5]);
234
235 let output = encoder.encode_step(&[2.0, -1.0]);
236 assert!(output.spikes.is_empty());
237 assert_eq!(encoder.last_values, vec![2.0, -1.0]);
238
239 let output = encoder.encode_step(&[0.0, 2.0]);
240 assert_eq!(output.spikes.len(), 2);
241 assert!(!output.spikes[0].polarity);
242 assert!(output.spikes[1].polarity);
243 assert_eq!(encoder.last_values, vec![0.0, 2.0]);
244 }
245
246 #[test]
247 fn derivative_encoder_does_not_fire_at_threshold() {
248 let mut encoder = DerivativeEncoder::new(vec![1.0, 2.0]);
249 let output = encoder.encode_step(&[1.0, -2.0]);
250 assert!(output.spikes.is_empty());
251 assert_eq!(encoder.last_values, vec![1.0, -2.0]);
252 }
253
254 #[test]
255 fn derivative_encoder_handles_channel_count_mismatches() {
256 let mut encoder = DerivativeEncoder::new(vec![1.0, 2.0]);
257 let output = encoder.encode_step(&[2.0, -3.0, 5.0]);
258 assert_eq!(output.spikes.len(), 2);
259 assert_eq!(encoder.last_values, vec![2.0, -3.0]);
260
261 let mut encoder = DerivativeEncoder::new(vec![1.0, 2.0]);
262 let output = encoder.encode_step(&[2.0]);
263 assert_eq!(output.spikes.len(), 1);
264 assert_eq!(output.spikes[0].channel, 0);
265 assert_eq!(encoder.last_values, vec![2.0, 0.0]);
266 }
267}