1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
16pub enum ConnectionPattern {
17 AllToAll,
19 OneToOne,
21 FixedProbability(f64),
23 FixedNumber(usize),
25 SmallWorld { k: usize, p: f64 },
27 ScaleFree { m: usize },
29 Gaussian { sigma: f64 },
31}
32
33impl ConnectionPattern {
34 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
67fn 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
78fn 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
92fn 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
119fn 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
148pub 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 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 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 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 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 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
224pub 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 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 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 while targets.len() < m {
275 let threshold = rng.gen::<f64>() * total_degree as f64;
276 let mut cumulative = 0.0;
277
278 for (i, °) 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 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 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
306fn 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 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
339pub fn network_statistics(connections: &[(usize, usize)], n_nodes: usize) -> NetworkStats {
341 let n_connections = connections.len();
342
343 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#[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 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 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 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 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 assert!(!connections.is_empty());
434
435 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 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); }
477}