sc_neurocore_engine/bindings/stochastic/
glm.rs1use pyo3::exceptions::PyValueError;
10use pyo3::prelude::*;
11use pyo3::types::PyDict;
12
13use crate::neurons;
14
15#[pyclass(name = "GLMNeuron", module = "sc_neurocore_engine.sc_neurocore_engine")]
16#[derive(Clone)]
17pub struct PyGLMNeuron {
18 inner: neurons::GLMNeuron,
19}
20
21#[pymethods]
22impl PyGLMNeuron {
23 #[new]
24 #[pyo3(signature = (n_k=10, n_h=20, seed=42))]
25 fn new(n_k: usize, n_h: usize, seed: u64) -> Self {
26 Self {
27 inner: neurons::GLMNeuron::new(n_k, n_h, seed),
28 }
29 }
30
31 #[staticmethod]
34 #[pyo3(signature = (n_k=10, n_h=20, seed=42))]
35 fn legacy_constant_filters(n_k: usize, n_h: usize, seed: u64) -> Self {
36 Self {
37 inner: neurons::GLMNeuron::new_legacy_constant_filters(n_k, n_h, seed),
38 }
39 }
40
41 #[pyo3(signature = (stimulus, uniform=None))]
45 fn step(&mut self, stimulus: f64, uniform: Option<f64>) -> PyResult<i32> {
46 self.inner
47 .try_step(stimulus, uniform)
48 .map_err(PyValueError::new_err)
49 }
50
51 fn reset(&mut self) {
53 self.inner.reset();
54 }
55
56 fn get_state(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
58 let state = PyDict::new(py);
59 state.set_item("mu", self.inner.mu)?;
60 state.set_item("dt_ms", self.inner.dt_ms)?;
61 state.set_item("k", self.inner.k.clone())?;
62 state.set_item("h", self.inner.h.clone())?;
63 state.set_item("stim_buf", self.inner.stim_buf_view())?;
64 state.set_item("spike_buf", self.inner.spike_buf_view())?;
65 Ok(state.into_any().unbind())
66 }
67}
68
69pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
70 module.add_class::<PyGLMNeuron>()?;
71 Ok(())
72}