Skip to main content

sc_neurocore_engine/neurons/rate/
parallel_spiking.rs

1// SPDX-License-Identifier: AGPL-3.0-or-later
2// Commercial license available
3// © Concepts 1996–2026 Miroslav Šotek. All rights reserved.
4// © Code 2020–2026 Miroslav Šotek. All rights reserved.
5// ORCID: 0009-0009-3560-0851
6// Contact: www.anulum.li | protoscience@anulum.li
7// SC-NeuroCore — Parallel spiking neuron model
8
9//! k-order sliding Parallel Spiking Neuron — Fang et al. (2023).
10
11/// k-order sliding PSN — Fang et al. (2023), NeurIPS.
12///
13/// Streaming form of the PSN family (paper Eqs. 14–15):
14///
15/// ```text
16/// H[t] = sum_{i=0}^{k-1} W_i * X[t-k+1+i],  X[j] = 0 for j < 0
17/// S[t] = Theta(H[t] - v_threshold)
18/// ```
19///
20/// `weights[k-1]` multiplies the newest input; the sum accumulates
21/// sequentially from i = 0 so every backend reproduces the same
22/// binary64 result bit-for-bit. `Theta(0) = 1` per the paper. No PSN
23/// variant has a reset: firing never clears the input history. The
24/// paper trains `W` and `v_threshold`; the uniform `1/k` defaults are
25/// repository defaults.
26#[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    /// Construct with the uniform repository default weights `1/k`.
36    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    /// Advance one step after validating the input and configuration.
55    ///
56    /// Computes the hidden state on a candidate window and commits only
57    /// on success: a non-finite input, an invalid configuration, or a
58    /// non-finite hidden state returns `Err` with the pre-step state
59    /// preserved exactly.
60    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    /// Fail-closed wrapper for legacy callers: returns 0 on any rejected
90    /// input without mutating state.
91    pub fn step(&mut self, current: f64) -> i32 {
92        self.try_step(current).unwrap_or(0)
93    }
94
95    /// Clear the retained inputs, preserving weights and threshold.
96    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    /// Independent oracle: explicit zero-padded window per paper Eq. 14.
107    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, &current) 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}