sc_neurocore_engine/bindings/
adex_neuron.rs1use numpy::{IntoPyArray, PyArray1};
10use pyo3::exceptions::PyFloatingPointError;
11use pyo3::prelude::*;
12use pyo3::types::PyDict;
13
14use crate::neuron;
15
16type CompleteTracePacket<'py> = (
17 Bound<'py, PyArray1<f64>>,
18 Bound<'py, PyArray1<f64>>,
19 Bound<'py, PyArray1<u8>>,
20 f64,
21 f64,
22);
23
24#[pyclass(
25 name = "AdExNeuron",
26 module = "sc_neurocore_engine.sc_neurocore_engine"
27)]
28#[derive(Clone)]
29pub struct PyAdExNeuron {
30 inner: neuron::AdExNeuron,
31}
32
33#[pymethods]
34impl PyAdExNeuron {
35 #[new]
36 fn new() -> Self {
37 Self {
38 inner: neuron::AdExNeuron::new(),
39 }
40 }
41 fn step(&mut self, current: f64) -> i32 {
42 self.inner.step(current)
43 }
44 fn reset(&mut self) {
45 self.inner.reset();
46 }
47 fn get_state(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
48 let d = PyDict::new(py);
49 d.set_item("v", self.inner.v)?;
50 d.set_item("w", self.inner.w)?;
51 Ok(d.into_any().unbind())
52 }
53}
54
55pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
57 module.add_class::<PyAdExNeuron>()?;
58 module.add_function(wrap_pyfunction!(adex_simulate_complete, module)?)?;
59 Ok(())
60}
61
62#[pyfunction]
64#[pyo3(signature = (
65 v, w, v_rest, v_reset, v_threshold, v_rh, delta_t, tau, tau_w,
66 a, b, c_m, dt, n_steps, current
67))]
68#[allow(clippy::too_many_arguments)]
69fn adex_simulate_complete<'py>(
70 py: Python<'py>,
71 v: f64,
72 w: f64,
73 v_rest: f64,
74 v_reset: f64,
75 v_threshold: f64,
76 v_rh: f64,
77 delta_t: f64,
78 tau: f64,
79 tau_w: f64,
80 a: f64,
81 b: f64,
82 c_m: f64,
83 dt: f64,
84 n_steps: usize,
85 current: f64,
86) -> PyResult<CompleteTracePacket<'py>> {
87 let mut model = neuron::AdExNeuron {
88 v,
89 w,
90 v_rest,
91 v_reset,
92 v_threshold,
93 v_rh,
94 delta_t,
95 tau,
96 tau_w,
97 a,
98 b,
99 c_m,
100 dt,
101 };
102 let (v_trace, w_trace, event_trace) = model
103 .simulate_complete(n_steps, current)
104 .map_err(PyFloatingPointError::new_err)?;
105 Ok((
106 v_trace.into_pyarray(py),
107 w_trace.into_pyarray(py),
108 event_trace.into_pyarray(py),
109 model.v,
110 model.w,
111 ))
112}