sc_neurocore_engine/neurons/rate/
parallel_spiking.rs1#[derive(Clone, Debug)]
27pub struct ParallelSpikingNeuron {
28 pub weights: Vec<f64>,
29 pub history: Vec<f64>,
30 pub v_threshold: f64,
31 pub hidden: f64,
32}
33
34impl ParallelSpikingNeuron {
35 pub fn new(kernel_size: usize, v_threshold: f64) -> Self {
37 let k = kernel_size.max(1);
38 Self {
39 weights: vec![1.0 / k as f64; k],
40 history: vec![0.0; k],
41 v_threshold,
42 hidden: 0.0,
43 }
44 }
45
46 fn valid(&self) -> bool {
47 !self.weights.is_empty()
48 && self.history.len() == self.weights.len()
49 && self.weights.iter().all(|w| w.is_finite())
50 && self.history.iter().all(|x| x.is_finite())
51 && self.v_threshold.is_finite()
52 }
53
54 pub fn try_step(&mut self, current: f64) -> Result<i32, &'static str> {
61 if !current.is_finite() {
62 return Err("current must be finite");
63 }
64 if !self.valid() {
65 return Err("sliding PSN state and parameters must be finite");
66 }
67
68 let mut hidden = 0.0_f64;
69 for (index, weight) in self.weights.iter().enumerate() {
70 let value = if index + 1 < self.history.len() {
71 self.history[index + 1]
72 } else {
73 current
74 };
75 hidden += weight * value;
76 }
77 if !hidden.is_finite() {
78 return Err("sliding PSN hidden state became non-finite");
79 }
80
81 self.history.rotate_left(1);
82 if let Some(last) = self.history.last_mut() {
83 *last = current;
84 }
85 self.hidden = hidden;
86 Ok(if hidden >= self.v_threshold { 1 } else { 0 })
87 }
88
89 pub fn step(&mut self, current: f64) -> i32 {
92 self.try_step(current).unwrap_or(0)
93 }
94
95 pub fn reset(&mut self) {
97 self.history.fill(0.0);
98 self.hidden = 0.0;
99 }
100}
101
102#[cfg(test)]
103mod tests {
104 use super::*;
105
106 fn oracle(weights: &[f64], v_th: f64, drive: &[f64]) -> (Vec<f64>, Vec<i32>) {
108 let k = weights.len();
109 let mut hidden_trace = Vec::new();
110 let mut spikes = Vec::new();
111 for t in 0..drive.len() {
112 let mut hidden = 0.0_f64;
113 for (i, w) in weights.iter().enumerate() {
114 let j = t as i64 - k as i64 + 1 + i as i64;
115 let x = if j < 0 { 0.0 } else { drive[j as usize] };
116 hidden += w * x;
117 }
118 hidden_trace.push(hidden);
119 spikes.push(if hidden >= v_th { 1 } else { 0 });
120 }
121 (hidden_trace, spikes)
122 }
123
124 #[test]
125 fn matches_paper_equation_oracle_bit_exactly() {
126 let drive: Vec<f64> = (0..64)
127 .map(|i| 0.4 + 0.3 * (i as f64 * 0.17).sin())
128 .collect();
129 let weights = [0.1, -0.2, 0.35, 0.75];
130 let mut n = ParallelSpikingNeuron::new(4, 0.4);
131 n.weights = weights.to_vec();
132 let (hidden_trace, spikes) = oracle(&weights, 0.4, &drive);
133 for (t, ¤t) in drive.iter().enumerate() {
134 let spike = n.try_step(current).expect("finite configured drive");
135 assert_eq!(n.hidden.to_bits(), hidden_trace[t].to_bits());
136 assert_eq!(spike, spikes[t]);
137 }
138 }
139
140 #[test]
141 fn firing_never_clears_history() {
142 let mut n = ParallelSpikingNeuron::new(4, 0.5);
143 let mut fired = 0;
144 for _ in 0..8 {
145 fired += n.step(1.0);
146 }
147 assert!(
148 fired >= 5,
149 "constant supra-threshold drive must keep firing"
150 );
151 assert!(n.history.iter().all(|&x| x == 1.0));
152 }
153
154 #[test]
155 fn theta_is_right_continuous_at_threshold() {
156 let mut n = ParallelSpikingNeuron::new(1, 1.0);
157 assert_eq!(n.step(1.0), 1);
158 }
159
160 #[test]
161 fn invalid_input_is_rejected_atomically() {
162 let mut n = ParallelSpikingNeuron::new(4, 0.5);
163 n.step(0.7);
164 let before = (n.history.clone(), n.hidden);
165 for bad in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
166 assert!(n.try_step(bad).is_err());
167 assert_eq!(n.step(bad), 0);
168 }
169 assert_eq!((n.history.clone(), n.hidden), before);
170 }
171
172 #[test]
173 fn overflowing_hidden_state_is_rejected_atomically() {
174 let mut n = ParallelSpikingNeuron::new(2, 0.5);
175 n.weights = vec![f64::MAX, f64::MAX];
176 n.history = vec![f64::MAX, f64::MAX];
177 assert!(n.try_step(f64::MAX).is_err());
178 assert_eq!(n.history, vec![f64::MAX, f64::MAX]);
179 }
180
181 #[test]
182 fn reset_clears_history_only() {
183 let mut n = ParallelSpikingNeuron::new(4, 0.5);
184 n.step(1.0);
185 n.reset();
186 assert!(n.history.iter().all(|&x| x == 0.0));
187 assert_eq!(n.hidden, 0.0);
188 assert_eq!(n.weights, vec![0.25; 4]);
189 }
190}