1use 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#[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 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, ¤t) 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, ¤t) 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, ¤t) 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}