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#[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; }
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 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}