Skip to main content

neural_dynamics/
connectivity.rs

1//! Network topology and connectivity pattern generation.
2//!
3//! This module implements various connectivity patterns including random,
4//! small-world, scale-free, and spatial networks.
5
6use crate::error::{NeuralDynamicsError, Result};
7use petgraph::graph::{Graph, NodeIndex};
8use petgraph::Undirected;
9use rand::Rng;
10use rand::seq::SliceRandom;
11use serde::{Deserialize, Serialize};
12use std::collections::HashSet;
13
14/// Connection pattern for projections between populations.
15#[derive(Debug, Clone, Serialize, Deserialize)]
16pub enum ConnectionPattern {
17    /// All neurons in source connect to all in target
18    AllToAll,
19    /// One-to-one mapping (requires equal sizes)
20    OneToOne,
21    /// Each pair connects with fixed probability
22    FixedProbability(f64),
23    /// Each target receives exactly n inputs
24    FixedNumber(usize),
25    /// Small-world network (Watts-Strogatz)
26    SmallWorld { k: usize, p: f64 },
27    /// Scale-free network (Barabási-Albert)
28    ScaleFree { m: usize },
29    /// Gaussian distance-dependent connectivity
30    Gaussian { sigma: f64 },
31}
32
33impl ConnectionPattern {
34    /// Generate list of connections (source_idx, target_idx) based on pattern.
35    pub fn generate<R: Rng>(
36        &self,
37        source_size: usize,
38        target_size: usize,
39        rng: &mut R,
40    ) -> Result<Vec<(usize, usize)>> {
41        match self {
42            ConnectionPattern::AllToAll => {
43                Ok(all_to_all_connections(source_size, target_size))
44            }
45            ConnectionPattern::OneToOne => {
46                one_to_one_connections(source_size, target_size)
47            }
48            ConnectionPattern::FixedProbability(p) => {
49                fixed_probability_connections(source_size, target_size, *p, rng)
50            }
51            ConnectionPattern::FixedNumber(n) => {
52                fixed_number_connections(source_size, target_size, *n, rng)
53            }
54            ConnectionPattern::SmallWorld { k, p } => {
55                small_world_connections(source_size, *k, *p, rng)
56            }
57            ConnectionPattern::ScaleFree { m } => {
58                scale_free_connections(source_size, *m, rng)
59            }
60            ConnectionPattern::Gaussian { sigma } => {
61                gaussian_connections(source_size, target_size, *sigma, rng)
62            }
63        }
64    }
65}
66
67/// Generate all-to-all connections.
68fn all_to_all_connections(source_size: usize, target_size: usize) -> Vec<(usize, usize)> {
69    let mut connections = Vec::with_capacity(source_size * target_size);
70    for i in 0..source_size {
71        for j in 0..target_size {
72            connections.push((i, j));
73        }
74    }
75    connections
76}
77
78/// Generate one-to-one connections.
79fn one_to_one_connections(source_size: usize, target_size: usize) -> Result<Vec<(usize, usize)>> {
80    if source_size != target_size {
81        return Err(NeuralDynamicsError::ConnectivityError {
82            reason: format!(
83                "OneToOne requires equal population sizes, got {} and {}",
84                source_size, target_size
85            ),
86        });
87    }
88
89    Ok((0..source_size).map(|i| (i, i)).collect())
90}
91
92/// Generate connections with fixed probability.
93fn fixed_probability_connections<R: Rng>(
94    source_size: usize,
95    target_size: usize,
96    probability: f64,
97    rng: &mut R,
98) -> Result<Vec<(usize, usize)>> {
99    if probability < 0.0 || probability > 1.0 {
100        return Err(NeuralDynamicsError::InvalidParameter {
101            parameter: "probability".to_string(),
102            value: probability,
103            reason: "must be in [0, 1]".to_string(),
104        });
105    }
106
107    let mut connections = Vec::new();
108    for i in 0..source_size {
109        for j in 0..target_size {
110            if rng.gen::<f64>() < probability {
111                connections.push((i, j));
112            }
113        }
114    }
115
116    Ok(connections)
117}
118
119/// Generate connections with fixed number of inputs per target neuron.
120fn fixed_number_connections<R: Rng>(
121    source_size: usize,
122    target_size: usize,
123    n_inputs: usize,
124    rng: &mut R,
125) -> Result<Vec<(usize, usize)>> {
126    if n_inputs > source_size {
127        return Err(NeuralDynamicsError::ConnectivityError {
128            reason: format!(
129                "Cannot have {} inputs from population of size {}",
130                n_inputs, source_size
131            ),
132        });
133    }
134
135    let mut connections = Vec::new();
136    let mut source_indices: Vec<usize> = (0..source_size).collect();
137
138    for target in 0..target_size {
139        source_indices.shuffle(rng);
140        for &source in source_indices.iter().take(n_inputs) {
141            connections.push((source, target));
142        }
143    }
144
145    Ok(connections)
146}
147
148/// Generate small-world network (Watts-Strogatz model).
149///
150/// # Arguments
151///
152/// * `n` - Number of nodes
153/// * `k` - Each node is connected to k nearest neighbors in ring topology
154/// * `p` - Probability of rewiring each edge
155/// * `rng` - Random number generator
156///
157/// # Returns
158///
159/// List of connections forming a small-world network
160pub fn small_world_connections<R: Rng>(
161    n: usize,
162    k: usize,
163    p: f64,
164    rng: &mut R,
165) -> Result<Vec<(usize, usize)>> {
166    if k >= n {
167        return Err(NeuralDynamicsError::ConnectivityError {
168            reason: "k must be less than n".to_string(),
169        });
170    }
171
172    if k % 2 != 0 {
173        return Err(NeuralDynamicsError::ConnectivityError {
174            reason: "k must be even".to_string(),
175        });
176    }
177
178    // Start with ring lattice
179    let mut edges: HashSet<(usize, usize)> = HashSet::new();
180
181    for i in 0..n {
182        for j in 1..=k / 2 {
183            let neighbor = (i + j) % n;
184            edges.insert((i.min(neighbor), i.max(neighbor)));
185        }
186    }
187
188    // Rewire with probability p
189    let edges_to_rewire: Vec<_> = edges.iter().cloned().collect();
190
191    for (u, v) in edges_to_rewire {
192        if rng.gen::<f64>() < p {
193            edges.remove(&(u, v));
194
195            // Find new target that doesn't create self-loop or duplicate
196            let mut new_target = rng.gen_range(0..n);
197            let mut attempts = 0;
198            while new_target == u || edges.contains(&(u.min(new_target), u.max(new_target))) {
199                new_target = rng.gen_range(0..n);
200                attempts += 1;
201                if attempts > 100 {
202                    // Give up and keep original edge
203                    edges.insert((u, v));
204                    break;
205                }
206            }
207
208            if attempts <= 100 {
209                edges.insert((u.min(new_target), u.max(new_target)));
210            }
211        }
212    }
213
214    // Convert to directed connections
215    let mut connections = Vec::new();
216    for (u, v) in edges {
217        connections.push((u, v));
218        connections.push((v, u));
219    }
220
221    Ok(connections)
222}
223
224/// Generate scale-free network (Barabási-Albert model).
225///
226/// # Arguments
227///
228/// * `n` - Final number of nodes
229/// * `m` - Number of edges to attach from new node to existing nodes
230/// * `rng` - Random number generator
231pub fn scale_free_connections<R: Rng>(
232    n: usize,
233    m: usize,
234    rng: &mut R,
235) -> Result<Vec<(usize, usize)>> {
236    if m >= n {
237        return Err(NeuralDynamicsError::ConnectivityError {
238            reason: "m must be less than n".to_string(),
239        });
240    }
241
242    if m == 0 {
243        return Ok(Vec::new());
244    }
245
246    let mut graph = Graph::<(), (), Undirected>::new_undirected();
247    let mut nodes: Vec<NodeIndex> = Vec::new();
248    let mut degrees: Vec<usize> = Vec::new();
249
250    // Start with m+1 nodes in a complete graph
251    for _ in 0..=m {
252        nodes.push(graph.add_node(()));
253        degrees.push(0);
254    }
255
256    for i in 0..=m {
257        for j in i + 1..=m {
258            graph.add_edge(nodes[i], nodes[j], ());
259            degrees[i] += 1;
260            degrees[j] += 1;
261        }
262    }
263
264    // Add remaining nodes with preferential attachment
265    for _ in (m + 1)..n {
266        let new_node = graph.add_node(());
267        nodes.push(new_node);
268        degrees.push(0);
269
270        let total_degree: usize = degrees.iter().sum();
271        let mut targets = HashSet::new();
272
273        // Select m targets using preferential attachment
274        while targets.len() < m {
275            let threshold = rng.gen::<f64>() * total_degree as f64;
276            let mut cumulative = 0.0;
277
278            for (i, &deg) in degrees.iter().enumerate() {
279                cumulative += deg as f64;
280                if cumulative >= threshold && !targets.contains(&i) {
281                    targets.insert(i);
282                    break;
283                }
284            }
285        }
286
287        // Add edges to selected targets
288        for &target in &targets {
289            graph.add_edge(new_node, nodes[target], ());
290            degrees[nodes.len() - 1] += 1;
291            degrees[target] += 1;
292        }
293    }
294
295    // Extract connections from graph edges
296    let mut connections = Vec::new();
297    for node_a in graph.node_indices() {
298        for node_b in graph.neighbors(node_a) {
299            connections.push((node_a.index(), node_b.index()));
300        }
301    }
302
303    Ok(connections)
304}
305
306/// Generate Gaussian distance-dependent connections.
307fn gaussian_connections<R: Rng>(
308    source_size: usize,
309    target_size: usize,
310    sigma: f64,
311    rng: &mut R,
312) -> Result<Vec<(usize, usize)>> {
313    if sigma <= 0.0 {
314        return Err(NeuralDynamicsError::InvalidParameter {
315            parameter: "sigma".to_string(),
316            value: sigma,
317            reason: "must be positive".to_string(),
318        });
319    }
320
321    // Assume 1D arrangement for simplicity
322    let mut connections = Vec::new();
323
324    for i in 0..source_size {
325        for j in 0..target_size {
326            let distance = ((i as f64 / source_size as f64) - (j as f64 / target_size as f64)).abs();
327            let probability = (-distance * distance / (2.0 * sigma * sigma)).exp();
328
329            if rng.gen::<f64>() < probability {
330                connections.push((i, j));
331            }
332        }
333    }
334
335    Ok(connections)
336}
337
338
339/// Calculate network statistics.
340pub fn network_statistics(connections: &[(usize, usize)], n_nodes: usize) -> NetworkStats {
341    let n_connections = connections.len();
342
343    // Calculate degree distribution
344    let mut in_degrees = vec![0; n_nodes];
345    let mut out_degrees = vec![0; n_nodes];
346
347    for &(source, target) in connections {
348        if source < n_nodes && target < n_nodes {
349            out_degrees[source] += 1;
350            in_degrees[target] += 1;
351        }
352    }
353
354    let mean_in_degree = in_degrees.iter().sum::<usize>() as f64 / n_nodes as f64;
355    let mean_out_degree = out_degrees.iter().sum::<usize>() as f64 / n_nodes as f64;
356
357    NetworkStats {
358        n_nodes,
359        n_connections,
360        mean_in_degree,
361        mean_out_degree,
362        max_in_degree: *in_degrees.iter().max().unwrap_or(&0),
363        max_out_degree: *out_degrees.iter().max().unwrap_or(&0),
364    }
365}
366
367/// Network statistics.
368#[derive(Debug, Clone, Serialize, Deserialize)]
369pub struct NetworkStats {
370    pub n_nodes: usize,
371    pub n_connections: usize,
372    pub mean_in_degree: f64,
373    pub mean_out_degree: f64,
374    pub max_in_degree: usize,
375    pub max_out_degree: usize,
376}
377
378#[cfg(test)]
379mod tests {
380    use super::*;
381    use approx::assert_relative_eq;
382
383    #[test]
384    fn test_all_to_all() {
385        let connections = all_to_all_connections(3, 2);
386        assert_eq!(connections.len(), 6);
387    }
388
389    #[test]
390    fn test_one_to_one() {
391        let connections = one_to_one_connections(5, 5).unwrap();
392        assert_eq!(connections.len(), 5);
393        assert_eq!(connections[0], (0, 0));
394        assert_eq!(connections[4], (4, 4));
395
396        // Unequal sizes should fail
397        assert!(one_to_one_connections(3, 5).is_err());
398    }
399
400    #[test]
401    fn test_fixed_probability() {
402        let mut rng = rand::thread_rng();
403        let connections = fixed_probability_connections(10, 10, 0.5, &mut rng).unwrap();
404
405        // Should have approximately 50 connections (10*10*0.5)
406        assert!(connections.len() > 30 && connections.len() < 70);
407    }
408
409    #[test]
410    fn test_fixed_number() {
411        let mut rng = rand::thread_rng();
412        let connections = fixed_number_connections(20, 10, 5, &mut rng).unwrap();
413
414        // Each of 10 targets gets 5 inputs
415        assert_eq!(connections.len(), 50);
416    }
417
418    #[test]
419    fn test_small_world() {
420        let mut rng = rand::thread_rng();
421        let connections = small_world_connections(20, 4, 0.3, &mut rng).unwrap();
422
423        // Should have connections (bidirectional from ring)
424        assert!(!connections.is_empty());
425    }
426
427    #[test]
428    fn test_scale_free() {
429        let mut rng = rand::thread_rng();
430        let connections = scale_free_connections(50, 3, &mut rng).unwrap();
431
432        // Should have connections
433        assert!(!connections.is_empty());
434
435        // Check that connections form a valid graph
436        let stats = network_statistics(&connections, 50);
437        assert_eq!(stats.n_nodes, 50);
438        assert!(stats.mean_in_degree > 0.0);
439    }
440
441    #[test]
442    fn test_gaussian_connections() {
443        let mut rng = rand::thread_rng();
444        let connections = gaussian_connections(20, 20, 0.2, &mut rng).unwrap();
445
446        // Should favor nearby connections
447        assert!(!connections.is_empty());
448    }
449
450    #[test]
451    fn test_network_statistics() {
452        let connections = vec![(0, 1), (0, 2), (1, 2), (2, 3)];
453        let stats = network_statistics(&connections, 4);
454
455        assert_eq!(stats.n_nodes, 4);
456        assert_eq!(stats.n_connections, 4);
457        assert_relative_eq!(stats.mean_out_degree, 1.0);
458    }
459
460
461    #[test]
462    fn test_connection_pattern_generate() {
463        let mut rng = rand::thread_rng();
464
465        let pattern = ConnectionPattern::AllToAll;
466        let connections = pattern.generate(3, 2, &mut rng).unwrap();
467        assert_eq!(connections.len(), 6);
468
469        let pattern = ConnectionPattern::OneToOne;
470        let connections = pattern.generate(4, 4, &mut rng).unwrap();
471        assert_eq!(connections.len(), 4);
472
473        let pattern = ConnectionPattern::FixedProbability(1.0);
474        let connections = pattern.generate(2, 3, &mut rng).unwrap();
475        assert_eq!(connections.len(), 6); // p=1.0 means all connections
476    }
477}