//! Temporal GNN for dynamic connectivity analysis. //! //! Provides models for analyzing time-varying brain connectivity graphs. use crate::error::{GnnError, GnnResult}; use crate::graph::BrainGraph; use crate::layers::{BrainPool, BrainPoolConfig, PoolMethod}; use crate::models::{BrainGNN, BrainNetCNN, BrainNetCNNConfig}; use ndarray::{Array1, Array2}; use serde::{Deserialize, Serialize}; /// Time window specification #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TimeWindow { /// Start time in seconds pub start: f64, /// End time in seconds pub end: f64, /// Window index pub index: usize, } impl TimeWindow { /// Create a new time window pub fn new(start: f64, end: f64, index: usize) -> Self { Self { start, end, index } } /// Duration of the window pub fn duration(&self) -> f64 { self.end - self.start } /// Center time of the window pub fn center(&self) -> f64 { f64::midpoint(self.start, self.end) } } /// Dynamic connectivity sequence #[derive(Debug, Clone)] pub struct DynamicConnectivity { /// Sequence of brain graphs (one per time window) pub graphs: Vec, /// Time windows pub windows: Vec, /// Sampling frequency pub sfreq: f64, } impl DynamicConnectivity { /// Create from a sequence of connectivity matrices pub fn from_matrices( matrices: &[Vec>], channel_names: &[&str], window_duration: f64, step: f64, threshold: f64, ) -> GnnResult { let mut graphs = Vec::with_capacity(matrices.len()); let mut windows = Vec::with_capacity(matrices.len()); for (i, matrix) in matrices.iter().enumerate() { let graph = BrainGraph::from_connectivity_matrix(matrix, channel_names, threshold)?; graphs.push(graph); let start = i as f64 * step; windows.push(TimeWindow::new(start, start + window_duration, i)); } Ok(Self { graphs, windows, sfreq: 1.0 / step, }) } /// Number of time points pub fn n_times(&self) -> usize { self.graphs.len() } /// Get graph at specific time index pub fn get(&self, index: usize) -> Option<&BrainGraph> { self.graphs.get(index) } /// Iterator over (time_window, graph) pairs pub fn iter(&self) -> impl Iterator { self.windows.iter().zip(self.graphs.iter()) } } /// Configuration for temporal GNN #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TemporalConfig { /// Spatial GNN configuration pub spatial: BrainNetCNNConfig, /// Number of temporal attention heads pub temporal_heads: usize, /// Temporal hidden dimension pub temporal_hidden: usize, /// Whether to use temporal positional encoding pub positional_encoding: bool, /// Maximum sequence length pub max_seq_len: usize, } impl Default for TemporalConfig { fn default() -> Self { Self { spatial: BrainNetCNNConfig { hidden_dims: vec![32, 64], n_classes: 2, ..Default::default() }, temporal_heads: 4, temporal_hidden: 64, positional_encoding: true, max_seq_len: 100, } } } /// Temporal attention mechanism #[derive(Debug, Clone)] pub struct TemporalAttention { /// Number of attention heads n_heads: usize, /// Dimension per head head_dim: usize, /// Query weights per head wq: Vec>, /// Key weights per head wk: Vec>, /// Value weights per head wv: Vec>, /// Output projection wo: Array2, /// Positional encoding matrix positional_encoding: Option>, } impl TemporalAttention { /// Create a new temporal attention layer pub fn new(d_model: usize, n_heads: usize, max_seq_len: usize, use_pe: bool) -> Self { let head_dim = d_model / n_heads; let scale = (1.0 / head_dim as f64).sqrt(); let mut wq = Vec::with_capacity(n_heads); let mut wk = Vec::with_capacity(n_heads); let mut wv = Vec::with_capacity(n_heads); for _ in 0..n_heads { wq.push(Array2::from_shape_fn((d_model, head_dim), |_| { rand_simple() * scale })); wk.push(Array2::from_shape_fn((d_model, head_dim), |_| { rand_simple() * scale })); wv.push(Array2::from_shape_fn((d_model, head_dim), |_| { rand_simple() * scale })); } let wo = Array2::from_shape_fn((d_model, d_model), |_| rand_simple() * scale); // Sinusoidal positional encoding let positional_encoding = if use_pe { let mut pe = Array2::zeros((max_seq_len, d_model)); for pos in 0..max_seq_len { for i in 0..(d_model / 2) { let angle = pos as f64 / (10000.0_f64).powf(2.0 * i as f64 / d_model as f64); pe[[pos, 2 * i]] = angle.sin(); pe[[pos, 2 * i + 1]] = angle.cos(); } } Some(pe) } else { None }; Self { n_heads, head_dim, wq, wk, wv, wo, positional_encoding, } } /// Forward pass on sequence of embeddings [n_times, d_model] pub fn forward(&self, x: &Array2) -> Array2 { let n_times = x.nrows(); let d_model = x.ncols(); // Add positional encoding let mut h = x.clone(); if let Some(ref pe) = self.positional_encoding { for t in 0..n_times { for d in 0..d_model { h[[t, d]] += pe[[t, d]]; } } } // Multi-head attention let mut head_outputs = Vec::with_capacity(self.n_heads); for head in 0..self.n_heads { let q = h.dot(&self.wq[head]); // [n_times, head_dim] let k = h.dot(&self.wk[head]); let v = h.dot(&self.wv[head]); // Attention scores let scale = 1.0 / (self.head_dim as f64).sqrt(); let scores = q.dot(&k.t()) * scale; // [n_times, n_times] // Softmax let mut attn = Array2::zeros((n_times, n_times)); for i in 0..n_times { let max_val = scores .row(i) .iter() .copied() .fold(f64::NEG_INFINITY, f64::max); let exp_sum: f64 = scores.row(i).iter().map(|&s| (s - max_val).exp()).sum(); for j in 0..n_times { attn[[i, j]] = (scores[[i, j]] - max_val).exp() / exp_sum; } } let out = attn.dot(&v); head_outputs.push(out); } // Concatenate heads let mut concat = Array2::zeros((n_times, d_model)); for (h, head_out) in head_outputs.iter().enumerate() { for t in 0..n_times { for d in 0..self.head_dim { concat[[t, h * self.head_dim + d]] = head_out[[t, d]]; } } } // Output projection + residual let output = concat.dot(&self.wo); x + &output } /// Get attention weights for interpretation pub fn get_attention_weights(&self, x: &Array2) -> Array2 { let n_times = x.nrows(); // Use first head for visualization let q = x.dot(&self.wq[0]); let k = x.dot(&self.wk[0]); let scale = 1.0 / (self.head_dim as f64).sqrt(); let scores = q.dot(&k.t()) * scale; // Softmax let mut attn = Array2::zeros((n_times, n_times)); for i in 0..n_times { let max_val = scores .row(i) .iter() .copied() .fold(f64::NEG_INFINITY, f64::max); let exp_sum: f64 = scores.row(i).iter().map(|&s| (s - max_val).exp()).sum(); for j in 0..n_times { attn[[i, j]] = (scores[[i, j]] - max_val).exp() / exp_sum; } } attn } } /// Temporal Brain GNN model /// /// Combines spatial GNN for each time window with temporal attention /// for capturing dynamic brain state changes. #[derive(Debug, Clone)] pub struct TemporalBrainGNN { config: TemporalConfig, /// Spatial GNN (applied to each time window) spatial_gnn: BrainNetCNN, /// Temporal attention temporal_attention: TemporalAttention, /// Temporal convolution weights temporal_conv: Array2, /// Pooling pool: BrainPool, /// Classifier classifier: Array2, classifier_bias: Array1, } impl TemporalBrainGNN { /// Create a new temporal brain GNN pub fn new(config: TemporalConfig) -> GnnResult { // Create spatial GNN let spatial_gnn = BrainNetCNN::new(config.spatial.clone())?; // Final dimension from spatial GNN let spatial_dim = *config .spatial .hidden_dims .last() .unwrap_or(&config.spatial.n_features); // Temporal attention let temporal_attention = TemporalAttention::new( spatial_dim, config.temporal_heads, config.max_seq_len, config.positional_encoding, ); // Temporal 1D convolution (simplified as matrix) let scale = (1.0 / config.temporal_hidden as f64).sqrt(); let temporal_conv = Array2::from_shape_fn((spatial_dim, config.temporal_hidden), |_| { rand_simple() * scale }); let pool = BrainPool::new(BrainPoolConfig { method: PoolMethod::Attention, ratio: 1.0, }); // Classifier let classifier = Array2::from_shape_fn((config.temporal_hidden, config.spatial.n_classes), |_| { rand_simple() * scale }); let classifier_bias = Array1::zeros(config.spatial.n_classes); Ok(Self { config, spatial_gnn, temporal_attention, temporal_conv, pool, classifier, classifier_bias, }) } /// Forward pass on dynamic connectivity pub fn forward(&self, dyn_conn: &DynamicConnectivity) -> GnnResult> { let temporal_embeddings = self.forward_temporal_embeddings(dyn_conn)?; // Global temporal pooling let pooled = self.pool.global_pool(&temporal_embeddings); // Apply temporal convolution let h = Array1::from_iter((0..self.config.temporal_hidden).map(|j| { (0..pooled.len()) .map(|i| pooled[i] * self.temporal_conv[[i, j]]) .sum::() })); // Classifier let mut logits = Array1::zeros(self.config.spatial.n_classes); for j in 0..self.config.spatial.n_classes { logits[j] = self.classifier_bias[j]; for i in 0..h.len() { logits[j] += h[i] * self.classifier[[i, j]]; } } Ok(logits) } /// Get temporal embeddings [n_times, spatial_dim] pub fn forward_temporal_embeddings( &self, dyn_conn: &DynamicConnectivity, ) -> GnnResult> { let n_times = dyn_conn.n_times(); if n_times == 0 { return Err(GnnError::InvalidGraph("No time windows".into())); } let spatial_dim = *self .config .spatial .hidden_dims .last() .unwrap_or(&self.config.spatial.n_features); // Get spatial embeddings for each time window let mut spatial_embeddings = Array2::zeros((n_times, spatial_dim)); for (t, graph) in dyn_conn.graphs.iter().enumerate() { let node_emb = self.spatial_gnn.forward_node_embeddings(graph)?; let pooled = self.pool.global_pool(&node_emb); for d in 0..spatial_dim.min(pooled.len()) { spatial_embeddings[[t, d]] = pooled[d]; } } // Apply temporal attention let temporal_out = self.temporal_attention.forward(&spatial_embeddings); Ok(temporal_out) } /// Get temporal attention weights for interpretation pub fn get_temporal_attention_weights( &self, dyn_conn: &DynamicConnectivity, ) -> GnnResult> { let temporal_embeddings = self.forward_temporal_embeddings(dyn_conn)?; Ok(self .temporal_attention .get_attention_weights(&temporal_embeddings)) } /// Predict brain state at each time window pub fn predict_states(&self, dyn_conn: &DynamicConnectivity) -> GnnResult> { let temporal_embeddings = self.forward_temporal_embeddings(dyn_conn)?; let mut states = Vec::with_capacity(dyn_conn.n_times()); for t in 0..dyn_conn.n_times() { // Get embedding for this time let h: Vec = temporal_embeddings.row(t).to_vec(); // Apply temporal conv let conv_out: Vec = (0..self.config.temporal_hidden) .map(|j| { (0..h.len()) .map(|i| h[i] * self.temporal_conv[[i, j]]) .sum() }) .collect(); // Classify let mut logits = vec![0.0; self.config.spatial.n_classes]; for j in 0..self.config.spatial.n_classes { logits[j] = self.classifier_bias[j]; for i in 0..conv_out.len() { logits[j] += conv_out[i] * self.classifier[[i, j]]; } } // Argmax let state = logits .iter() .enumerate() .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap()) .map_or(0, |(i, _)| i); states.push(state); } Ok(states) } } // Simple pseudo-random for initialization fn rand_simple() -> f64 { use std::sync::atomic::{AtomicU64, Ordering}; static SEED: AtomicU64 = AtomicU64::new(98765); let mut s = SEED.fetch_add(1, Ordering::Relaxed); s ^= s >> 12; s ^= s << 25; s ^= s >> 27; s = s.wrapping_mul(0x2545F4914F6CDD1D); (s as f64) / (u64::MAX as f64) } #[cfg(test)] mod tests { use super::*; use crate::graph::{BrainRegion, GraphBuilder, Hemisphere}; fn create_test_graph() -> BrainGraph { let mut builder = GraphBuilder::new(); builder.add_node("F3", Hemisphere::Left, BrainRegion::Frontal); builder.add_node("F4", Hemisphere::Right, BrainRegion::Frontal); builder.add_node("O1", Hemisphere::Left, BrainRegion::Occipital); builder.add_node("O2", Hemisphere::Right, BrainRegion::Occipital); builder.add_edge(0, 1, 0.8); builder.add_edge(0, 2, 0.5); builder.add_edge(1, 3, 0.6); builder.add_edge(2, 3, 0.9); let mut graph = builder.build().unwrap(); for node in &mut graph.nodes { node.features = vec![0.5, 0.3, 0.2, 0.1]; } graph } #[test] fn test_dynamic_connectivity() { let channels = ["F3", "F4", "O1", "O2"]; let matrices: Vec>> = (0..5) .map(|t| { (0..4) .map(|i| { (0..4) .map(|j| if i == j { 1.0 } else { 0.5 + t as f64 * 0.05 }) .collect() }) .collect() }) .collect(); let dyn_conn = DynamicConnectivity::from_matrices(&matrices, &channels, 1.0, 0.5, 0.4).unwrap(); assert_eq!(dyn_conn.n_times(), 5); assert!(dyn_conn.get(2).is_some()); } #[test] fn test_temporal_attention() { let attn = TemporalAttention::new(16, 4, 50, true); let x = Array2::from_shape_fn((10, 16), |_| rand_simple()); let out = attn.forward(&x); assert_eq!(out.shape(), &[10, 16]); } #[test] fn test_temporal_brain_gnn() { // Create dynamic connectivity let channels = ["F3", "F4", "O1", "O2"]; let matrices: Vec>> = (0..5) .map(|_| { (0..4) .map(|i| (0..4).map(|j| if i == j { 1.0 } else { 0.6 }).collect()) .collect() }) .collect(); let dyn_conn = DynamicConnectivity::from_matrices(&matrices, &channels, 1.0, 0.5, 0.4).unwrap(); // Create model let model = TemporalBrainGNN::new(TemporalConfig { spatial: BrainNetCNNConfig { n_channels: 4, n_features: 1, // Will use degree hidden_dims: vec![8, 16], n_classes: 2, ..Default::default() }, temporal_heads: 2, temporal_hidden: 8, max_seq_len: 10, ..Default::default() }) .unwrap(); let logits = model.forward(&dyn_conn).unwrap(); assert_eq!(logits.len(), 2); let states = model.predict_states(&dyn_conn).unwrap(); assert_eq!(states.len(), 5); } }