1use 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
12pub struct Network {
14 populations: Vec<NeuralPopulation>,
16 projections: Vec<Projection>,
18 current_time: f64,
20 dt: f64,
22 stimulations: Vec<(usize, Box<dyn Stimulation>)>, pub spike_recorder: Option<SpikeRecorder>,
26 pub voltage_recorder: Option<VoltageRecorder>,
27 pub rate_recorder: Option<PopulationRateRecorder>,
28}
29
30impl Network {
31 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 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 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 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 pub fn enable_spike_recording(&mut self) {
81 self.spike_recorder = Some(SpikeRecorder::new());
82 }
83
84 pub fn enable_voltage_recording(&mut self) {
86 self.voltage_recorder = Some(VoltageRecorder::new(self.dt));
87 }
88
89 pub fn enable_rate_recording(&mut self, window: f64) {
91 self.rate_recorder = Some(PopulationRateRecorder::new(window));
92 }
93
94 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 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 pub fn current_time(&self) -> f64 {
113 self.current_time
114 }
115
116 pub fn num_populations(&self) -> usize {
118 self.populations.len()
119 }
120
121 pub fn num_projections(&self) -> usize {
123 self.projections.len()
124 }
125
126 pub fn step(&mut self) -> Result<()> {
128 if self.populations.is_empty() {
129 return Err(NeuralDynamicsError::EmptyNetwork);
130 }
131
132 for pop in self.populations.iter_mut() {
134 pop.reset_synaptic_currents();
135 }
136
137 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 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 for (target_idx, current) in synaptic_currents {
157 self.populations[proj.target_pop].add_external_current(target_idx, current)?;
158 }
159 }
160
161 if self.populations.len() > 1 {
163 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 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 if last_spike >= self.current_time - self.dt && last_spike < self.current_time {
181 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 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 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 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 self.current_time += self.dt;
217
218 Ok(())
219 }
220
221 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 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 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#[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
288pub struct NetworkBuilder {
290 network: Network,
291 rng: rand::rngs::ThreadRng,
292}
293
294impl NetworkBuilder {
295 pub fn new(dt: f64) -> Result<Self> {
297 Ok(Self {
298 network: Network::new(dt)?,
299 rng: rand::thread_rng(),
300 })
301 }
302
303 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 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 pub fn add_population(mut self, population: NeuralPopulation) -> Self {
327 self.network.add_population(population);
328 self
329 }
330
331 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 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 pub fn with_spike_recording(mut self) -> Self {
377 self.network.enable_spike_recording();
378 self
379 }
380
381 pub fn with_voltage_recording(mut self) -> Self {
383 self.network.enable_voltage_recording();
384 self
385 }
386
387 pub fn with_rate_recording(mut self, window: f64) -> Self {
389 self.network.enable_rate_recording(window);
390 self
391 }
392
393 pub fn build(self) -> Network {
395 self.network
396 }
397}
398
399#[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 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 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}