Skip to main content

neural_dynamics/
network.rs

1//! Neural network architecture and simulation control.
2
3use crate::connectivity::ConnectionPattern;
4use crate::error::{NeuralDynamicsError, Result};
5use crate::population::NeuralPopulation;
6use crate::projection::{DelayInit, Projection, WeightInit};
7use crate::recording::{PopulationRateRecorder, SpikeRecorder, VoltageRecorder};
8use crate::stimulation::Stimulation;
9use serde::{Deserialize, Serialize};
10use synapse_models::synapse::Synapse;
11
12/// Neural network with multiple populations and projections.
13pub struct Network {
14    /// Neural populations
15    populations: Vec<NeuralPopulation>,
16    /// Inter-population projections
17    projections: Vec<Projection>,
18    /// Current simulation time (ms)
19    current_time: f64,
20    /// Time step (ms)
21    dt: f64,
22    /// External stimulation protocols
23    stimulations: Vec<(usize, Box<dyn Stimulation>)>, // (pop_idx, stimulation)
24    /// Recorders
25    pub spike_recorder: Option<SpikeRecorder>,
26    pub voltage_recorder: Option<VoltageRecorder>,
27    pub rate_recorder: Option<PopulationRateRecorder>,
28}
29
30impl Network {
31    /// Create a new empty network.
32    pub fn new(dt: f64) -> Result<Self> {
33        if dt <= 0.0 {
34            return Err(NeuralDynamicsError::InvalidParameter {
35                parameter: "dt".to_string(),
36                value: dt,
37                reason: "must be positive".to_string(),
38            });
39        }
40
41        Ok(Self {
42            populations: Vec::new(),
43            projections: Vec::new(),
44            current_time: 0.0,
45            dt,
46            stimulations: Vec::new(),
47            spike_recorder: None,
48            voltage_recorder: None,
49            rate_recorder: None,
50        })
51    }
52
53    /// Add a population to the network.
54    pub fn add_population(&mut self, population: NeuralPopulation) -> usize {
55        let idx = self.populations.len();
56        self.populations.push(population);
57        idx
58    }
59
60    /// Add a projection between populations.
61    pub fn add_projection(&mut self, projection: Projection) -> usize {
62        let idx = self.projections.len();
63        self.projections.push(projection);
64        idx
65    }
66
67    /// Add external stimulation to a population.
68    pub fn add_stimulation(&mut self, pop_idx: usize, stim: Box<dyn Stimulation>) -> Result<()> {
69        if pop_idx >= self.populations.len() {
70            return Err(NeuralDynamicsError::InvalidPopulationIndex {
71                index: pop_idx,
72                max: self.populations.len() - 1,
73            });
74        }
75        self.stimulations.push((pop_idx, stim));
76        Ok(())
77    }
78
79    /// Enable spike recording.
80    pub fn enable_spike_recording(&mut self) {
81        self.spike_recorder = Some(SpikeRecorder::new());
82    }
83
84    /// Enable voltage recording.
85    pub fn enable_voltage_recording(&mut self) {
86        self.voltage_recorder = Some(VoltageRecorder::new(self.dt));
87    }
88
89    /// Enable population rate recording.
90    pub fn enable_rate_recording(&mut self, window: f64) {
91        self.rate_recorder = Some(PopulationRateRecorder::new(window));
92    }
93
94    /// Get a population by index.
95    pub fn get_population(&self, idx: usize) -> Result<&NeuralPopulation> {
96        self.populations.get(idx).ok_or(NeuralDynamicsError::InvalidPopulationIndex {
97            index: idx,
98            max: self.populations.len().saturating_sub(1),
99        })
100    }
101
102    /// Get a mutable population by index.
103    pub fn get_population_mut(&mut self, idx: usize) -> Result<&mut NeuralPopulation> {
104        let max = self.populations.len().saturating_sub(1);
105        self.populations.get_mut(idx).ok_or(NeuralDynamicsError::InvalidPopulationIndex {
106            index: idx,
107            max,
108        })
109    }
110
111    /// Get current simulation time.
112    pub fn current_time(&self) -> f64 {
113        self.current_time
114    }
115
116    /// Get number of populations.
117    pub fn num_populations(&self) -> usize {
118        self.populations.len()
119    }
120
121    /// Get number of projections.
122    pub fn num_projections(&self) -> usize {
123        self.projections.len()
124    }
125
126    /// Step the simulation forward by one time step.
127    pub fn step(&mut self) -> Result<()> {
128        if self.populations.is_empty() {
129            return Err(NeuralDynamicsError::EmptyNetwork);
130        }
131
132        // 1. Reset synaptic currents
133        for pop in self.populations.iter_mut() {
134            pop.reset_synaptic_currents();
135        }
136
137        // 2. Apply external stimulations
138        for (pop_idx, stim) in self.stimulations.iter_mut() {
139            let pop_size = self.populations[*pop_idx].size;
140            for neuron_idx in 0..pop_size {
141                let current = stim.current(neuron_idx, self.current_time, self.dt);
142                if current != 0.0 {
143                    self.populations[*pop_idx].add_external_current(neuron_idx, current)?;
144                }
145            }
146        }
147
148        // 3. Process spikes through projections and accumulate synaptic currents
149        for proj in self.projections.iter_mut() {
150            let target_pop = &self.populations[proj.target_pop];
151            let target_voltages = target_pop.get_voltages();
152
153            let synaptic_currents = proj.process_spikes(self.current_time, &target_voltages, self.dt)?;
154
155            // Apply synaptic currents to target population
156            for (target_idx, current) in synaptic_currents {
157                self.populations[proj.target_pop].add_external_current(target_idx, current)?;
158            }
159        }
160
161        // 4. Update all populations (parallel if beneficial)
162        if self.populations.len() > 1 {
163            // Parallel update - more complex due to borrow checker
164            // For now, use sequential update
165            for pop in self.populations.iter_mut() {
166                pop.update(self.dt, self.current_time)?;
167            }
168        } else {
169            for pop in self.populations.iter_mut() {
170                pop.update(self.dt, self.current_time)?;
171            }
172        }
173
174        // 5. Register spikes with projections
175        for (pop_idx, pop) in self.populations.iter().enumerate() {
176            for neuron_idx in 0..pop.size {
177                let spike_times = pop.get_spike_times(neuron_idx)?;
178                if let Some(&last_spike) = spike_times.last() {
179                    // Check if spike occurred in this time step
180                    if last_spike >= self.current_time - self.dt && last_spike < self.current_time {
181                        // Register spike with all relevant projections
182                        for proj in self.projections.iter_mut() {
183                            if proj.source_pop == pop_idx {
184                                proj.register_spike(neuron_idx, last_spike);
185                            }
186                        }
187
188                        // Record spike if recording enabled
189                        if let Some(ref mut recorder) = self.spike_recorder {
190                            recorder.record_spike(pop_idx, neuron_idx, last_spike);
191                        }
192                    }
193                }
194            }
195        }
196
197        // 6. Record voltages if enabled
198        if let Some(ref mut recorder) = self.voltage_recorder {
199            for (pop_idx, pop) in self.populations.iter().enumerate() {
200                for neuron_idx in 0..pop.size {
201                    let voltage = pop.get_voltage(neuron_idx)?;
202                    recorder.record(pop_idx, neuron_idx, self.current_time, voltage);
203                }
204            }
205        }
206
207        // 7. Record population rates if enabled
208        if let Some(ref mut recorder) = self.rate_recorder {
209            for (pop_idx, pop) in self.populations.iter().enumerate() {
210                let rate = pop.instantaneous_rate(self.current_time, recorder.window);
211                recorder.record(pop_idx, self.current_time, rate);
212            }
213        }
214
215        // Advance time
216        self.current_time += self.dt;
217
218        Ok(())
219    }
220
221    /// Run simulation for a specified duration.
222    pub fn run(&mut self, duration: f64) -> Result<()> {
223        let n_steps = (duration / self.dt).ceil() as usize;
224
225        for _ in 0..n_steps {
226            self.step()?;
227        }
228
229        Ok(())
230    }
231
232    /// Reset the network to initial state.
233    pub fn reset(&mut self) {
234        self.current_time = 0.0;
235
236        for pop in self.populations.iter_mut() {
237            pop.reset();
238        }
239
240        for (_pop_idx, stim) in self.stimulations.iter_mut() {
241            stim.reset();
242        }
243
244        if let Some(ref mut recorder) = self.spike_recorder {
245            recorder.clear();
246        }
247        if let Some(ref mut recorder) = self.voltage_recorder {
248            recorder.clear();
249        }
250        if let Some(ref mut recorder) = self.rate_recorder {
251            recorder.clear();
252        }
253    }
254
255    /// Get network statistics.
256    pub fn statistics(&self) -> NetworkStats {
257        let total_neurons: usize = self.populations.iter().map(|p| p.size).sum();
258        let total_connections: usize = self.projections.iter().map(|p| p.num_connections()).sum();
259
260        let total_spikes = if let Some(ref recorder) = self.spike_recorder {
261            recorder.total_spikes()
262        } else {
263            0
264        };
265
266        NetworkStats {
267            n_populations: self.populations.len(),
268            n_projections: self.projections.len(),
269            total_neurons,
270            total_connections,
271            total_spikes,
272            current_time: self.current_time,
273        }
274    }
275}
276
277/// Network statistics.
278#[derive(Debug, Clone, Serialize, Deserialize)]
279pub struct NetworkStats {
280    pub n_populations: usize,
281    pub n_projections: usize,
282    pub total_neurons: usize,
283    pub total_connections: usize,
284    pub total_spikes: usize,
285    pub current_time: f64,
286}
287
288/// Builder pattern for constructing networks.
289pub struct NetworkBuilder {
290    network: Network,
291    rng: rand::rngs::ThreadRng,
292}
293
294impl NetworkBuilder {
295    /// Create a new network builder.
296    pub fn new(dt: f64) -> Result<Self> {
297        Ok(Self {
298            network: Network::new(dt)?,
299            rng: rand::thread_rng(),
300        })
301    }
302
303    /// Add an excitatory population.
304    pub fn add_excitatory_population(
305        mut self,
306        id: impl Into<String>,
307        size: usize,
308    ) -> Result<Self> {
309        let pop = NeuralPopulation::excitatory(id, size)?;
310        self.network.add_population(pop);
311        Ok(self)
312    }
313
314    /// Add an inhibitory population.
315    pub fn add_inhibitory_population(
316        mut self,
317        id: impl Into<String>,
318        size: usize,
319    ) -> Result<Self> {
320        let pop = NeuralPopulation::inhibitory(id, size)?;
321        self.network.add_population(pop);
322        Ok(self)
323    }
324
325    /// Add a custom population.
326    pub fn add_population(mut self, population: NeuralPopulation) -> Self {
327        self.network.add_population(population);
328        self
329    }
330
331    /// Connect two populations.
332    pub fn connect(
333        mut self,
334        source_idx: usize,
335        target_idx: usize,
336        pattern: ConnectionPattern,
337        synapse_type: SynapseType,
338        weight: f64,
339        delay: f64,
340    ) -> Result<Self> {
341        let source_size = self.network.get_population(source_idx)?.size;
342        let target_size = self.network.get_population(target_idx)?.size;
343
344        let synapse = match synapse_type {
345            SynapseType::Excitatory => Synapse::excitatory(weight, delay)?,
346            SynapseType::Inhibitory => Synapse::inhibitory(weight, delay)?,
347        };
348
349        let projection = Projection::new(
350            source_idx,
351            target_idx,
352            source_size,
353            target_size,
354            &pattern,
355            &synapse,
356            WeightInit::Constant(weight),
357            DelayInit::Constant(delay),
358            &mut self.rng,
359        )?;
360
361        self.network.add_projection(projection);
362        Ok(self)
363    }
364
365    /// Add stimulation to a population.
366    pub fn add_stimulation(
367        mut self,
368        pop_idx: usize,
369        stim: Box<dyn Stimulation>,
370    ) -> Result<Self> {
371        self.network.add_stimulation(pop_idx, stim)?;
372        Ok(self)
373    }
374
375    /// Enable spike recording.
376    pub fn with_spike_recording(mut self) -> Self {
377        self.network.enable_spike_recording();
378        self
379    }
380
381    /// Enable voltage recording.
382    pub fn with_voltage_recording(mut self) -> Self {
383        self.network.enable_voltage_recording();
384        self
385    }
386
387    /// Enable rate recording.
388    pub fn with_rate_recording(mut self, window: f64) -> Self {
389        self.network.enable_rate_recording(window);
390        self
391    }
392
393    /// Build the network.
394    pub fn build(self) -> Network {
395        self.network
396    }
397}
398
399/// Synapse type for builder.
400#[derive(Debug, Clone, Copy)]
401pub enum SynapseType {
402    Excitatory,
403    Inhibitory,
404}
405
406#[cfg(test)]
407mod tests {
408    use super::*;
409    use crate::stimulation::CurrentInjection;
410
411    #[test]
412    fn test_network_creation() {
413        let network = Network::new(0.1).unwrap();
414        assert_eq!(network.num_populations(), 0);
415        assert_eq!(network.current_time(), 0.0);
416    }
417
418    #[test]
419    fn test_add_population() {
420        let mut network = Network::new(0.1).unwrap();
421        let pop = NeuralPopulation::excitatory("E", 10).unwrap();
422        let idx = network.add_population(pop);
423
424        assert_eq!(idx, 0);
425        assert_eq!(network.num_populations(), 1);
426    }
427
428    #[test]
429    fn test_network_step() {
430        let mut network = Network::new(0.1).unwrap();
431        let pop = NeuralPopulation::excitatory("E", 5).unwrap();
432        network.add_population(pop);
433
434        let time_before = network.current_time();
435        network.step().unwrap();
436        let time_after = network.current_time();
437
438        assert_eq!(time_after, time_before + 0.1);
439    }
440
441    #[test]
442    fn test_network_run() {
443        let mut network = Network::new(0.1).unwrap();
444        let pop = NeuralPopulation::excitatory("E", 5).unwrap();
445        network.add_population(pop);
446
447        network.run(10.0).unwrap();
448        assert!((network.current_time() - 10.0).abs() < 0.01);
449    }
450
451    #[test]
452    fn test_network_with_projection() {
453        let mut rng = rand::thread_rng();
454        let mut network = Network::new(0.1).unwrap();
455
456        let pop1 = NeuralPopulation::excitatory("E", 3).unwrap();
457        let pop2 = NeuralPopulation::excitatory("E2", 2).unwrap();
458
459        let idx1 = network.add_population(pop1);
460        let idx2 = network.add_population(pop2);
461
462        let synapse = Synapse::excitatory(1.0, 0.5).unwrap();
463        let proj = Projection::all_to_all(idx1, idx2, 3, 2, &synapse, 1.0, 0.5, &mut rng).unwrap();
464
465        network.add_projection(proj);
466
467        assert_eq!(network.num_projections(), 1);
468
469        // Should run without error
470        network.run(5.0).unwrap();
471    }
472
473    #[test]
474    fn test_network_with_stimulation() {
475        let mut network = Network::new(0.1).unwrap();
476        let pop = NeuralPopulation::excitatory("E", 5).unwrap();
477        let idx = network.add_population(pop);
478
479        let stim = CurrentInjection::new(10.0, 0.0, 10.0);
480        network.add_stimulation(idx, Box::new(stim)).unwrap();
481
482        network.run(15.0).unwrap();
483
484        // Neurons should have spiked with strong stimulation
485        let pop = network.get_population(idx).unwrap();
486        let has_spikes = (0..pop.size).any(|i| !pop.get_spike_times(i).unwrap().is_empty());
487        assert!(has_spikes);
488    }
489
490    #[test]
491    fn test_network_reset() {
492        let mut network = Network::new(0.1).unwrap();
493        let pop = NeuralPopulation::excitatory("E", 5).unwrap();
494        network.add_population(pop);
495
496        network.run(10.0).unwrap();
497        assert!(network.current_time() > 0.0);
498
499        network.reset();
500        assert_eq!(network.current_time(), 0.0);
501    }
502
503    #[test]
504    fn test_network_builder() {
505        let network = NetworkBuilder::new(0.1)
506            .unwrap()
507            .add_excitatory_population("E", 10)
508            .unwrap()
509            .add_inhibitory_population("I", 5)
510            .unwrap()
511            .with_spike_recording()
512            .build();
513
514        assert_eq!(network.num_populations(), 2);
515        assert!(network.spike_recorder.is_some());
516    }
517
518    #[test]
519    fn test_network_builder_with_connections() {
520        let network = NetworkBuilder::new(0.1)
521            .unwrap()
522            .add_excitatory_population("E", 10)
523            .unwrap()
524            .add_inhibitory_population("I", 5)
525            .unwrap()
526            .connect(
527                0,
528                1,
529                ConnectionPattern::FixedProbability(0.5),
530                SynapseType::Excitatory,
531                1.0,
532                1.0,
533            )
534            .unwrap()
535            .build();
536
537        assert_eq!(network.num_projections(), 1);
538    }
539
540    #[test]
541    fn test_spike_recording() {
542        let mut network = Network::new(0.1).unwrap();
543        network.enable_spike_recording();
544
545        let pop = NeuralPopulation::excitatory("E", 3).unwrap();
546        let idx = network.add_population(pop);
547
548        let stim = CurrentInjection::new(15.0, 0.0, 50.0);
549        network.add_stimulation(idx, Box::new(stim)).unwrap();
550
551        network.run(50.0).unwrap();
552
553        let recorder = network.spike_recorder.as_ref().unwrap();
554        assert!(recorder.total_spikes() > 0);
555    }
556
557    #[test]
558    fn test_network_statistics() {
559        let mut network = Network::new(0.1).unwrap();
560        let pop = NeuralPopulation::excitatory("E", 10).unwrap();
561        network.add_population(pop);
562
563        network.run(5.0).unwrap();
564
565        let stats = network.statistics();
566        assert_eq!(stats.n_populations, 1);
567        assert_eq!(stats.total_neurons, 10);
568    }
569}