Skip to main content

sc_neurocore_engine/
rk4_neurons.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 — RK4 neuron integrator ports
8
9//! Explicit RK4 ports for the priority neuron integrator paths.
10
11use numpy::{IntoPyArray, PyReadonlyArray1};
12use pyo3::exceptions::PyValueError;
13use pyo3::prelude::*;
14use pyo3::types::PyDict;
15
16const IZH_SPIKE_THRESHOLD: f64 = 30.0;
17
18#[derive(Clone, Debug)]
19pub struct IzhikevichRk4 {
20    pub v: f64,
21    pub u: f64,
22    pub a: f64,
23    pub b: f64,
24    pub c: f64,
25    pub d: f64,
26    pub dt: f64,
27}
28
29impl IzhikevichRk4 {
30    pub fn new(dt: f64) -> Self {
31        let c = -65.0;
32        let b = 0.2;
33        Self {
34            v: c,
35            u: b * c,
36            a: 0.02,
37            b,
38            c,
39            d: 8.0,
40            dt,
41        }
42    }
43
44    fn rhs(&self, v: f64, u: f64, current: f64) -> (f64, f64) {
45        let dv = 0.04 * v.powi(2) + 5.0 * v + 140.0 - u + current;
46        let du = self.a * (self.b * v - u);
47        (dv, du)
48    }
49
50    pub fn step(&mut self, current: f64) -> i32 {
51        let (k1_v, k1_u) = self.rhs(self.v, self.u, current);
52        let (k2_v, k2_u) = self.rhs(
53            self.v + 0.5 * self.dt * k1_v,
54            self.u + 0.5 * self.dt * k1_u,
55            current,
56        );
57        let (k3_v, k3_u) = self.rhs(
58            self.v + 0.5 * self.dt * k2_v,
59            self.u + 0.5 * self.dt * k2_u,
60            current,
61        );
62        let (k4_v, k4_u) = self.rhs(self.v + self.dt * k3_v, self.u + self.dt * k3_u, current);
63
64        self.v += (self.dt / 6.0) * (k1_v + 2.0 * k2_v + 2.0 * k3_v + k4_v);
65        self.u += (self.dt / 6.0) * (k1_u + 2.0 * k2_u + 2.0 * k3_u + k4_u);
66
67        if self.v >= IZH_SPIKE_THRESHOLD {
68            self.v = self.c;
69            self.u += self.d;
70            1
71        } else {
72            0
73        }
74    }
75}
76
77/// Izhikevich 2007 biophysical parameterisation (NeuroML `izhikevich2007Cell`):
78/// `C dv/dt = k (v - vr)(v - vt) - u + I`, `du/dt = a (b (v - vr) - u)`, with a
79/// `v >= vpeak -> v = c, u += d` reset. RK4 over the coupled ODE. The right-hand
80/// side is exact arithmetic (products, a sum, a division — no transcendental
81/// functions), so `simulate` matches the Python reference bit-for-bit.
82#[derive(Clone, Debug)]
83pub struct Izhikevich2007Rk4 {
84    pub v: f64,
85    pub u: f64,
86    pub cap: f64,
87    pub k: f64,
88    pub vr: f64,
89    pub vt: f64,
90    pub vpeak: f64,
91    pub a: f64,
92    pub b: f64,
93    pub c: f64,
94    pub d: f64,
95    pub dt: f64,
96}
97
98impl Izhikevich2007Rk4 {
99    pub fn new() -> Self {
100        Self {
101            v: -60.0,
102            u: 0.0,
103            cap: 100.0,
104            k: 0.7,
105            vr: -60.0,
106            vt: -40.0,
107            vpeak: 35.0,
108            a: 0.03,
109            b: -2.0,
110            c: -50.0,
111            d: 100.0,
112            dt: 0.1,
113        }
114    }
115
116    fn is_valid(&self) -> bool {
117        [
118            self.v, self.u, self.cap, self.k, self.vr, self.vt, self.vpeak, self.a, self.b, self.c,
119            self.d, self.dt,
120        ]
121        .iter()
122        .all(|value| value.is_finite())
123            && self.cap > 0.0
124            && self.dt > 0.0
125    }
126
127    fn rhs(&self, v: f64, u: f64, current: f64) -> (f64, f64) {
128        let dv = (self.k * (v - self.vr) * (v - self.vt) - u + current) / self.cap;
129        let du = self.a * (self.b * (v - self.vr) - u);
130        (dv, du)
131    }
132
133    pub fn try_step(&mut self, current: f64) -> Result<i32, &'static str> {
134        if !self.is_valid() || !current.is_finite() {
135            return Err("invalid Izhikevich 2007 runtime state or current");
136        }
137        let (k1v, k1u) = self.rhs(self.v, self.u, current);
138        let (k2v, k2u) = self.rhs(
139            self.v + 0.5 * self.dt * k1v,
140            self.u + 0.5 * self.dt * k1u,
141            current,
142        );
143        let (k3v, k3u) = self.rhs(
144            self.v + 0.5 * self.dt * k2v,
145            self.u + 0.5 * self.dt * k2u,
146            current,
147        );
148        let (k4v, k4u) = self.rhs(self.v + self.dt * k3v, self.u + self.dt * k3u, current);
149        let dt6 = self.dt / 6.0;
150        let mut v_next = self.v + dt6 * (k1v + 2.0 * k2v + 2.0 * k3v + k4v);
151        let mut u_next = self.u + dt6 * (k1u + 2.0 * k2u + 2.0 * k3u + k4u);
152        let event = if v_next >= self.vpeak {
153            v_next = self.c;
154            u_next += self.d;
155            1
156        } else {
157            0
158        };
159        if !v_next.is_finite() || !u_next.is_finite() {
160            return Err("Izhikevich 2007 candidate state became non-finite");
161        }
162        self.v = v_next;
163        self.u = u_next;
164        Ok(event)
165    }
166
167    pub fn step(&mut self, current: f64) -> i32 {
168        self.try_step(current).unwrap_or(0)
169    }
170
171    /// Run `n_steps` RK4 updates under a constant input, returning the `v` trace
172    /// (already reset to `c` on spiking steps) and the spike count. Reuses
173    /// `step`, so the trace is bit-identical to the per-step path and to the
174    /// Python reference (the right-hand side is exact arithmetic). The final
175    /// state is left in `self.v` / `self.u`.
176    pub fn simulate(
177        &mut self,
178        n_steps: usize,
179        current: f64,
180    ) -> Result<(Vec<f64>, i64), &'static str> {
181        let mut trace = Vec::with_capacity(n_steps);
182        let mut spikes: i64 = 0;
183        for _ in 0..n_steps {
184            let spiked = self.try_step(current)?;
185            trace.push(self.v);
186            if spiked == 1 {
187                spikes += 1;
188            }
189        }
190        Ok((trace, spikes))
191    }
192
193    pub fn try_reset(&mut self) -> Result<(), &'static str> {
194        if !self.is_valid() {
195            return Err("invalid Izhikevich 2007 parameters for reset");
196        }
197        let v_next = self.vr;
198        let u_next = self.b * (v_next - self.vr);
199        if !v_next.is_finite() || !u_next.is_finite() {
200            return Err("Izhikevich 2007 reset state became non-finite");
201        }
202        self.v = v_next;
203        self.u = u_next;
204        Ok(())
205    }
206
207    pub fn reset(&mut self) {
208        let _ = self.try_reset();
209    }
210}
211
212impl Default for Izhikevich2007Rk4 {
213    fn default() -> Self {
214        Self::new()
215    }
216}
217
218#[derive(Clone, Debug)]
219pub struct AdExRk4 {
220    pub v: f64,
221    pub w: f64,
222    pub v_rest: f64,
223    pub v_reset: f64,
224    pub v_threshold: f64,
225    pub v_rh: f64,
226    pub delta_t: f64,
227    pub tau: f64,
228    pub tau_w: f64,
229    pub a: f64,
230    pub b: f64,
231    pub c_m: f64,
232    pub dt: f64,
233}
234
235impl AdExRk4 {
236    pub fn new(dt: f64) -> Self {
237        Self {
238            v: -65.0,
239            w: 0.0,
240            v_rest: -65.0,
241            v_reset: -68.0,
242            v_threshold: -50.0,
243            v_rh: -55.0,
244            delta_t: 2.0,
245            tau: 20.0,
246            tau_w: 100.0,
247            a: 0.5,
248            b: 7.0,
249            c_m: 200.0,
250            dt,
251        }
252    }
253
254    fn rhs(&self, v: f64, w: f64, current: f64) -> (f64, f64) {
255        let exp_arg = ((v - self.v_rh) / self.delta_t).clamp(-20.0, 20.0);
256        let exp_term = self.delta_t * exp_arg.exp();
257        let dv = (-(v - self.v_rest) + exp_term) / self.tau + (-w + current) / self.c_m;
258        let dw = (self.a * (v - self.v_rest) - w) / self.tau_w;
259        (dv, dw)
260    }
261
262    pub fn step(&mut self, current: f64) -> i32 {
263        let (k1_v, k1_w) = self.rhs(self.v, self.w, current);
264        let (k2_v, k2_w) = self.rhs(
265            self.v + 0.5 * self.dt * k1_v,
266            self.w + 0.5 * self.dt * k1_w,
267            current,
268        );
269        let (k3_v, k3_w) = self.rhs(
270            self.v + 0.5 * self.dt * k2_v,
271            self.w + 0.5 * self.dt * k2_w,
272            current,
273        );
274        let (k4_v, k4_w) = self.rhs(self.v + self.dt * k3_v, self.w + self.dt * k3_w, current);
275
276        self.v += (self.dt / 6.0) * (k1_v + 2.0 * k2_v + 2.0 * k3_v + k4_v);
277        self.w += (self.dt / 6.0) * (k1_w + 2.0 * k2_w + 2.0 * k3_w + k4_w);
278
279        if self.v >= self.v_threshold {
280            self.v = self.v_reset;
281            self.w += self.b;
282            1
283        } else {
284            0
285        }
286    }
287}
288
289#[derive(Clone, Debug)]
290pub struct HodgkinHuxleyRk4 {
291    pub v: f64,
292    pub m: f64,
293    pub h: f64,
294    pub n: f64,
295    pub c_m: f64,
296    pub g_na: f64,
297    pub g_k: f64,
298    pub g_l: f64,
299    pub e_na: f64,
300    pub e_k: f64,
301    pub e_l: f64,
302    pub dt: f64,
303    pub v_threshold: f64,
304}
305
306impl HodgkinHuxleyRk4 {
307    pub fn new(dt: f64) -> Self {
308        Self {
309            v: -65.0,
310            m: 0.05,
311            h: 0.6,
312            n: 0.32,
313            c_m: 1.0,
314            g_na: 120.0,
315            g_k: 36.0,
316            g_l: 0.3,
317            e_na: 50.0,
318            e_k: -77.0,
319            e_l: -54.4,
320            dt,
321            v_threshold: 0.0,
322        }
323    }
324
325    fn alpha_m(v: f64) -> f64 {
326        let d = v + 40.0;
327        if d.abs() < 1e-7 {
328            1.0
329        } else {
330            0.1 * d / (1.0 - (-d / 10.0).exp())
331        }
332    }
333
334    fn beta_m(v: f64) -> f64 {
335        4.0 * (-(v + 65.0) / 18.0).exp()
336    }
337
338    fn alpha_h(v: f64) -> f64 {
339        0.07 * (-(v + 65.0) / 20.0).exp()
340    }
341
342    fn beta_h(v: f64) -> f64 {
343        1.0 / (1.0 + (-(v + 35.0) / 10.0).exp())
344    }
345
346    fn alpha_n(v: f64) -> f64 {
347        let d = v + 55.0;
348        if d.abs() < 1e-7 {
349            0.1
350        } else {
351            0.01 * d / (1.0 - (-d / 10.0).exp())
352        }
353    }
354
355    fn beta_n(v: f64) -> f64 {
356        0.125 * (-(v + 65.0) / 80.0).exp()
357    }
358
359    fn rhs(&self, state: [f64; 4], current: f64) -> [f64; 4] {
360        let [v, m, h, n] = state;
361        let am = Self::alpha_m(v);
362        let bm = Self::beta_m(v);
363        let ah = Self::alpha_h(v);
364        let bh = Self::beta_h(v);
365        let an = Self::alpha_n(v);
366        let bn = Self::beta_n(v);
367
368        let dm = am * (1.0 - m) - bm * m;
369        let dh = ah * (1.0 - h) - bh * h;
370        let dn = an * (1.0 - n) - bn * n;
371        let i_na = self.g_na * m.powi(3) * h * (v - self.e_na);
372        let i_k = self.g_k * n.powi(4) * (v - self.e_k);
373        let i_l = self.g_l * (v - self.e_l);
374        let dv = (-i_na - i_k - i_l + current) / self.c_m;
375        [dv, dm, dh, dn]
376    }
377
378    pub fn step(&mut self, current: f64) -> i32 {
379        let v_prev = self.v;
380        let mut state = [self.v, self.m, self.h, self.n];
381        let substeps = (1.0 / self.dt).round() as usize;
382        for _ in 0..substeps {
383            let k1 = self.rhs(state, current);
384            let k2 = self.rhs(add_scaled(state, k1, 0.5 * self.dt), current);
385            let k3 = self.rhs(add_scaled(state, k2, 0.5 * self.dt), current);
386            let k4 = self.rhs(add_scaled(state, k3, self.dt), current);
387            for idx in 0..4 {
388                state[idx] += (self.dt / 6.0) * (k1[idx] + 2.0 * k2[idx] + 2.0 * k3[idx] + k4[idx]);
389            }
390        }
391        self.v = state[0];
392        self.m = state[1];
393        self.h = state[2];
394        self.n = state[3];
395
396        if self.v >= self.v_threshold && v_prev < self.v_threshold {
397            1
398        } else {
399            0
400        }
401    }
402}
403
404fn add_scaled(state: [f64; 4], deriv: [f64; 4], scale: f64) -> [f64; 4] {
405    [
406        state[0] + scale * deriv[0],
407        state[1] + scale * deriv[1],
408        state[2] + scale * deriv[2],
409        state[3] + scale * deriv[3],
410    ]
411}
412
413#[pyfunction]
414#[pyo3(signature = (model_name, current_trace, dt=None))]
415pub fn py_rk4_neuron_simulate<'py>(
416    py: Python<'py>,
417    model_name: &str,
418    current_trace: PyReadonlyArray1<'py, f64>,
419    dt: Option<f64>,
420) -> PyResult<Py<PyAny>> {
421    let currents = current_trace.as_slice()?;
422    match normalise_model_name(model_name).as_str() {
423        "izhikevich" | "scizhikevichneuron" | "izhikevichneuron" => {
424            let dt = validate_trace_dt(currents, dt.unwrap_or(1.0))?;
425            simulate_izhikevich(py, currents, dt)
426        }
427        "hodgkinhuxley" | "hodgkinhuxleyneuron" => {
428            let dt = validate_trace_dt(currents, dt.unwrap_or(0.01))?;
429            simulate_hodgkin_huxley(py, currents, dt)
430        }
431        "adex" | "adexneuron" => {
432            let dt = validate_trace_dt(currents, dt.unwrap_or(0.1))?;
433            simulate_adex(py, currents, dt)
434        }
435        _ => Err(PyValueError::new_err(format!(
436            "unsupported RK4 neuron model {model_name:?}"
437        ))),
438    }
439}
440
441fn validate_trace_dt(currents: &[f64], dt: f64) -> PyResult<f64> {
442    if !dt.is_finite() || dt <= 0.0 {
443        return Err(PyValueError::new_err("dt must be a positive finite scalar"));
444    }
445    if currents.is_empty() {
446        return Err(PyValueError::new_err("current_trace must be non-empty"));
447    }
448    if currents.iter().any(|current| !current.is_finite()) {
449        return Err(PyValueError::new_err(
450            "current_trace must contain only finite values",
451        ));
452    }
453    Ok(dt)
454}
455
456fn normalise_model_name(name: &str) -> String {
457    name.chars()
458        .filter(|ch| ch.is_ascii_alphanumeric())
459        .flat_map(char::to_lowercase)
460        .collect()
461}
462
463fn simulate_izhikevich<'py>(py: Python<'py>, currents: &[f64], dt: f64) -> PyResult<Py<PyAny>> {
464    let mut neuron = IzhikevichRk4::new(dt);
465    let mut v = Vec::with_capacity(currents.len());
466    let mut u = Vec::with_capacity(currents.len());
467    let mut spikes = Vec::new();
468    for (idx, &current) in currents.iter().enumerate() {
469        if neuron.step(current) != 0 {
470            spikes.push(idx as u64);
471        }
472        v.push(neuron.v);
473        u.push(neuron.u);
474    }
475    let d = PyDict::new(py);
476    d.set_item("v", v.into_pyarray(py))?;
477    d.set_item("u", u.into_pyarray(py))?;
478    d.set_item("spikes", spikes.into_pyarray(py))?;
479    d.set_item("n_steps", currents.len())?;
480    Ok(d.into_any().unbind())
481}
482
483fn simulate_adex<'py>(py: Python<'py>, currents: &[f64], dt: f64) -> PyResult<Py<PyAny>> {
484    let mut neuron = AdExRk4::new(dt);
485    let mut v = Vec::with_capacity(currents.len());
486    let mut w = Vec::with_capacity(currents.len());
487    let mut spikes = Vec::new();
488    for (idx, &current) in currents.iter().enumerate() {
489        if neuron.step(current) != 0 {
490            spikes.push(idx as u64);
491        }
492        v.push(neuron.v);
493        w.push(neuron.w);
494    }
495    let d = PyDict::new(py);
496    d.set_item("v", v.into_pyarray(py))?;
497    d.set_item("w", w.into_pyarray(py))?;
498    d.set_item("spikes", spikes.into_pyarray(py))?;
499    d.set_item("n_steps", currents.len())?;
500    Ok(d.into_any().unbind())
501}
502
503fn simulate_hodgkin_huxley<'py>(py: Python<'py>, currents: &[f64], dt: f64) -> PyResult<Py<PyAny>> {
504    let mut neuron = HodgkinHuxleyRk4::new(dt);
505    let mut v = Vec::with_capacity(currents.len());
506    let mut m = Vec::with_capacity(currents.len());
507    let mut h = Vec::with_capacity(currents.len());
508    let mut n = Vec::with_capacity(currents.len());
509    let mut spikes = Vec::new();
510    for (idx, &current) in currents.iter().enumerate() {
511        if neuron.step(current) != 0 {
512            spikes.push(idx as u64);
513        }
514        v.push(neuron.v);
515        m.push(neuron.m);
516        h.push(neuron.h);
517        n.push(neuron.n);
518    }
519    let d = PyDict::new(py);
520    d.set_item("v", v.into_pyarray(py))?;
521    d.set_item("m", m.into_pyarray(py))?;
522    d.set_item("h", h.into_pyarray(py))?;
523    d.set_item("n", n.into_pyarray(py))?;
524    d.set_item("spikes", spikes.into_pyarray(py))?;
525    d.set_item("n_steps", currents.len())?;
526    Ok(d.into_any().unbind())
527}
528
529#[cfg(test)]
530mod tests {
531    use super::*;
532
533    #[test]
534    fn izhikevich_rk4_is_deterministic_and_spikes() {
535        let mut a = IzhikevichRk4::new(1.0);
536        let mut b = IzhikevichRk4::new(1.0);
537        let mut spikes = 0;
538        for _ in 0..100 {
539            spikes += a.step(10.0);
540            b.step(10.0);
541        }
542        assert!(spikes > 0);
543        assert_eq!(a.v, b.v);
544        assert_eq!(a.u, b.u);
545    }
546
547    #[test]
548    fn izhikevich2007_checked_trace_and_reset_are_fail_closed() {
549        let mut neuron = Izhikevich2007Rk4::new();
550        let (trace, events) = neuron.simulate(2_000, 100.0).unwrap();
551        assert_eq!(trace.len(), 2_000);
552        assert_eq!(events, 3);
553        assert_eq!(trace.last().copied(), Some(neuron.v));
554
555        let before = (neuron.v, neuron.u);
556        assert!(neuron.try_step(f64::NAN).is_err());
557        assert_eq!((neuron.v, neuron.u), before);
558
559        neuron.vr = f64::NAN;
560        assert!(neuron.try_reset().is_err());
561        assert_eq!((neuron.v, neuron.u), before);
562    }
563
564    #[test]
565    fn adex_rk4_remains_finite_under_sustained_current() {
566        let mut neuron = AdExRk4::new(0.1);
567        let mut spikes = 0;
568        for _ in 0..3000 {
569            spikes += neuron.step(500.0);
570        }
571        assert!(spikes > 0);
572        assert!(neuron.v.is_finite());
573        assert!(neuron.w.is_finite());
574    }
575
576    #[test]
577    fn hodgkin_huxley_rk4_keeps_gates_bounded() {
578        let mut neuron = HodgkinHuxleyRk4::new(0.01);
579        let mut spikes = 0;
580        for _ in 0..1000 {
581            spikes += neuron.step(10.0);
582        }
583        assert!(spikes > 0);
584        assert!(neuron.v.is_finite());
585        assert!((0.0..=1.0).contains(&neuron.m));
586        assert!((0.0..=1.0).contains(&neuron.h));
587        assert!((0.0..=1.0).contains(&neuron.n));
588    }
589
590    #[test]
591    fn model_name_normalisation_accepts_common_aliases() {
592        assert_eq!(
593            normalise_model_name("Hodgkin-HuxleyNeuron"),
594            "hodgkinhuxleyneuron"
595        );
596        assert_eq!(normalise_model_name("AdEx"), "adex");
597    }
598}