Skip to main content

axon_encoder/encoders/
derivative.rs

1use crate::prelude::*;
2
3/// Encodes based on the rate of change (derivative) of the input.
4///
5/// Fires an excitatory spike when the positive change exceeds a threshold,
6/// and an inhibitory spike when the negative change exceeds the threshold.
7///
8/// # Examples
9///
10/// ```rust
11/// use axon_encoder::prelude::*;
12/// # fn main() -> Result<(), EncoderError> {
13/// let mut enc = DerivativeEncoder::try_new(vec![0.2])?;
14/// let _ = enc.encode_step(&[0.0]); // seed previous sample
15/// let out = enc.encode_step(&[0.5]); // +0.5 change exceeds threshold
16/// assert_eq!(out.spikes.len(), 1);
17/// assert!(out.spikes[0].polarity);
18/// # Ok(())
19/// # }
20/// ```
21#[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    /// Creates a new `DerivativeEncoder`, panicking if configuration is invalid.
59    ///
60    /// Prefer [`try_new`](Self::try_new) for typed validation errors.
61    ///
62    /// # Panics
63    ///
64    /// Panics if any threshold is non-finite/negative or the channel count is too large.
65    pub fn new(thresholds: Vec<f32>) -> Self {
66        Self::try_new(thresholds).expect("invalid DerivativeEncoder configuration")
67    }
68
69    /// Creates a new `DerivativeEncoder`, returning an [`EncoderError`] for invalid configuration.
70    ///
71    /// Each threshold must be finite and non-negative; channel count must fit `u16` IDs.
72    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, &current_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            // Excitatory spike on positive jump exceeding threshold
101            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            // Inhibitory/Negative spike on sudden drop
109            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        // Initial jump
138        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        // Stay same
144        let output = encoder.encode(&[1.5, 1.5]);
145        assert!(output.spikes.is_empty());
146
147        // Jump down
148        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        // Jump up on channel 1
154        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); // Only channel 0 should be processed
176    }
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}