Skip to main content

neural_dynamics/
population.rs

1//! Neural population management.
2//!
3//! This module provides structures for managing groups of neurons with similar properties.
4
5use 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/// Neural population with homogeneous or heterogeneous neuron parameters.
14#[derive(Clone)]
15pub struct NeuralPopulation {
16    /// Population identifier
17    pub id: String,
18    /// Number of neurons in population
19    pub size: usize,
20    /// Neurons in this population
21    neurons: Vec<HodgkinHuxleyNeuron>,
22    /// External currents for each neuron (µA/cm²)
23    external_currents: Vec<f64>,
24    /// Synaptic currents from network (µA/cm²)
25    synaptic_currents: Vec<f64>,
26    /// Spike times for each neuron
27    spike_times: Vec<Vec<f64>>,
28    /// Spike threshold for detection (mV)
29    spike_threshold: f64,
30    /// Last spike detection state
31    last_above_threshold: Vec<bool>,
32    /// Voltage recording (optional)
33    voltage_trace: Option<Vec<Vec<f64>>>,
34    /// Time points for recordings
35    recorded_times: Vec<f64>,
36}
37
38impl NeuralPopulation {
39    /// Create a homogeneous population with identical neurons.
40    ///
41    /// # Arguments
42    ///
43    /// * `id` - Population identifier
44    /// * `size` - Number of neurons
45    /// * `config` - Neuron configuration (same for all neurons)
46    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        // Initialize all neurons at rest
57        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    /// Create a heterogeneous population with varied parameters.
76    ///
77    /// # Arguments
78    ///
79    /// * `id` - Population identifier
80    /// * `size` - Number of neurons
81    /// * `base_config` - Base neuron configuration
82    /// * `variability` - Parameter variability (0.0 = homogeneous, 0.2 = 20% std dev)
83    /// * `rng` - Random number generator
84    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            // Add variability to conductances
109            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    /// Create an excitatory population (regular spiking).
147    pub fn excitatory(id: impl Into<String>, size: usize) -> Result<Self> {
148        Self::new_homogeneous(id, size, NeuronConfig::regular_spiking())
149    }
150
151    /// Create an inhibitory population (fast spiking).
152    pub fn inhibitory(id: impl Into<String>, size: usize) -> Result<Self> {
153        Self::new_homogeneous(id, size, NeuronConfig::fast_spiking())
154    }
155
156    /// Enable voltage recording for all neurons.
157    pub fn enable_recording(&mut self) {
158        self.voltage_trace = Some(vec![Vec::new(); self.size]);
159        self.recorded_times.clear();
160    }
161
162    /// Disable voltage recording.
163    pub fn disable_recording(&mut self) {
164        self.voltage_trace = None;
165        self.recorded_times.clear();
166    }
167
168    /// Set external current for a specific neuron.
169    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    /// Set external current for all neurons.
181    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    /// Add to external current for a specific neuron.
193    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    /// Set synaptic current for a specific neuron (called by network).
205    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    /// Reset synaptic currents (typically called before network updates).
217    pub fn reset_synaptic_currents(&mut self) {
218        self.synaptic_currents.fill(0.0);
219    }
220
221    /// Get voltage of a specific neuron.
222    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    /// Get voltages of all neurons.
233    pub fn get_voltages(&self) -> Vec<f64> {
234        self.neurons.iter().map(|n| n.voltage()).collect()
235    }
236
237    /// Get spike times for a specific neuron.
238    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    /// Get all spike times.
249    pub fn get_all_spike_times(&self) -> &[Vec<f64>] {
250        &self.spike_times
251    }
252
253    /// Update population for one time step (sequential).
254    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            // Detect spikes
260            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        // Record voltages if enabled
270        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    /// Update population for one time step (parallel).
281    pub fn update_parallel(&mut self, dt: f64, current_time: f64) -> Result<()> {
282        // Parallel update of neuron dynamics
283        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        // Check for errors and collect voltages
295        let voltages: Result<Vec<_>> = results.into_iter().collect();
296        let voltages = voltages?;
297
298        // Sequential spike detection (typically fast)
299        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        // Record voltages if enabled
309        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    /// Get recorded voltage traces.
320    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    /// Clear all recordings and spike times.
325    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    /// Reset all neurons to resting state.
334    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    /// Calculate population statistics.
345    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        // Count spikes in time window
354        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, // Calculated if time window provided
378        }
379    }
380
381    /// Calculate instantaneous population firing rate.
382    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 // Hz
392    }
393}
394
395/// Population statistics.
396#[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        // Out of bounds
446        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        // Neuron 0 should have spiked with strong current
459        assert!(!pop.spike_times[0].is_empty());
460        // Other neurons should be quiet
461        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); // Should be at rest
496    }
497
498    #[test]
499    fn test_instantaneous_rate() {
500        let mut pop = NeuralPopulation::excitatory("E", 10).unwrap();
501
502        // Inject currents to cause spikes
503        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); // Should have some activity
513    }
514}