Skip to main content

neural_dynamics/
projection.rs

1//! Inter-population connectivity (projections).
2//!
3//! This module defines how populations connect to each other through synaptic projections.
4
5use 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/// A projection from a source population to a target population.
14#[derive(Clone)]
15pub struct Projection {
16    /// Source population index
17    pub source_pop: usize,
18    /// Target population index
19    pub target_pop: usize,
20    /// Connection matrix (sparse representation)
21    /// Each element is (source_idx, target_idx, synapse)
22    connections: Vec<Connection>,
23    /// Delay queue for spike propagation
24    delay_queue: VecDeque<DelayedSpike>,
25    /// Maximum delay in ms
26    max_delay: f64,
27}
28
29/// A single synaptic connection with delay.
30#[derive(Clone)]
31pub struct Connection {
32    /// Source neuron index in source population
33    pub source: usize,
34    /// Target neuron index in target population
35    pub target: usize,
36    /// Synaptic model
37    pub synapse: Synapse,
38    /// Transmission delay (ms)
39    pub delay: f64,
40    /// Connection weight (scaling factor)
41    pub weight: f64,
42}
43
44/// A spike waiting to be delivered.
45#[derive(Clone, Debug)]
46struct DelayedSpike {
47    /// Time when spike should be delivered
48    delivery_time: f64,
49    /// Source neuron index
50    #[allow(dead_code)]
51    source: usize,
52    /// Target neuron index
53    #[allow(dead_code)]
54    target: usize,
55    /// Connection index in connections vector
56    connection_idx: usize,
57}
58
59impl Projection {
60    /// Create a new projection with specified connection pattern.
61    ///
62    /// # Arguments
63    ///
64    /// * `source_pop` - Source population index
65    /// * `target_pop` - Target population index
66    /// * `source_size` - Number of neurons in source population
67    /// * `target_size` - Number of neurons in target population
68    /// * `pattern` - Connection pattern
69    /// * `synapse_template` - Template synapse (will be cloned for each connection)
70    /// * `weight_init` - Weight initialization scheme
71    /// * `delay_init` - Delay initialization scheme
72    /// * `rng` - Random number generator
73    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        // Generate connection list based on pattern
85        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    /// Create an all-to-all projection.
115    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    /// Create a one-to-one projection.
139    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    /// Create a random fixed-probability projection.
162    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    /// Register a spike from source population.
187    pub fn register_spike(&mut self, source_idx: usize, current_time: f64) {
188        // Find all connections from this source neuron
189        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    /// Process delayed spikes and update synapses.
202    ///
203    /// Returns list of (target_idx, synaptic_current) pairs.
204    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        // Deliver spikes whose time has come
208        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                // Trigger presynaptic spike in synapse
212                let conn = &mut self.connections[spike.connection_idx];
213                conn.synapse.presynaptic_spike(current_time)?;
214            } else {
215                break; // Queue is sorted by time
216            }
217        }
218
219        // Update all synapses and collect currents
220        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        // Convert to sparse representation
229        let result: Vec<(usize, f64)> = target_currents
230            .iter()
231            .enumerate()
232            .filter(|(_, &current)| current.abs() > 1e-12)
233            .map(|(idx, &current)| (idx, current))
234            .collect();
235
236        Ok(result)
237    }
238
239    /// Get number of connections.
240    pub fn num_connections(&self) -> usize {
241        self.connections.len()
242    }
243
244    /// Get maximum delay.
245    pub fn max_delay(&self) -> f64 {
246        self.max_delay
247    }
248
249    /// Get all connections.
250    pub fn connections(&self) -> &[Connection] {
251        &self.connections
252    }
253
254    /// Apply STDP learning rule to all synapses.
255    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            // Apply STDP based on spike timing
261            for &t_pre in source_spikes {
262                for &t_post in target_spikes {
263                    let dt = t_post - t_pre;
264                    // STDP is handled by the synapse itself
265                    // We just notify it of spike pairs
266                    if dt.abs() < 40.0 { // STDP window typically ±20-40 ms
267                        conn.synapse.presynaptic_spike(t_pre)?;
268                        conn.synapse.postsynaptic_spike(t_post)?;
269                    }
270                }
271            }
272        }
273        Ok(())
274    }
275
276    /// Scale all weights by a factor.
277    pub fn scale_weights(&mut self, factor: f64) {
278        for conn in self.connections.iter_mut() {
279            conn.weight *= factor;
280        }
281    }
282
283    /// Get weight statistics.
284    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/// Weight initialization schemes.
316#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
317pub enum WeightInit {
318    /// Constant weight for all connections
319    Constant(f64),
320    /// Uniform distribution [min, max]
321    Uniform { min: f64, max: f64 },
322    /// Normal distribution
323    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/// Delay initialization schemes.
341#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
342pub enum DelayInit {
343    /// Constant delay for all connections
344    Constant(f64),
345    /// Uniform distribution [min, max]
346    Uniform { min: f64, max: f64 },
347    /// Distance-dependent delay
348    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                // For now, return constant (distance calculation requires spatial info)
360                1.0
361            }
362        }
363    }
364}
365
366/// Weight statistics.
367#[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); // 3 * 2
388        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        // Probabilistic, but should have roughly 50 connections (10*10*0.5)
412        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        // Spike at t=0
438        proj.register_spike(0, 0.0);
439
440        // Process at t=1.0 (before delay)
441        let voltages = vec![-65.0; 2];
442        let currents = proj.process_spikes(1.0, &voltages, 0.1).unwrap();
443        assert!(currents.is_empty()); // Delay not yet elapsed
444
445        // Process at t=2.5 (after delay)
446        let currents = proj.process_spikes(2.5, &voltages, 0.1).unwrap();
447        // Spike should have been delivered
448        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}