sc_neurocore_engine/bindings/
sigmoid_rate.rs1use numpy::{IntoPyArray, PyArray1};
12use pyo3::exceptions::PyValueError;
13use pyo3::prelude::*;
14use pyo3::types::PyDict;
15
16use crate::neurons;
17
18pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
20 module.add_class::<PySigmoidRateNeuron>()?;
21 module.add_function(wrap_pyfunction!(py_sigmoid_rate_simulate, module)?)?;
22 Ok(())
23}
24
25#[pyclass(
27 name = "SigmoidRateNeuron",
28 module = "sc_neurocore_engine.sc_neurocore_engine"
29)]
30#[derive(Clone)]
31pub struct PySigmoidRateNeuron {
32 inner: neurons::SigmoidRateNeuron,
33}
34
35#[pymethods]
36impl PySigmoidRateNeuron {
37 #[new]
38 #[pyo3(signature = (r=0.0, tau=10.0, beta=1.0, theta=0.0, dt=0.1))]
39 fn new(r: f64, tau: f64, beta: f64, theta: f64, dt: f64) -> PyResult<Self> {
40 Ok(Self {
41 inner: neurons::SigmoidRateNeuron::with_parameters(r, tau, beta, theta, dt)
42 .map_err(PyValueError::new_err)?,
43 })
44 }
45 fn step(&mut self, current: f64) -> PyResult<f64> {
46 self.inner.try_step(current).map_err(PyValueError::new_err)
47 }
48 fn reset(&mut self) {
49 self.inner.reset();
50 }
51 fn get_state(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
52 let d = PyDict::new(py);
53 d.set_item("r", self.inner.r)?;
54 d.set_item("tau", self.inner.tau)?;
55 d.set_item("beta", self.inner.beta)?;
56 d.set_item("theta", self.inner.theta)?;
57 d.set_item("dt", self.inner.dt)?;
58 Ok(d.into_any().unbind())
59 }
60}
61
62fn simulate_sigmoid_rate(
63 r: f64,
64 tau: f64,
65 beta: f64,
66 theta: f64,
67 dt: f64,
68 n_steps: usize,
69 current: f64,
70) -> Result<(Vec<f64>, f64), String> {
71 let mut neuron = crate::neurons::SigmoidRateNeuron::with_parameters(r, tau, beta, theta, dt)?;
72 let mut trace = Vec::with_capacity(n_steps);
73 for _ in 0..n_steps {
74 trace.push(neuron.try_step(current)?);
75 }
76 Ok((trace, neuron.r))
77}
78
79#[pyfunction]
81#[pyo3(signature = (r, tau, beta, theta, dt, n_steps, current))]
82fn py_sigmoid_rate_simulate<'py>(
83 py: Python<'py>,
84 r: f64,
85 tau: f64,
86 beta: f64,
87 theta: f64,
88 dt: f64,
89 n_steps: usize,
90 current: f64,
91) -> PyResult<(Bound<'py, PyArray1<f64>>, f64)> {
92 let (trace, final_rate) = simulate_sigmoid_rate(r, tau, beta, theta, dt, n_steps, current)
93 .map_err(PyValueError::new_err)?;
94 Ok((trace.into_pyarray(py), final_rate))
95}
96
97#[cfg(test)]
98mod tests {
99 use super::*;
100
101 #[test]
102 fn batch_matches_python_exact_relaxation_golden() {
103 let (trace, final_rate) = simulate_sigmoid_rate(0.25, 10.0, 2.0, 1.0, 0.5, 6, 3.0).unwrap();
104 let expected = [
105 0.2857007338135623,
106 0.3196603222932904,
107 0.3519636820991432,
108 0.38269158845670403,
109 0.41192087713731845,
110 0.43972463658754457,
111 ];
112 assert_eq!(trace.len(), expected.len());
113 for (actual, target) in trace.into_iter().zip(expected) {
114 assert!((actual - target).abs() <= 2.0e-15, "{actual} != {target}");
115 }
116 assert!((final_rate - expected[5]).abs() <= 2.0e-15);
117 }
118
119 #[test]
120 fn empty_batch_preserves_initial_rate() {
121 let (trace, final_rate) = simulate_sigmoid_rate(0.25, 10.0, 2.0, 1.0, 0.5, 0, 3.0).unwrap();
122 assert!(trace.is_empty());
123 assert_eq!(final_rate, 0.25);
124 }
125
126 #[test]
127 fn batch_rejects_invalid_contracts() {
128 assert!(simulate_sigmoid_rate(-0.1, 10.0, 2.0, 1.0, 0.5, 1, 3.0).is_err());
129 assert!(simulate_sigmoid_rate(0.25, 0.0, 2.0, 1.0, 0.5, 1, 3.0).is_err());
130 assert!(simulate_sigmoid_rate(0.25, 10.0, 2.0, 1.0, 0.5, 1, f64::NAN).is_err());
131 }
132
133 #[test]
134 fn large_timestep_batch_remains_in_unit_interval() {
135 let (trace, final_rate) =
136 simulate_sigmoid_rate(1.0, 0.1, 1.0, 0.0, 5.0, 2, -100.0).unwrap();
137 assert!(trace.iter().all(|rate| (0.0..=1.0).contains(rate)));
138 assert!((0.0..=1.0).contains(&final_rate));
139 }
140}