1use crate::connectivity::ConnectionPattern;
6use crate::error::Result;
7use rand::Rng;
8use rand_distr::{Distribution, Normal, Uniform};
9use serde::{Deserialize, Serialize};
10use std::collections::VecDeque;
11use synapse_models::synapse::Synapse;
12
13#[derive(Clone)]
15pub struct Projection {
16 pub source_pop: usize,
18 pub target_pop: usize,
20 connections: Vec<Connection>,
23 delay_queue: VecDeque<DelayedSpike>,
25 max_delay: f64,
27}
28
29#[derive(Clone)]
31pub struct Connection {
32 pub source: usize,
34 pub target: usize,
36 pub synapse: Synapse,
38 pub delay: f64,
40 pub weight: f64,
42}
43
44#[derive(Clone, Debug)]
46struct DelayedSpike {
47 delivery_time: f64,
49 #[allow(dead_code)]
51 source: usize,
52 #[allow(dead_code)]
54 target: usize,
55 connection_idx: usize,
57}
58
59impl Projection {
60 pub fn new<R: Rng>(
74 source_pop: usize,
75 target_pop: usize,
76 source_size: usize,
77 target_size: usize,
78 pattern: &ConnectionPattern,
79 synapse_template: &Synapse,
80 weight_init: WeightInit,
81 delay_init: DelayInit,
82 rng: &mut R,
83 ) -> Result<Self> {
84 let connection_pairs = pattern.generate(source_size, target_size, rng)?;
86
87 let mut connections = Vec::with_capacity(connection_pairs.len());
88 let mut max_delay: f64 = 0.0;
89
90 for (source, target) in connection_pairs {
91 let synapse = synapse_template.clone();
92 let weight = weight_init.sample(rng);
93 let delay = delay_init.sample(rng);
94 max_delay = max_delay.max(delay);
95
96 connections.push(Connection {
97 source,
98 target,
99 synapse,
100 delay,
101 weight,
102 });
103 }
104
105 Ok(Self {
106 source_pop,
107 target_pop,
108 connections,
109 delay_queue: VecDeque::new(),
110 max_delay,
111 })
112 }
113
114 pub fn all_to_all<R: Rng>(
116 source_pop: usize,
117 target_pop: usize,
118 source_size: usize,
119 target_size: usize,
120 synapse_template: &Synapse,
121 weight: f64,
122 delay: f64,
123 rng: &mut R,
124 ) -> Result<Self> {
125 Self::new(
126 source_pop,
127 target_pop,
128 source_size,
129 target_size,
130 &ConnectionPattern::AllToAll,
131 synapse_template,
132 WeightInit::Constant(weight),
133 DelayInit::Constant(delay),
134 rng,
135 )
136 }
137
138 pub fn one_to_one<R: Rng>(
140 source_pop: usize,
141 target_pop: usize,
142 size: usize,
143 synapse_template: &Synapse,
144 weight: f64,
145 delay: f64,
146 rng: &mut R,
147 ) -> Result<Self> {
148 Self::new(
149 source_pop,
150 target_pop,
151 size,
152 size,
153 &ConnectionPattern::OneToOne,
154 synapse_template,
155 WeightInit::Constant(weight),
156 DelayInit::Constant(delay),
157 rng,
158 )
159 }
160
161 pub fn fixed_probability<R: Rng>(
163 source_pop: usize,
164 target_pop: usize,
165 source_size: usize,
166 target_size: usize,
167 probability: f64,
168 synapse_template: &Synapse,
169 weight: f64,
170 delay: f64,
171 rng: &mut R,
172 ) -> Result<Self> {
173 Self::new(
174 source_pop,
175 target_pop,
176 source_size,
177 target_size,
178 &ConnectionPattern::FixedProbability(probability),
179 synapse_template,
180 WeightInit::Constant(weight),
181 DelayInit::Constant(delay),
182 rng,
183 )
184 }
185
186 pub fn register_spike(&mut self, source_idx: usize, current_time: f64) {
188 for (conn_idx, conn) in self.connections.iter().enumerate() {
190 if conn.source == source_idx {
191 self.delay_queue.push_back(DelayedSpike {
192 delivery_time: current_time + conn.delay,
193 source: source_idx,
194 target: conn.target,
195 connection_idx: conn_idx,
196 });
197 }
198 }
199 }
200
201 pub fn process_spikes(&mut self, current_time: f64, target_voltages: &[f64], dt: f64) -> Result<Vec<(usize, f64)>> {
205 let mut target_currents: Vec<f64> = vec![0.0; target_voltages.len()];
206
207 while let Some(spike) = self.delay_queue.front() {
209 if spike.delivery_time <= current_time {
210 let spike = self.delay_queue.pop_front().unwrap();
211 let conn = &mut self.connections[spike.connection_idx];
213 conn.synapse.presynaptic_spike(current_time)?;
214 } else {
215 break; }
217 }
218
219 for conn in self.connections.iter_mut() {
221 let target_voltage = target_voltages[conn.target];
222 conn.synapse.update(current_time, target_voltage, dt)?;
223
224 let current = conn.synapse.current(target_voltage) * conn.weight;
225 target_currents[conn.target] += current;
226 }
227
228 let result: Vec<(usize, f64)> = target_currents
230 .iter()
231 .enumerate()
232 .filter(|(_, ¤t)| current.abs() > 1e-12)
233 .map(|(idx, ¤t)| (idx, current))
234 .collect();
235
236 Ok(result)
237 }
238
239 pub fn num_connections(&self) -> usize {
241 self.connections.len()
242 }
243
244 pub fn max_delay(&self) -> f64 {
246 self.max_delay
247 }
248
249 pub fn connections(&self) -> &[Connection] {
251 &self.connections
252 }
253
254 pub fn apply_stdp(&mut self, source_spike_times: &[Vec<f64>], target_spike_times: &[Vec<f64>]) -> Result<()> {
256 for conn in self.connections.iter_mut() {
257 let source_spikes = &source_spike_times[conn.source];
258 let target_spikes = &target_spike_times[conn.target];
259
260 for &t_pre in source_spikes {
262 for &t_post in target_spikes {
263 let dt = t_post - t_pre;
264 if dt.abs() < 40.0 { conn.synapse.presynaptic_spike(t_pre)?;
268 conn.synapse.postsynaptic_spike(t_post)?;
269 }
270 }
271 }
272 }
273 Ok(())
274 }
275
276 pub fn scale_weights(&mut self, factor: f64) {
278 for conn in self.connections.iter_mut() {
279 conn.weight *= factor;
280 }
281 }
282
283 pub fn weight_statistics(&self) -> WeightStats {
285 if self.connections.is_empty() {
286 return WeightStats {
287 mean: 0.0,
288 std: 0.0,
289 min: 0.0,
290 max: 0.0,
291 count: 0,
292 };
293 }
294
295 let weights: Vec<f64> = self.connections.iter().map(|c| c.weight).collect();
296 let mean = weights.iter().sum::<f64>() / weights.len() as f64;
297 let variance = weights.iter()
298 .map(|w| (w - mean).powi(2))
299 .sum::<f64>() / weights.len() as f64;
300 let std = variance.sqrt();
301
302 let min = weights.iter().cloned().fold(f64::INFINITY, f64::min);
303 let max = weights.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
304
305 WeightStats {
306 mean,
307 std,
308 min,
309 max,
310 count: weights.len(),
311 }
312 }
313}
314
315#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
317pub enum WeightInit {
318 Constant(f64),
320 Uniform { min: f64, max: f64 },
322 Normal { mean: f64, std: f64 },
324}
325
326impl WeightInit {
327 fn sample<R: Rng>(&self, rng: &mut R) -> f64 {
328 match self {
329 WeightInit::Constant(w) => *w,
330 WeightInit::Uniform { min, max } => {
331 Uniform::new(*min, *max).sample(rng)
332 }
333 WeightInit::Normal { mean, std } => {
334 Normal::new(*mean, *std).unwrap().sample(rng).max(0.0)
335 }
336 }
337 }
338}
339
340#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
342pub enum DelayInit {
343 Constant(f64),
345 Uniform { min: f64, max: f64 },
347 DistanceDependent { speed: f64 },
349}
350
351impl DelayInit {
352 fn sample<R: Rng>(&self, rng: &mut R) -> f64 {
353 match self {
354 DelayInit::Constant(d) => *d,
355 DelayInit::Uniform { min, max } => {
356 Uniform::new(*min, *max).sample(rng)
357 }
358 DelayInit::DistanceDependent { speed: _ } => {
359 1.0
361 }
362 }
363 }
364}
365
366#[derive(Debug, Clone, Serialize, Deserialize)]
368pub struct WeightStats {
369 pub mean: f64,
370 pub std: f64,
371 pub min: f64,
372 pub max: f64,
373 pub count: usize,
374}
375
376#[cfg(test)]
377mod tests {
378 use super::*;
379
380 #[test]
381 fn test_all_to_all_projection() {
382 let mut rng = rand::thread_rng();
383 let synapse = Synapse::excitatory(1.0, 1.0).unwrap();
384
385 let proj = Projection::all_to_all(0, 1, 3, 2, &synapse, 1.0, 0.5, &mut rng).unwrap();
386
387 assert_eq!(proj.num_connections(), 6); assert_eq!(proj.source_pop, 0);
389 assert_eq!(proj.target_pop, 1);
390 }
391
392 #[test]
393 fn test_one_to_one_projection() {
394 let mut rng = rand::thread_rng();
395 let synapse = Synapse::excitatory(1.0, 1.0).unwrap();
396
397 let proj = Projection::one_to_one(0, 1, 5, &synapse, 1.0, 0.5, &mut rng).unwrap();
398
399 assert_eq!(proj.num_connections(), 5);
400 }
401
402 #[test]
403 fn test_fixed_probability_projection() {
404 let mut rng = rand::thread_rng();
405 let synapse = Synapse::excitatory(1.0, 1.0).unwrap();
406
407 let proj = Projection::fixed_probability(
408 0, 1, 10, 10, 0.5, &synapse, 1.0, 0.5, &mut rng
409 ).unwrap();
410
411 let n_conn = proj.num_connections();
413 assert!(n_conn > 20 && n_conn < 80);
414 }
415
416 #[test]
417 fn test_spike_registration() {
418 let mut rng = rand::thread_rng();
419 let synapse = Synapse::excitatory(1.0, 1.0).unwrap();
420
421 let mut proj = Projection::one_to_one(0, 1, 3, &synapse, 1.0, 1.0, &mut rng).unwrap();
422
423 proj.register_spike(0, 0.0);
424 assert_eq!(proj.delay_queue.len(), 1);
425
426 proj.register_spike(1, 1.0);
427 assert_eq!(proj.delay_queue.len(), 2);
428 }
429
430 #[test]
431 fn test_spike_processing_with_delay() {
432 let mut rng = rand::thread_rng();
433 let synapse = Synapse::excitatory(1.0, 1.0).unwrap();
434
435 let mut proj = Projection::one_to_one(0, 1, 2, &synapse, 1.0, 2.0, &mut rng).unwrap();
436
437 proj.register_spike(0, 0.0);
439
440 let voltages = vec![-65.0; 2];
442 let currents = proj.process_spikes(1.0, &voltages, 0.1).unwrap();
443 assert!(currents.is_empty()); let currents = proj.process_spikes(2.5, &voltages, 0.1).unwrap();
447 assert_eq!(proj.delay_queue.len(), 0);
449 }
450
451 #[test]
452 fn test_weight_scaling() {
453 let mut rng = rand::thread_rng();
454 let synapse = Synapse::excitatory(1.0, 1.0).unwrap();
455
456 let mut proj = Projection::all_to_all(0, 1, 2, 2, &synapse, 2.0, 0.5, &mut rng).unwrap();
457
458 proj.scale_weights(0.5);
459
460 for conn in proj.connections() {
461 assert!((conn.weight - 1.0).abs() < 1e-10);
462 }
463 }
464
465 #[test]
466 fn test_weight_statistics() {
467 let mut rng = rand::thread_rng();
468 let synapse = Synapse::excitatory(1.0, 1.0).unwrap();
469
470 let proj = Projection::new(
471 0, 1, 5, 5,
472 &ConnectionPattern::AllToAll,
473 &synapse,
474 WeightInit::Normal { mean: 1.0, std: 0.1 },
475 DelayInit::Constant(1.0),
476 &mut rng,
477 ).unwrap();
478
479 let stats = proj.weight_statistics();
480 assert_eq!(stats.count, 25);
481 assert!(stats.mean > 0.0);
482 }
483}