1use crate::error::{NeuralDynamicsError, Result};
6use hodgkin_huxley::{HodgkinHuxleyNeuron, neuron_types::NeuronConfig};
7use rand::Rng;
8use rand_distr::{Distribution, Normal};
9use rayon::prelude::*;
10use serde::{Deserialize, Serialize};
11use std::sync::{Arc, Mutex};
12
13#[derive(Clone)]
15pub struct NeuralPopulation {
16 pub id: String,
18 pub size: usize,
20 neurons: Vec<HodgkinHuxleyNeuron>,
22 external_currents: Vec<f64>,
24 synaptic_currents: Vec<f64>,
26 spike_times: Vec<Vec<f64>>,
28 spike_threshold: f64,
30 last_above_threshold: Vec<bool>,
32 voltage_trace: Option<Vec<Vec<f64>>>,
34 recorded_times: Vec<f64>,
36}
37
38impl NeuralPopulation {
39 pub fn new_homogeneous(id: impl Into<String>, size: usize, config: NeuronConfig) -> Result<Self> {
47 if size == 0 {
48 return Err(NeuralDynamicsError::EmptyPopulation);
49 }
50
51 let neurons: Result<Vec<HodgkinHuxleyNeuron>> = (0..size)
52 .map(|_| HodgkinHuxleyNeuron::new(config.clone()).map_err(|e| e.into()))
53 .collect();
54 let mut neurons: Vec<HodgkinHuxleyNeuron> = neurons?;
55
56 for neuron in neurons.iter_mut() {
58 neuron.initialize_rest();
59 }
60
61 Ok(Self {
62 id: id.into(),
63 size,
64 neurons,
65 external_currents: vec![0.0; size],
66 synaptic_currents: vec![0.0; size],
67 spike_times: vec![Vec::new(); size],
68 spike_threshold: -20.0,
69 last_above_threshold: vec![false; size],
70 voltage_trace: None,
71 recorded_times: Vec::new(),
72 })
73 }
74
75 pub fn new_heterogeneous<R: Rng>(
85 id: impl Into<String>,
86 size: usize,
87 base_config: NeuronConfig,
88 variability: f64,
89 rng: &mut R,
90 ) -> Result<Self> {
91 if size == 0 {
92 return Err(NeuralDynamicsError::EmptyPopulation);
93 }
94
95 if variability < 0.0 {
96 return Err(NeuralDynamicsError::InvalidParameter {
97 parameter: "variability".to_string(),
98 value: variability,
99 reason: "must be non-negative".to_string(),
100 });
101 }
102
103 let mut neurons = Vec::with_capacity(size);
104
105 for _ in 0..size {
106 let mut config = base_config.clone();
107
108 if variability > 0.0 {
110 let g_na_dist = Normal::new(config.na_channel.g_max, config.na_channel.g_max * variability)
111 .map_err(|e| NeuralDynamicsError::InvalidParameter {
112 parameter: "g_na variability".to_string(),
113 value: variability,
114 reason: e.to_string(),
115 })?;
116 config.na_channel.g_max = g_na_dist.sample(rng).max(0.0);
117
118 let g_k_dist = Normal::new(config.k_channel.g_max, config.k_channel.g_max * variability)
119 .map_err(|e| NeuralDynamicsError::InvalidParameter {
120 parameter: "g_k variability".to_string(),
121 value: variability,
122 reason: e.to_string(),
123 })?;
124 config.k_channel.g_max = g_k_dist.sample(rng).max(0.0);
125 }
126
127 let mut neuron = HodgkinHuxleyNeuron::new(config)?;
128 neuron.initialize_rest();
129 neurons.push(neuron);
130 }
131
132 Ok(Self {
133 id: id.into(),
134 size,
135 neurons,
136 external_currents: vec![0.0; size],
137 synaptic_currents: vec![0.0; size],
138 spike_times: vec![Vec::new(); size],
139 spike_threshold: -20.0,
140 last_above_threshold: vec![false; size],
141 voltage_trace: None,
142 recorded_times: Vec::new(),
143 })
144 }
145
146 pub fn excitatory(id: impl Into<String>, size: usize) -> Result<Self> {
148 Self::new_homogeneous(id, size, NeuronConfig::regular_spiking())
149 }
150
151 pub fn inhibitory(id: impl Into<String>, size: usize) -> Result<Self> {
153 Self::new_homogeneous(id, size, NeuronConfig::fast_spiking())
154 }
155
156 pub fn enable_recording(&mut self) {
158 self.voltage_trace = Some(vec![Vec::new(); self.size]);
159 self.recorded_times.clear();
160 }
161
162 pub fn disable_recording(&mut self) {
164 self.voltage_trace = None;
165 self.recorded_times.clear();
166 }
167
168 pub fn set_external_current(&mut self, index: usize, current: f64) -> Result<()> {
170 if index >= self.size {
171 return Err(NeuralDynamicsError::InvalidNeuronIndex {
172 index,
173 max: self.size - 1,
174 });
175 }
176 self.external_currents[index] = current;
177 Ok(())
178 }
179
180 pub fn set_external_currents(&mut self, currents: &[f64]) -> Result<()> {
182 if currents.len() != self.size {
183 return Err(NeuralDynamicsError::SizeMismatch {
184 expected: self.size,
185 actual: currents.len(),
186 });
187 }
188 self.external_currents.copy_from_slice(currents);
189 Ok(())
190 }
191
192 pub fn add_external_current(&mut self, index: usize, current: f64) -> Result<()> {
194 if index >= self.size {
195 return Err(NeuralDynamicsError::InvalidNeuronIndex {
196 index,
197 max: self.size - 1,
198 });
199 }
200 self.external_currents[index] += current;
201 Ok(())
202 }
203
204 pub fn set_synaptic_current(&mut self, index: usize, current: f64) -> Result<()> {
206 if index >= self.size {
207 return Err(NeuralDynamicsError::InvalidNeuronIndex {
208 index,
209 max: self.size - 1,
210 });
211 }
212 self.synaptic_currents[index] = current;
213 Ok(())
214 }
215
216 pub fn reset_synaptic_currents(&mut self) {
218 self.synaptic_currents.fill(0.0);
219 }
220
221 pub fn get_voltage(&self, index: usize) -> Result<f64> {
223 if index >= self.size {
224 return Err(NeuralDynamicsError::InvalidNeuronIndex {
225 index,
226 max: self.size - 1,
227 });
228 }
229 Ok(self.neurons[index].voltage())
230 }
231
232 pub fn get_voltages(&self) -> Vec<f64> {
234 self.neurons.iter().map(|n| n.voltage()).collect()
235 }
236
237 pub fn get_spike_times(&self, index: usize) -> Result<&[f64]> {
239 if index >= self.size {
240 return Err(NeuralDynamicsError::InvalidNeuronIndex {
241 index,
242 max: self.size - 1,
243 });
244 }
245 Ok(&self.spike_times[index])
246 }
247
248 pub fn get_all_spike_times(&self) -> &[Vec<f64>] {
250 &self.spike_times
251 }
252
253 pub fn update(&mut self, dt: f64, current_time: f64) -> Result<()> {
255 for i in 0..self.size {
256 let total_current = self.external_currents[i] + self.synaptic_currents[i];
257 self.neurons[i].step(dt, total_current)?;
258
259 let v = self.neurons[i].voltage();
261 if !self.last_above_threshold[i] && v > self.spike_threshold {
262 self.spike_times[i].push(current_time);
263 self.last_above_threshold[i] = true;
264 } else if self.last_above_threshold[i] && v <= self.spike_threshold {
265 self.last_above_threshold[i] = false;
266 }
267 }
268
269 if let Some(ref mut traces) = self.voltage_trace {
271 for (i, neuron) in self.neurons.iter().enumerate() {
272 traces[i].push(neuron.voltage());
273 }
274 self.recorded_times.push(current_time);
275 }
276
277 Ok(())
278 }
279
280 pub fn update_parallel(&mut self, dt: f64, current_time: f64) -> Result<()> {
282 let neurons_mutex = Arc::new(Mutex::new(&mut self.neurons));
284 let results: Vec<Result<f64>> = (0..self.size)
285 .into_par_iter()
286 .map(|i| {
287 let total_current = self.external_currents[i] + self.synaptic_currents[i];
288 let mut neurons = neurons_mutex.lock().unwrap();
289 neurons[i].step(dt, total_current)?;
290 Ok(neurons[i].voltage())
291 })
292 .collect();
293
294 let voltages: Result<Vec<_>> = results.into_iter().collect();
296 let voltages = voltages?;
297
298 for (i, &v) in voltages.iter().enumerate() {
300 if !self.last_above_threshold[i] && v > self.spike_threshold {
301 self.spike_times[i].push(current_time);
302 self.last_above_threshold[i] = true;
303 } else if self.last_above_threshold[i] && v <= self.spike_threshold {
304 self.last_above_threshold[i] = false;
305 }
306 }
307
308 if let Some(ref mut traces) = self.voltage_trace {
310 for (i, &v) in voltages.iter().enumerate() {
311 traces[i].push(v);
312 }
313 self.recorded_times.push(current_time);
314 }
315
316 Ok(())
317 }
318
319 pub fn get_voltage_traces(&self) -> Option<(&[Vec<f64>], &[f64])> {
321 self.voltage_trace.as_ref().map(|traces| (traces.as_slice(), self.recorded_times.as_slice()))
322 }
323
324 pub fn clear_history(&mut self) {
326 self.spike_times.iter_mut().for_each(|v| v.clear());
327 if let Some(ref mut traces) = self.voltage_trace {
328 traces.iter_mut().for_each(|v| v.clear());
329 }
330 self.recorded_times.clear();
331 }
332
333 pub fn reset(&mut self) {
335 for neuron in self.neurons.iter_mut() {
336 neuron.initialize_rest();
337 }
338 self.external_currents.fill(0.0);
339 self.synaptic_currents.fill(0.0);
340 self.last_above_threshold.fill(false);
341 self.clear_history();
342 }
343
344 pub fn statistics(&self, time_window: Option<(f64, f64)>) -> PopulationStats {
346 let voltages = self.get_voltages();
347 let mean_voltage = voltages.iter().sum::<f64>() / self.size as f64;
348 let voltage_std = (voltages.iter()
349 .map(|v| (v - mean_voltage).powi(2))
350 .sum::<f64>() / self.size as f64)
351 .sqrt();
352
353 let (total_spikes, active_neurons) = if let Some((t_start, t_end)) = time_window {
355 let mut total = 0;
356 let mut active = 0;
357 for spike_train in &self.spike_times {
358 let count = spike_train.iter().filter(|&&t| t >= t_start && t < t_end).count();
359 if count > 0 {
360 active += 1;
361 total += count;
362 }
363 }
364 (total, active)
365 } else {
366 let total: usize = self.spike_times.iter().map(|v| v.len()).sum();
367 let active = self.spike_times.iter().filter(|v| !v.is_empty()).count();
368 (total, active)
369 };
370
371 PopulationStats {
372 size: self.size,
373 mean_voltage,
374 voltage_std,
375 total_spikes,
376 active_neurons,
377 firing_rate: 0.0, }
379 }
380
381 pub fn instantaneous_rate(&self, time: f64, window: f64) -> f64 {
383 let t_start = time - window / 2.0;
384 let t_end = time + window / 2.0;
385
386 let spike_count: usize = self.spike_times
387 .iter()
388 .map(|spikes| spikes.iter().filter(|&&t| t >= t_start && t < t_end).count())
389 .sum();
390
391 (spike_count as f64 / (self.size as f64 * window)) * 1000.0 }
393}
394
395#[derive(Debug, Clone, Serialize, Deserialize)]
397pub struct PopulationStats {
398 pub size: usize,
399 pub mean_voltage: f64,
400 pub voltage_std: f64,
401 pub total_spikes: usize,
402 pub active_neurons: usize,
403 pub firing_rate: f64,
404}
405
406#[cfg(test)]
407mod tests {
408 use super::*;
409 use approx::assert_relative_eq;
410
411 #[test]
412 fn test_create_homogeneous_population() {
413 let pop = NeuralPopulation::excitatory("E", 10).unwrap();
414 assert_eq!(pop.size, 10);
415 assert_eq!(pop.id, "E");
416 assert_eq!(pop.neurons.len(), 10);
417 }
418
419 #[test]
420 fn test_create_heterogeneous_population() {
421 let mut rng = rand::thread_rng();
422 let pop = NeuralPopulation::new_heterogeneous(
423 "E",
424 20,
425 NeuronConfig::regular_spiking(),
426 0.2,
427 &mut rng,
428 )
429 .unwrap();
430 assert_eq!(pop.size, 20);
431 }
432
433 #[test]
434 fn test_empty_population_fails() {
435 let result = NeuralPopulation::excitatory("E", 0);
436 assert!(result.is_err());
437 }
438
439 #[test]
440 fn test_set_external_current() {
441 let mut pop = NeuralPopulation::excitatory("E", 5).unwrap();
442 pop.set_external_current(2, 10.0).unwrap();
443 assert_eq!(pop.external_currents[2], 10.0);
444
445 assert!(pop.set_external_current(10, 5.0).is_err());
447 }
448
449 #[test]
450 fn test_population_update() {
451 let mut pop = NeuralPopulation::excitatory("E", 3).unwrap();
452 pop.set_external_current(0, 20.0).unwrap();
453
454 for i in 0..1000 {
455 pop.update(0.01, i as f64 * 0.01).unwrap();
456 }
457
458 assert!(!pop.spike_times[0].is_empty());
460 assert!(pop.spike_times[1].is_empty());
462 }
463
464 #[test]
465 fn test_voltage_recording() {
466 let mut pop = NeuralPopulation::excitatory("E", 2).unwrap();
467 pop.enable_recording();
468
469 for i in 0..10 {
470 pop.update(0.1, i as f64 * 0.1).unwrap();
471 }
472
473 let (traces, times) = pop.get_voltage_traces().unwrap();
474 assert_eq!(traces.len(), 2);
475 assert_eq!(times.len(), 10);
476 assert_eq!(traces[0].len(), 10);
477 }
478
479 #[test]
480 fn test_population_reset() {
481 let mut pop = NeuralPopulation::excitatory("E", 3).unwrap();
482 pop.set_external_current(0, 10.0).unwrap();
483 pop.update(0.01, 0.01).unwrap();
484
485 pop.reset();
486 assert_eq!(pop.external_currents[0], 0.0);
487 assert!(pop.spike_times[0].is_empty());
488 }
489
490 #[test]
491 fn test_population_statistics() {
492 let mut pop = NeuralPopulation::excitatory("E", 5).unwrap();
493 let stats = pop.statistics(None);
494 assert_eq!(stats.size, 5);
495 assert!(stats.mean_voltage < 0.0); }
497
498 #[test]
499 fn test_instantaneous_rate() {
500 let mut pop = NeuralPopulation::excitatory("E", 10).unwrap();
501
502 for i in 0..10 {
504 pop.set_external_current(i, 15.0).unwrap();
505 }
506
507 for i in 0..5000 {
508 pop.update(0.01, i as f64 * 0.01).unwrap();
509 }
510
511 let rate = pop.instantaneous_rate(25.0, 10.0);
512 assert!(rate > 0.0); }
514}