//! Brain graph construction from connectivity matrices. //! //! This module provides the bridge between connectivity analysis (coherence, PLV, etc.) //! and graph neural networks. use crate::error::{GnnError, GnnResult}; use ndarray::Array2; use rtx_neuro_connectivity::ConnectivityResult2D; use serde::{Deserialize, Serialize}; /// Brain hemisphere #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum Hemisphere { /// Left hemisphere Left, /// Right hemisphere Right, /// Midline (Fz, Cz, Pz, etc.) Midline, /// Unknown or unspecified Unknown, } impl Hemisphere { /// Parse from channel name (e.g., "Fp1" -> Left, "Fp2" -> Right) pub fn from_channel_name(name: &str) -> Self { let name = name.trim(); if name.ends_with('z') || name.ends_with('Z') { Hemisphere::Midline } else if let Some(last) = name.chars().last() { if last.is_ascii_digit() { let digit = last.to_digit(10).unwrap_or(0); if digit % 2 == 1 { Hemisphere::Left } else { Hemisphere::Right } } else { Hemisphere::Unknown } } else { Hemisphere::Unknown } } } /// Brain region classification #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum BrainRegion { /// Frontal lobe (Fp, F) Frontal, /// Central (C) Central, /// Temporal (T) Temporal, /// Parietal (P) Parietal, /// Occipital (O) Occipital, /// Unknown region Unknown, } impl BrainRegion { /// Parse from channel name (e.g., "Fp1" -> Frontal, "O2" -> Occipital) pub fn from_channel_name(name: &str) -> Self { let name = name.to_uppercase(); if name.starts_with("FP") || name.starts_with("AF") { BrainRegion::Frontal } else if name.starts_with('F') { BrainRegion::Frontal } else if name.starts_with('C') { BrainRegion::Central } else if name.starts_with('T') { BrainRegion::Temporal } else if name.starts_with('P') && !name.starts_with("PO") { BrainRegion::Parietal } else if name.starts_with('O') || name.starts_with("PO") { BrainRegion::Occipital } else { BrainRegion::Unknown } } /// Get approximate 3D coordinates for visualization (normalized) pub fn approx_position(&self) -> [f32; 3] { match self { BrainRegion::Frontal => [0.0, 0.8, 0.3], BrainRegion::Central => [0.0, 0.0, 0.5], BrainRegion::Temporal => [0.9, 0.0, 0.0], BrainRegion::Parietal => [0.0, -0.5, 0.5], BrainRegion::Occipital => [0.0, -0.9, 0.1], BrainRegion::Unknown => [0.0, 0.0, 0.0], } } } /// A node in the brain graph (represents a channel/electrode) #[derive(Debug, Clone, Serialize, Deserialize)] pub struct BrainNode { /// Node index pub index: usize, /// Channel name (e.g., "Fp1", "O2") pub name: String, /// Hemisphere pub hemisphere: Hemisphere, /// Brain region pub region: BrainRegion, /// 3D position [x, y, z] in normalized coordinates pub position: [f32; 3], /// Node features (e.g., power spectrum, time series features) pub features: Vec, } impl BrainNode { /// Create a new brain node pub fn new(index: usize, name: &str) -> Self { Self { index, name: name.to_string(), hemisphere: Hemisphere::from_channel_name(name), region: BrainRegion::from_channel_name(name), position: [0.0, 0.0, 0.0], features: Vec::new(), } } /// Create with explicit hemisphere and region pub fn with_location( index: usize, name: &str, hemisphere: Hemisphere, region: BrainRegion, ) -> Self { let mut node = Self::new(index, name); node.hemisphere = hemisphere; node.region = region; // Compute approximate position from region and hemisphere let base_pos = region.approx_position(); node.position = match hemisphere { Hemisphere::Left => [-0.5 + base_pos[0], base_pos[1], base_pos[2]], Hemisphere::Right => [0.5 + base_pos[0], base_pos[1], base_pos[2]], Hemisphere::Midline => [0.0, base_pos[1], base_pos[2]], Hemisphere::Unknown => base_pos, }; node } /// Set node features pub fn with_features(mut self, features: Vec) -> Self { self.features = features; self } } /// An edge in the brain graph (represents connectivity between channels) #[derive(Debug, Clone, Serialize, Deserialize)] pub struct BrainEdge { /// Source node index pub source: usize, /// Target node index pub target: usize, /// Connectivity weight (e.g., coherence, PLV) pub weight: f64, /// Edge features (e.g., multi-frequency connectivity) pub features: Vec, /// Whether this is an interhemispheric connection pub interhemispheric: bool, } impl BrainEdge { /// Create a new brain edge pub fn new(source: usize, target: usize, weight: f64) -> Self { Self { source, target, weight, features: vec![weight], interhemispheric: false, } } /// Add multi-frequency features pub fn with_frequency_features(mut self, features: Vec) -> Self { self.features = features; self } /// Mark as interhemispheric pub fn with_interhemispheric(mut self, flag: bool) -> Self { self.interhemispheric = flag; self } } /// A brain connectivity graph #[derive(Debug, Clone, Serialize, Deserialize)] pub struct BrainGraph { /// Nodes (channels/electrodes) pub nodes: Vec, /// Edges (connectivity) pub edges: Vec, /// Adjacency matrix (dense) adjacency: Option>, /// Metadata pub metadata: GraphMetadata, } /// Graph metadata #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct GraphMetadata { /// Subject ID pub subject_id: Option, /// Condition/task pub condition: Option, /// Time window [start, end] in seconds pub time_window: Option<[f64; 2]>, /// Frequency band [fmin, fmax] in Hz pub freq_band: Option<[f64; 2]>, /// Connectivity method used pub connectivity_method: Option, /// Label for classification pub label: Option, /// Label name (e.g., "Control", "Patient") pub label_name: Option, } impl BrainGraph { /// Create an empty brain graph pub fn new() -> Self { Self { nodes: Vec::new(), edges: Vec::new(), adjacency: None, metadata: GraphMetadata::default(), } } /// Number of nodes pub fn n_nodes(&self) -> usize { self.nodes.len() } /// Number of edges pub fn n_edges(&self) -> usize { self.edges.len() } /// Create from a connectivity matrix pub fn from_connectivity_matrix( matrix: &[Vec], channel_names: &[&str], threshold: f64, ) -> GnnResult { let n = matrix.len(); if n == 0 { return Err(GnnError::InvalidGraph("Empty connectivity matrix".into())); } if channel_names.len() != n { return Err(GnnError::DimensionMismatch(format!( "Matrix size {} != channel names {}", n, channel_names.len() ))); } let mut graph = BrainGraph::new(); // Add nodes for (i, name) in channel_names.iter().enumerate() { graph.nodes.push(BrainNode::new(i, name)); } // Add edges (above threshold) for i in 0..n { for j in 0..n { if i != j && matrix[i][j] >= threshold { let is_inter = graph.nodes[i].hemisphere != graph.nodes[j].hemisphere && graph.nodes[i].hemisphere != Hemisphere::Midline && graph.nodes[j].hemisphere != Hemisphere::Midline; graph .edges .push(BrainEdge::new(i, j, matrix[i][j]).with_interhemispheric(is_inter)); } } } // Store adjacency let mut adj = Array2::zeros((n, n)); for edge in &graph.edges { adj[[edge.source, edge.target]] = edge.weight; } graph.adjacency = Some(adj); Ok(graph) } /// Create from ConnectivityResult2D at a specific frequency pub fn from_connectivity_result( result: &ConnectivityResult2D, channel_names: &[&str], freq_idx: usize, threshold: f64, ) -> GnnResult { let n_channels = channel_names.len(); let matrix = result.to_matrix(freq_idx, n_channels); let mut graph = Self::from_connectivity_matrix(&matrix, channel_names, threshold)?; // Add metadata if let Some(&freq) = result.freqs.get(freq_idx) { graph.metadata.freq_band = Some([freq, freq]); } graph.metadata.connectivity_method = Some(format!("{:?}", result.method)); Ok(graph) } /// Create multi-frequency graph with edge features pub fn from_connectivity_multifreq( result: &ConnectivityResult2D, channel_names: &[&str], freq_indices: &[usize], threshold: f64, ) -> GnnResult { let n_channels = channel_names.len(); // Start with first frequency let matrix0 = result.to_matrix(freq_indices[0], n_channels); let mut graph = Self::from_connectivity_matrix(&matrix0, channel_names, threshold)?; // Add multi-frequency features to edges for edge in &mut graph.edges { let mut features = Vec::with_capacity(freq_indices.len()); for &freq_idx in freq_indices { let matrix = result.to_matrix(freq_idx, n_channels); features.push(matrix[edge.source][edge.target]); } edge.features = features; } // Update metadata if !freq_indices.is_empty() { let fmin = result.freqs.get(freq_indices[0]).copied().unwrap_or(0.0); let fmax = result .freqs .get(*freq_indices.last().unwrap()) .copied() .unwrap_or(0.0); graph.metadata.freq_band = Some([fmin, fmax]); } Ok(graph) } /// Get adjacency matrix pub fn adjacency_matrix(&self) -> Array2 { if let Some(ref adj) = self.adjacency { adj.clone() } else { let n = self.n_nodes(); let mut adj = Array2::zeros((n, n)); for edge in &self.edges { adj[[edge.source, edge.target]] = edge.weight; } adj } } /// Get degree of each node pub fn degrees(&self) -> Vec { let adj = self.adjacency_matrix(); adj.sum_axis(ndarray::Axis(1)).to_vec() } /// Get node feature matrix [n_nodes, n_features] pub fn node_feature_matrix(&self) -> GnnResult> { if self.nodes.is_empty() { return Err(GnnError::InvalidGraph("No nodes".into())); } let n_features = self.nodes[0].features.len(); if n_features == 0 { // Use degree as default feature let degrees = self.degrees(); let mut features = Array2::zeros((self.n_nodes(), 1)); for (i, &d) in degrees.iter().enumerate() { features[[i, 0]] = d; } return Ok(features); } let mut features = Array2::zeros((self.n_nodes(), n_features)); for (i, node) in self.nodes.iter().enumerate() { if node.features.len() != n_features { return Err(GnnError::DimensionMismatch(format!( "Node {} has {} features, expected {}", i, node.features.len(), n_features ))); } for (j, &f) in node.features.iter().enumerate() { features[[i, j]] = f; } } Ok(features) } /// Get edge index tensor [2, n_edges] for sparse GNN pub fn edge_index(&self) -> (Vec, Vec) { let sources: Vec = self.edges.iter().map(|e| e.source).collect(); let targets: Vec = self.edges.iter().map(|e| e.target).collect(); (sources, targets) } /// Get edge weights pub fn edge_weights(&self) -> Vec { self.edges.iter().map(|e| e.weight).collect() } /// Get edge feature matrix [n_edges, n_features] pub fn edge_feature_matrix(&self) -> Array2 { if self.edges.is_empty() { return Array2::zeros((0, 0)); } let n_features = self.edges[0].features.len(); let mut features = Array2::zeros((self.n_edges(), n_features)); for (i, edge) in self.edges.iter().enumerate() { for (j, &f) in edge.features.iter().enumerate() { features[[i, j]] = f; } } features } /// Set node features from power spectrum pub fn set_node_features_from_psd(&mut self, psd: &[Vec]) -> GnnResult<()> { if psd.len() != self.n_nodes() { return Err(GnnError::DimensionMismatch(format!( "PSD has {} channels, graph has {} nodes", psd.len(), self.n_nodes() ))); } for (node, features) in self.nodes.iter_mut().zip(psd.iter()) { node.features = features.clone(); } Ok(()) } /// Get subgraph for a specific brain region pub fn subgraph_by_region(&self, region: BrainRegion) -> BrainGraph { let node_indices: Vec = self .nodes .iter() .filter(|n| n.region == region) .map(|n| n.index) .collect(); self.subgraph(&node_indices) } /// Get subgraph for a specific hemisphere pub fn subgraph_by_hemisphere(&self, hemisphere: Hemisphere) -> BrainGraph { let node_indices: Vec = self .nodes .iter() .filter(|n| n.hemisphere == hemisphere) .map(|n| n.index) .collect(); self.subgraph(&node_indices) } /// Create a subgraph with only specified nodes pub fn subgraph(&self, node_indices: &[usize]) -> BrainGraph { let node_set: std::collections::HashSet = node_indices.iter().copied().collect(); let index_map: std::collections::HashMap = node_indices .iter() .enumerate() .map(|(new, &old)| (old, new)) .collect(); let nodes: Vec = node_indices .iter() .enumerate() .map(|(new_idx, &old_idx)| { let mut node = self.nodes[old_idx].clone(); node.index = new_idx; node }) .collect(); let edges: Vec = self .edges .iter() .filter(|e| node_set.contains(&e.source) && node_set.contains(&e.target)) .map(|e| BrainEdge { source: *index_map.get(&e.source).unwrap(), target: *index_map.get(&e.target).unwrap(), ..e.clone() }) .collect(); BrainGraph { nodes, edges, adjacency: None, metadata: self.metadata.clone(), } } /// Convert to rtx-geom Graph for GNN processing pub fn to_geom_graph(&self) -> GnnResult { use rtx_tensor::{Device, Tensor}; let device = Device::default(); let mut graph = rtx_geom::Graph::new(); // Add nodes and track their IDs let mut node_ids = Vec::with_capacity(self.nodes.len()); for node in &self.nodes { let features = if node.features.is_empty() { // Default: one-hot region encoding let mut feat = vec![0.0f32; 5]; match node.region { BrainRegion::Frontal => feat[0] = 1.0, BrainRegion::Central => feat[1] = 1.0, BrainRegion::Temporal => feat[2] = 1.0, BrainRegion::Parietal => feat[3] = 1.0, BrainRegion::Occipital => feat[4] = 1.0, BrainRegion::Unknown => {} } feat } else { node.features.iter().map(|&x| x as f32).collect() }; let tensor = Tensor::from_data(features.clone(), [features.len()], &device) .map_err(|e| GnnError::TensorError(e.to_string()))?; let node_id = graph.add_node(rtx_geom::Node::new(tensor)); node_ids.push(node_id); } // Add edges using tracked node IDs for edge in &self.edges { let features: Vec = edge.features.iter().map(|&x| x as f32).collect(); let tensor = Tensor::from_data(features.clone(), [features.len()], &device) .map_err(|e| GnnError::TensorError(e.to_string()))?; let source_id = node_ids[edge.source]; let target_id = node_ids[edge.target]; graph .add_edge(source_id, target_id, rtx_geom::Edge::new(tensor)) .map_err(|e| GnnError::GeomError(e.to_string()))?; } Ok(graph) } } impl Default for BrainGraph { fn default() -> Self { Self::new() } } /// Builder for constructing brain graphs programmatically pub struct GraphBuilder { nodes: Vec, edges: Vec, metadata: GraphMetadata, } impl GraphBuilder { /// Create a new graph builder pub fn new() -> Self { Self { nodes: Vec::new(), edges: Vec::new(), metadata: GraphMetadata::default(), } } /// Add a node pub fn add_node(&mut self, name: &str, hemisphere: Hemisphere, region: BrainRegion) -> usize { let index = self.nodes.len(); self.nodes .push(BrainNode::with_location(index, name, hemisphere, region)); index } /// Add a node with auto-detected hemisphere and region pub fn add_channel(&mut self, name: &str) -> usize { let index = self.nodes.len(); self.nodes.push(BrainNode::new(index, name)); index } /// Add an edge pub fn add_edge(&mut self, source: usize, target: usize, weight: f64) { self.edges.push(BrainEdge::new(source, target, weight)); } /// Set metadata pub fn with_metadata(mut self, metadata: GraphMetadata) -> Self { self.metadata = metadata; self } /// Build the graph pub fn build(self) -> GnnResult { // Validate edges let n = self.nodes.len(); for edge in &self.edges { if edge.source >= n || edge.target >= n { return Err(GnnError::InvalidGraph(format!( "Edge {}->{} references invalid node (n_nodes={})", edge.source, edge.target, n ))); } } // Mark interhemispheric edges let mut edges = self.edges; for edge in &mut edges { edge.interhemispheric = self.nodes[edge.source].hemisphere != self.nodes[edge.target].hemisphere && self.nodes[edge.source].hemisphere != Hemisphere::Midline && self.nodes[edge.target].hemisphere != Hemisphere::Midline; } Ok(BrainGraph { nodes: self.nodes, edges, adjacency: None, metadata: self.metadata, }) } } impl Default for GraphBuilder { fn default() -> Self { Self::new() } } /// Create a standard 10-20 system graph pub fn standard_10_20_graph() -> BrainGraph { let mut builder = GraphBuilder::new(); // Add standard 10-20 channels let channels = [ "Fp1", "Fp2", "F7", "F3", "Fz", "F4", "F8", "T3", "C3", "Cz", "C4", "T4", "T5", "P3", "Pz", "P4", "T6", "O1", "O2", ]; for name in &channels { builder.add_channel(name); } // Add full connectivity (will be thresholded later) let n = channels.len(); for i in 0..n { for j in (i + 1)..n { // Default weight based on distance builder.add_edge(i, j, 0.5); builder.add_edge(j, i, 0.5); } } builder.build().unwrap() } #[cfg(test)] mod tests { use super::*; #[test] fn test_hemisphere_detection() { assert_eq!(Hemisphere::from_channel_name("Fp1"), Hemisphere::Left); assert_eq!(Hemisphere::from_channel_name("Fp2"), Hemisphere::Right); assert_eq!(Hemisphere::from_channel_name("Fz"), Hemisphere::Midline); assert_eq!(Hemisphere::from_channel_name("Cz"), Hemisphere::Midline); assert_eq!(Hemisphere::from_channel_name("O1"), Hemisphere::Left); assert_eq!(Hemisphere::from_channel_name("O2"), Hemisphere::Right); } #[test] fn test_region_detection() { assert_eq!(BrainRegion::from_channel_name("Fp1"), BrainRegion::Frontal); assert_eq!(BrainRegion::from_channel_name("F3"), BrainRegion::Frontal); assert_eq!(BrainRegion::from_channel_name("C3"), BrainRegion::Central); assert_eq!(BrainRegion::from_channel_name("T3"), BrainRegion::Temporal); assert_eq!(BrainRegion::from_channel_name("P3"), BrainRegion::Parietal); assert_eq!(BrainRegion::from_channel_name("O1"), BrainRegion::Occipital); } #[test] fn test_graph_builder() { let mut builder = GraphBuilder::new(); builder.add_channel("Fp1"); builder.add_channel("Fp2"); builder.add_edge(0, 1, 0.8); let graph = builder.build().unwrap(); assert_eq!(graph.n_nodes(), 2); assert_eq!(graph.n_edges(), 1); assert!(graph.edges[0].interhemispheric); } #[test] fn test_subgraph() { 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(2, 3, 0.9); let graph = builder.build().unwrap(); // Left hemisphere subgraph let left = graph.subgraph_by_hemisphere(Hemisphere::Left); assert_eq!(left.n_nodes(), 2); assert_eq!(left.n_edges(), 1); // F3->O1 // Frontal subgraph let frontal = graph.subgraph_by_region(BrainRegion::Frontal); assert_eq!(frontal.n_nodes(), 2); assert_eq!(frontal.n_edges(), 1); // F3->F4 } #[test] fn test_standard_10_20() { let graph = standard_10_20_graph(); assert_eq!(graph.n_nodes(), 19); assert!(graph.n_edges() > 0); } }