1#[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 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#[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 (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 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 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}