Skip to main content

sc_neurocore_engine/
wilson_cowan.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 — Rust N-step simulator for the Wilson-Cowan 1972 E/I rate model
8
9//! Batch parity with `WilsonCowanUnit.step` in
10//! `src/sc_neurocore/neurons/models/wilson_cowan.py` (Wilson & Cowan
11//! 1972, Biophys. J. 12:1–24).
12//!
13//! Per step:
14//!   dE/dt = (−E + sigmoid(w_ee · E − w_ei · I + ext)) / τ_e
15//!   dI/dt = (−I + sigmoid(w_ie · E − w_ii · I)) / τ_i
16//!   (E, I) advance through one fixed-step RK4 update.
17//!
18//! where `sigmoid(x) = logistic(a·(x − θ)) − logistic(−a·θ)`.
19//!
20//! The model is deterministic (no noise), so parity needs no pre-drawn RNG
21//! buffer. Transcendental library implementations are compared under the
22//! public bounded floating-point trajectory contract.
23
24#[inline]
25fn logistic(z: f64) -> f64 {
26    if z >= 0.0 {
27        1.0 / (1.0 + (-z).exp())
28    } else {
29        let exp_z = z.exp();
30        exp_z / (1.0 + exp_z)
31    }
32}
33
34#[inline]
35fn sigmoid(a: f64, theta: f64, x: f64) -> f64 {
36    // Published Wilson-Cowan 1972 two-term form:
37    //   S(x) = 1/(1+exp(-a(x-θ))) − 1/(1+exp(aθ))
38    // Range is [-β, 1-β] where β = 1/(1+exp(aθ)).
39    logistic(a * (x - theta)) - logistic(-a * theta)
40}
41
42#[inline]
43fn finite_rate(value: f64, a: f64, theta: f64) -> bool {
44    let baseline = logistic(-a * theta);
45    value.is_finite() && value >= -baseline && value <= 1.0
46}
47
48#[expect(
49    clippy::too_many_arguments,
50    reason = "native parity surface passes the complete scientific configuration"
51)]
52fn valid_configuration(
53    e: f64,
54    i: f64,
55    w_ee: f64,
56    w_ei: f64,
57    w_ie: f64,
58    w_ii: f64,
59    tau_e: f64,
60    tau_i: f64,
61    a: f64,
62    theta: f64,
63    dt: f64,
64) -> bool {
65    [w_ee, w_ei, w_ie, w_ii]
66        .into_iter()
67        .all(|value| value.is_finite() && value >= 0.0)
68        && tau_e.is_finite()
69        && tau_e > 0.0
70        && tau_i.is_finite()
71        && tau_i > 0.0
72        && a.is_finite()
73        && a > 0.0
74        && theta.is_finite()
75        && dt.is_finite()
76        && dt > 0.0
77        && finite_rate(e, a, theta)
78        && finite_rate(i, a, theta)
79}
80
81#[inline]
82fn derivatives(
83    e: f64,
84    i: f64,
85    ext: f64,
86    params: (f64, f64, f64, f64, f64, f64, f64, f64),
87) -> (f64, f64) {
88    let (w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta) = params;
89    let s_e = sigmoid(a, theta, w_ee * e - w_ei * i + ext);
90    let s_i = sigmoid(a, theta, w_ie * e - w_ii * i);
91    ((-e + s_e) / tau_e, (-i + s_i) / tau_i)
92}
93
94/// Simulate `ext_input.len()` Wilson-Cowan iterations, writing per-step
95/// `E` and `I` traces into caller-allocated buffers. Returns final
96/// `(E, I)` for convenience.
97#[expect(
98    clippy::too_many_arguments,
99    reason = "Python extension parity surface passes canonical scalar parameters"
100)]
101pub fn simulate(
102    mut e: f64,
103    mut i: f64,
104    w_ee: f64,
105    w_ei: f64,
106    w_ie: f64,
107    w_ii: f64,
108    tau_e: f64,
109    tau_i: f64,
110    a: f64,
111    theta: f64,
112    dt: f64,
113    ext_input: &[f64],
114    e_out: &mut [f64],
115    i_out: &mut [f64],
116) -> Result<(f64, f64), &'static str> {
117    let n = ext_input.len();
118    if e_out.len() != n {
119        return Err("e_out length mismatch");
120    }
121    if i_out.len() != n {
122        return Err("i_out length mismatch");
123    }
124    if !valid_configuration(e, i, w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta, dt) {
125        return Err("invalid Wilson-Cowan numerical configuration");
126    }
127    if !ext_input.iter().all(|value| value.is_finite()) {
128        return Err("external input must be finite");
129    }
130    let params = (w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta);
131    let mut next_e_out = Vec::with_capacity(n);
132    let mut next_i_out = Vec::with_capacity(n);
133
134    for &ext in ext_input {
135        let (k1_e, k1_i) = derivatives(e, i, ext, params);
136        let (k2_e, k2_i) = derivatives(e + 0.5 * dt * k1_e, i + 0.5 * dt * k1_i, ext, params);
137        let (k3_e, k3_i) = derivatives(e + 0.5 * dt * k2_e, i + 0.5 * dt * k2_i, ext, params);
138        let (k4_e, k4_i) = derivatives(e + dt * k3_e, i + dt * k3_i, ext, params);
139        if ![k1_e, k1_i, k2_e, k2_i, k3_e, k3_i, k4_e, k4_i]
140            .into_iter()
141            .all(f64::is_finite)
142        {
143            return Err("invalid Wilson-Cowan derivative");
144        }
145        let next_e = e + dt * (k1_e + 2.0 * k2_e + 2.0 * k3_e + k4_e) / 6.0;
146        let next_i = i + dt * (k1_i + 2.0 * k2_i + 2.0 * k3_i + k4_i) / 6.0;
147        if !finite_rate(next_e, a, theta) || !finite_rate(next_i, a, theta) {
148            return Err("invalid Wilson-Cowan candidate state");
149        }
150        e = next_e;
151        i = next_i;
152        next_e_out.push(e);
153        next_i_out.push(i);
154    }
155    e_out.copy_from_slice(&next_e_out);
156    i_out.copy_from_slice(&next_i_out);
157    Ok((e, i))
158}
159
160#[cfg(test)]
161mod tests {
162    use super::*;
163
164    fn defaults() -> (f64, f64, f64, f64, f64, f64, f64, f64, f64) {
165        // w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta, dt
166        (10.0, 6.0, 10.0, 1.0, 1.0, 2.0, 1.2, 4.0, 0.1)
167    }
168
169    #[test]
170    fn sigmoid_monotone_increasing() {
171        let (_, _, _, _, _, _, a, theta, _) = defaults();
172        let lo = sigmoid(a, theta, 0.0);
173        let mid = sigmoid(a, theta, 4.0);
174        let hi = sigmoid(a, theta, 10.0);
175        assert!(lo < mid && mid < hi);
176    }
177
178    #[test]
179    fn sigmoid_at_zero_is_zero() {
180        // Two-term form zeroes the baseline so S(0) = 0 exactly.
181        let (_, _, _, _, _, _, a, theta, _) = defaults();
182        assert!(sigmoid(a, theta, 0.0).abs() < 1e-12);
183    }
184
185    #[test]
186    fn sigmoid_at_theta_equals_half_minus_baseline() {
187        let (_, _, _, _, _, _, a, theta, _) = defaults();
188        let baseline = 1.0 / (1.0 + (a * theta).exp());
189        let r = sigmoid(a, theta, theta);
190        assert!((r - (0.5 - baseline)).abs() < 1e-12);
191    }
192
193    #[test]
194    fn sigmoid_asymptotes_respect_baseline() {
195        // As x → +∞, S(x) → 1 − baseline. As x → −∞, S(x) → −baseline.
196        let (_, _, _, _, _, _, a, theta, _) = defaults();
197        let baseline = 1.0 / (1.0 + (a * theta).exp());
198        assert!((sigmoid(a, theta, 1e6) - (1.0 - baseline)).abs() < 1e-50);
199        assert!((sigmoid(a, theta, -1e6) - (-baseline)).abs() < 1e-50);
200    }
201
202    #[test]
203    fn quiescent_converges() {
204        let (w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta, dt) = defaults();
205        let n = 20_000;
206        let ext = vec![0.0_f64; n];
207        let mut e_out = vec![0.0_f64; n];
208        let mut i_out = vec![0.0_f64; n];
209        let (e_f, i_f) = simulate(
210            0.1, 0.05, w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta, dt, &ext, &mut e_out,
211            &mut i_out,
212        )
213        .unwrap();
214        assert!(e_f.is_finite() && i_f.is_finite());
215        assert!(e_f < 0.2, "quiescent E must stay low, got {e_f}");
216        assert!(i_f < 0.2, "quiescent I must stay low, got {i_f}");
217    }
218
219    #[test]
220    fn high_drive_elevates_activity() {
221        let (w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta, dt) = defaults();
222        let n = 10_000;
223        let ext = vec![10.0_f64; n];
224        let mut e_out = vec![0.0_f64; n];
225        let mut i_out = vec![0.0_f64; n];
226        let (e_f, _) = simulate(
227            0.1, 0.05, w_ee, w_ei, w_ie, w_ii, tau_e, tau_i, a, theta, dt, &ext, &mut e_out,
228            &mut i_out,
229        )
230        .unwrap();
231        assert!(e_f > 0.3, "high external drive must elevate E, got {e_f}");
232    }
233
234    #[test]
235    fn rk4_step_matches_reference_and_separates_from_euler() {
236        let mut e_out = vec![0.0_f64; 1];
237        let mut i_out = vec![0.0_f64; 1];
238        let ext = vec![3.0_f64; 1];
239        simulate(
240            0.24, 0.11, 10.0, 6.0, 10.0, 1.0, 1.0, 2.0, 1.2, 4.0, 0.35, &ext, &mut e_out,
241            &mut i_out,
242        )
243        .unwrap();
244        let euler_e = 0.40111014473980233_f64;
245        let euler_i = 0.10924537850891547_f64;
246        assert!((e_out[0] - 0.42143718680097664_f64).abs() < 1e-15);
247        assert!((i_out[0] - 0.13798020053932203_f64).abs() < 1e-15);
248        assert!((e_out[0] - euler_e).abs() > 1e-2);
249        assert!((i_out[0] - euler_i).abs() > 1e-2);
250    }
251
252    #[test]
253    fn output_trace_shape_matches_input() {
254        let n = 64;
255        let ext = vec![1.0_f64; n];
256        let mut e_out = vec![f64::NAN; n];
257        let mut i_out = vec![f64::NAN; n];
258        simulate(
259            0.1, 0.05, 10.0, 6.0, 10.0, 1.0, 1.0, 2.0, 1.2, 4.0, 0.1, &ext, &mut e_out, &mut i_out,
260        )
261        .unwrap();
262        assert!(e_out.iter().all(|v| v.is_finite()));
263        assert!(i_out.iter().all(|v| v.is_finite()));
264    }
265
266    #[test]
267    fn mismatched_e_out_is_rejected() {
268        let n = 10;
269        let ext = vec![0.0_f64; n];
270        let mut e_out = vec![0.0_f64; n + 1];
271        let mut i_out = vec![0.0_f64; n];
272        let error = simulate(
273            0.1, 0.05, 10.0, 6.0, 10.0, 1.0, 1.0, 2.0, 1.2, 4.0, 0.1, &ext, &mut e_out, &mut i_out,
274        )
275        .unwrap_err();
276        assert_eq!(error, "e_out length mismatch");
277    }
278
279    #[test]
280    fn invalid_contract_preserves_caller_buffers() {
281        let ext = vec![1.0_f64, f64::NAN, 1.0];
282        let mut e_out = vec![-999.0_f64; ext.len()];
283        let mut i_out = vec![-999.0_f64; ext.len()];
284        let result = simulate(
285            0.1, 0.05, 10.0, 6.0, 10.0, 1.0, 1.0, 2.0, 1.2, 4.0, 0.1, &ext, &mut e_out, &mut i_out,
286        );
287        assert!(result.is_err());
288        assert_eq!(e_out, vec![-999.0; ext.len()]);
289        assert_eq!(i_out, vec![-999.0; ext.len()]);
290    }
291}