Skip to main content

axon_encoder/
encoder.rs

1use crate::types::{EncodedOutput, SpikeEvent};
2
3#[derive(Clone, Debug, PartialEq)]
4#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
5#[cfg_attr(feature = "serde", serde(try_from = "EmbeddingEncoderConfigRepr"))]
6pub struct EmbeddingEncoderConfig {
7    pub v_th: f32,
8}
9
10#[cfg(feature = "serde")]
11#[derive(serde::Deserialize)]
12struct EmbeddingEncoderConfigRepr {
13    v_th: f32,
14}
15
16#[cfg(feature = "serde")]
17impl TryFrom<EmbeddingEncoderConfigRepr> for EmbeddingEncoderConfig {
18    type Error = String;
19
20    fn try_from(r: EmbeddingEncoderConfigRepr) -> Result<Self, String> {
21        if r.v_th.partial_cmp(&0.0) != Some(core::cmp::Ordering::Greater) {
22            return Err("v_th must be positive".into());
23        }
24        Ok(Self { v_th: r.v_th })
25    }
26}
27
28#[derive(Clone, Debug, PartialEq)]
29#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
30pub struct EncoderState {
31    pub membrane_potentials: Vec<f32>,
32}
33
34impl EncoderState {
35    pub fn new_zeros(len: usize) -> Self {
36        Self {
37            membrane_potentials: vec![0.0; len],
38        }
39    }
40}
41
42/// Rate-style membrane encoder driven by a fixed embedding vector.
43///
44/// Accumulates normalized embedding components into per-channel membrane
45/// potentials and emits a spike when the threshold is crossed (soft reset).
46///
47/// # Examples
48///
49/// ```rust
50/// use axon_encoder::encoder::{EmbeddingEncoderConfig, EmbeddingRateEncoder, EncoderState};
51///
52/// let enc = EmbeddingRateEncoder::new(&[0.5, 1.0], EmbeddingEncoderConfig { v_th: 0.4 });
53/// let state = EncoderState::new_zeros(2);
54/// let (out, next) = enc.forward(&state);
55/// assert!(!out.spikes.is_empty());
56/// assert_eq!(next.membrane_potentials.len(), 2);
57/// ```
58#[derive(Clone, Debug, PartialEq)]
59#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
60#[cfg_attr(feature = "serde", serde(try_from = "EmbeddingRateEncoderRepr"))]
61pub struct EmbeddingRateEncoder {
62    pub config: EmbeddingEncoderConfig,
63    pub normalized_embeddings: Vec<f32>,
64}
65
66#[cfg(feature = "serde")]
67#[derive(serde::Deserialize)]
68struct EmbeddingRateEncoderRepr {
69    config: EmbeddingEncoderConfig,
70    normalized_embeddings: Vec<f32>,
71}
72
73#[cfg(feature = "serde")]
74impl TryFrom<EmbeddingRateEncoderRepr> for EmbeddingRateEncoder {
75    type Error = String;
76
77    fn try_from(r: EmbeddingRateEncoderRepr) -> Result<Self, String> {
78        if r.normalized_embeddings.iter().any(|v| !v.is_finite()) {
79            return Err("normalized_embeddings must be finite".into());
80        }
81        if r.normalized_embeddings.len() > u16::MAX as usize + 1 {
82            return Err("too many channels (max 65536)".into());
83        }
84        Ok(Self {
85            config: r.config,
86            normalized_embeddings: r.normalized_embeddings,
87        })
88    }
89}
90
91impl EmbeddingRateEncoder {
92    pub fn new(embeddings: &[f32], config: EmbeddingEncoderConfig) -> Self {
93        if config.v_th.partial_cmp(&0.0) != Some(core::cmp::Ordering::Greater) {
94            panic!("v_th must be positive");
95        }
96        assert!(
97            embeddings.len() <= u16::MAX as usize + 1,
98            "too many channels (max 65536)"
99        );
100
101        let min_val = embeddings.iter().copied().fold(f32::INFINITY, f32::min);
102        let max_val = embeddings.iter().copied().fold(f32::NEG_INFINITY, f32::max);
103        let range = max_val - min_val;
104        let epsilon = 1e-5f32;
105        let safe_range = range + epsilon;
106
107        let normalized: Vec<f32> = embeddings
108            .iter()
109            .map(|&x| (x - min_val) / safe_range)
110            .collect();
111
112        Self {
113            config,
114            normalized_embeddings: normalized,
115        }
116    }
117
118    pub fn forward(&self, prev_state: &EncoderState) -> (EncodedOutput, EncoderState) {
119        let mut new_potentials = prev_state.membrane_potentials.clone();
120        let mut output = EncodedOutput::new();
121
122        for (i, (pot, &emb)) in new_potentials
123            .iter_mut()
124            .zip(self.normalized_embeddings.iter())
125            .enumerate()
126        {
127            *pot += emb;
128
129            if *pot >= self.config.v_th {
130                output.spikes.push(SpikeEvent {
131                    channel: u16::try_from(i).expect("channel index exceeds u16::MAX"),
132                    timestamp: 0,
133                    polarity: true,
134                });
135                *pot -= self.config.v_th; // Soft reset
136            }
137        }
138
139        (
140            output,
141            EncoderState {
142                membrane_potentials: new_potentials,
143            },
144        )
145    }
146}
147
148#[cfg(test)]
149mod tests {
150    use super::*;
151
152    #[test]
153    fn test_embedding_rate_encoder_basic() {
154        let config = EmbeddingEncoderConfig { v_th: 0.9 };
155        let embeddings = [0.5, 1.0, 0.0];
156        let encoder = EmbeddingRateEncoder::new(&embeddings, config);
157
158        let state = EncoderState::new_zeros(3);
159        let (output, next_state) = encoder.forward(&state);
160
161        assert_eq!(output.spikes.len(), 1);
162        assert_eq!(output.spikes[0].channel, 1);
163
164        let (output2, _) = encoder.forward(&next_state);
165        // Channel 0: 0.5 + 0.5 = 1.0 > 0.9 -> spike
166        // Channel 1: (1.0-0.9) + 1.0 = 1.1 > 0.9 -> spike
167        assert_eq!(output2.spikes.len(), 2);
168    }
169
170    #[test]
171    #[should_panic(expected = "v_th must be positive")]
172    fn test_embedding_encoder_config_invalid_vth() {
173        let _ = EmbeddingRateEncoder::new(&[0.5], EmbeddingEncoderConfig { v_th: 0.0 });
174    }
175
176    #[test]
177    fn test_encoder_state_new_zeros() {
178        let state = EncoderState::new_zeros(5);
179        assert_eq!(state.membrane_potentials.len(), 5);
180        assert!(state.membrane_potentials.iter().all(|&v| v == 0.0));
181    }
182
183    #[cfg(feature = "serde")]
184    #[test]
185    fn test_embedding_rate_encoder_deserialize_rejects_too_many_channels() {
186        let embeddings: Vec<f32> = vec![0.0; (u16::MAX as usize) + 2];
187        let value = serde_json::json!({
188            "config": {"v_th": 1.0},
189            "normalized_embeddings": embeddings,
190        });
191        let res: Result<EmbeddingRateEncoder, _> = serde_json::from_value(value);
192        assert!(res.is_err());
193    }
194}
195
196#[cfg(test)]
197mod forward_coverage_tests {
198    use super::*;
199
200    #[test]
201    fn embedding_rate_encoder_initializes_normalized_values() {
202        let embeddings = vec![1.0, 2.0, 3.0, 5.0];
203        let encoder = EmbeddingRateEncoder::new(&embeddings, EmbeddingEncoderConfig { v_th: 1.0 });
204        assert_eq!(encoder.normalized_embeddings.len(), 4);
205        assert!((encoder.normalized_embeddings[0] - 0.0).abs() < 1e-5);
206        assert!((encoder.normalized_embeddings[3] - (4.0 / 4.00001)).abs() < 1e-5);
207    }
208
209    #[test]
210    fn embedding_rate_encoder_forward_without_spikes() {
211        let embeddings = vec![1.0, 2.0, 3.0];
212        let encoder = EmbeddingRateEncoder::new(&embeddings, EmbeddingEncoderConfig { v_th: 10.0 });
213        let state = EncoderState::new_zeros(3);
214        let (output, next_state) = encoder.forward(&state);
215
216        assert!(output.spikes.is_empty());
217        assert_eq!(next_state.membrane_potentials.len(), 3);
218        assert_eq!(
219            next_state.membrane_potentials,
220            encoder.normalized_embeddings
221        );
222    }
223
224    #[test]
225    fn embedding_rate_encoder_forward_soft_resets_spikes() {
226        let embeddings = vec![1.0, 2.0, 3.0];
227        let encoder = EmbeddingRateEncoder::new(&embeddings, EmbeddingEncoderConfig { v_th: 0.4 });
228        let state = EncoderState::new_zeros(3);
229        let (output, next_state) = encoder.forward(&state);
230
231        assert_eq!(output.spikes.len(), 2);
232        assert_eq!(output.spikes[0].channel, 1);
233        assert_eq!(output.spikes[1].channel, 2);
234        assert!((next_state.membrane_potentials[0] - 0.0).abs() < 1e-5);
235        assert!(
236            (next_state.membrane_potentials[1] - (encoder.normalized_embeddings[1] - 0.4)).abs()
237                < 1e-5
238        );
239        assert!(
240            (next_state.membrane_potentials[2] - (encoder.normalized_embeddings[2] - 0.4)).abs()
241                < 1e-5
242        );
243    }
244
245    #[test]
246    fn embedding_rate_encoder_accumulates_across_steps() {
247        let embeddings = vec![1.0, 3.0];
248        let encoder = EmbeddingRateEncoder::new(&embeddings, EmbeddingEncoderConfig { v_th: 0.6 });
249        let mut state = EncoderState::new_zeros(2);
250
251        for _ in 0..3 {
252            let (output, next_state) = encoder.forward(&state);
253            assert_eq!(output.spikes.len(), 1);
254            assert_eq!(output.spikes[0].channel, 1);
255            state = next_state;
256        }
257
258        assert!((state.membrane_potentials[0] - 0.0).abs() < 1e-5);
259        assert!(
260            (state.membrane_potentials[1] - (3.0 * encoder.normalized_embeddings[1] - 1.8)).abs()
261                < 1e-5
262        );
263    }
264
265    #[test]
266    fn embedding_rate_encoder_handles_equal_embeddings() {
267        let embeddings = vec![2.5, 2.5, 2.5];
268        let encoder = EmbeddingRateEncoder::new(&embeddings, EmbeddingEncoderConfig { v_th: 0.5 });
269        assert!(
270            encoder
271                .normalized_embeddings
272                .iter()
273                .all(|value| *value == 0.0)
274        );
275
276        let (output, next_state) = encoder.forward(&EncoderState::new_zeros(3));
277        assert!(output.spikes.is_empty());
278        assert_eq!(next_state.membrane_potentials, vec![0.0, 0.0, 0.0]);
279    }
280}