sc_neurocore_engine/bindings/rate/
parallel_spiking.rs1use pyo3::exceptions::PyValueError;
10use pyo3::prelude::*;
11use pyo3::types::PyDict;
12
13use crate::neurons;
14
15#[pyclass(
16 name = "ParallelSpikingNeuron",
17 module = "sc_neurocore_engine.sc_neurocore_engine"
18)]
19#[derive(Clone)]
20pub struct PyParallelSpikingNeuron {
21 inner: neurons::ParallelSpikingNeuron,
22}
23
24#[pymethods]
25impl PyParallelSpikingNeuron {
26 #[new]
27 #[pyo3(signature = (kernel_size=8, v_threshold=1.0))]
28 fn new(kernel_size: usize, v_threshold: f64) -> PyResult<Self> {
29 if kernel_size < 1 {
30 return Err(PyValueError::new_err(
31 "kernel_size must be a positive integer",
32 ));
33 }
34 Ok(Self {
35 inner: neurons::ParallelSpikingNeuron::new(kernel_size, v_threshold),
36 })
37 }
38
39 fn step(&mut self, current: f64) -> PyResult<i32> {
42 self.inner.try_step(current).map_err(PyValueError::new_err)
43 }
44
45 fn set_weights(&mut self, weights: Vec<f64>) -> PyResult<()> {
47 if weights.len() != self.inner.weights.len() {
48 return Err(PyValueError::new_err(
49 "weights must have exactly kernel_size entries",
50 ));
51 }
52 if !weights.iter().all(|w| w.is_finite()) {
53 return Err(PyValueError::new_err("weights must be finite"));
54 }
55 self.inner.weights = weights;
56 Ok(())
57 }
58
59 fn reset(&mut self) {
61 self.inner.reset();
62 }
63
64 fn get_state(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
66 let state = PyDict::new(py);
67 state.set_item("hidden", self.inner.hidden)?;
68 state.set_item("history", self.inner.history.clone())?;
69 Ok(state.into_any().unbind())
70 }
71}
72
73pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
74 module.add_class::<PyParallelSpikingNeuron>()?;
75 Ok(())
76}