116 lines
3.6 KiB
Rust
116 lines
3.6 KiB
Rust
//! # rtx-neuro-gnn
|
|
//!
|
|
//! Brain Connectome Graph Neural Networks for MEG/EEG analysis.
|
|
//!
|
|
//! This crate provides GPU-accelerated GNN models for brain connectivity analysis,
|
|
//! disorder classification, brain state prediction, and biomarker discovery.
|
|
//!
|
|
//! ## Features
|
|
//!
|
|
//! - **Graph Construction**: Convert connectivity matrices to graph representations
|
|
//! - **Brain GNN Models**: BrainNetCNN, GraphTransformer, and custom architectures
|
|
//! - **Temporal GNNs**: Dynamic connectivity analysis with temporal attention
|
|
//! - **Explainability**: Edge importance attribution and biomarker discovery
|
|
//!
|
|
//! ## Applications
|
|
//!
|
|
//! | Task | Input | Output |
|
|
//! |------|-------|--------|
|
|
//! | Disorder classification | Resting-state connectivity | Autism/Schizophrenia/Control |
|
|
//! | Brain state prediction | Dynamic connectivity | Sleep stage / Attention level |
|
|
//! | Biomarker discovery | Multi-frequency graphs | Important connections |
|
|
//!
|
|
//! ## Example
|
|
//!
|
|
//! ```ignore
|
|
//! use rtx_neuro_gnn::{BrainGraph, BrainNetCNN, BrainGNN};
|
|
//!
|
|
//! // Convert connectivity matrix to graph
|
|
//! let graph = BrainGraph::from_connectivity_matrix(&conn_matrix, &channel_names)?;
|
|
//!
|
|
//! // Create classifier
|
|
//! let mut model = BrainNetCNN::new(BrainNetCNNConfig {
|
|
//! n_channels: 64,
|
|
//! n_classes: 3,
|
|
//! ..Default::default()
|
|
//! })?;
|
|
//!
|
|
//! // Forward pass
|
|
//! let logits = model.forward(&graph)?;
|
|
//! ```
|
|
|
|
#![warn(missing_docs)]
|
|
|
|
pub mod error;
|
|
pub mod explain;
|
|
pub mod graph;
|
|
pub mod layers;
|
|
pub mod models;
|
|
pub mod temporal;
|
|
|
|
// Re-export main types
|
|
pub use error::{GnnError, GnnResult};
|
|
pub use explain::{
|
|
EdgeImportance, ExplainerConfig, GnnExplainer, InterpretationMethod, NodeImportance,
|
|
};
|
|
pub use graph::{BrainEdge, BrainGraph, BrainNode, BrainRegion, GraphBuilder, Hemisphere};
|
|
pub use layers::{
|
|
BrainAttention, BrainAttentionConfig, BrainConv, BrainConvConfig, BrainPool, BrainPoolConfig,
|
|
EdgeConv,
|
|
};
|
|
pub use models::{
|
|
BrainGAT, BrainGATConfig, BrainGNN, BrainNetCNN, BrainNetCNNConfig, BrainTransformer,
|
|
BrainTransformerConfig,
|
|
};
|
|
pub use temporal::{
|
|
DynamicConnectivity, TemporalAttention, TemporalBrainGNN, TemporalConfig, TimeWindow,
|
|
};
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_basic_graph_creation() {
|
|
// Create a simple 4-node brain graph
|
|
let mut builder = GraphBuilder::new();
|
|
builder.add_node("Fp1", Hemisphere::Left, BrainRegion::Frontal);
|
|
builder.add_node("Fp2", Hemisphere::Right, BrainRegion::Frontal);
|
|
builder.add_node("O1", Hemisphere::Left, BrainRegion::Occipital);
|
|
builder.add_node("O2", Hemisphere::Right, BrainRegion::Occipital);
|
|
|
|
// Add connectivity edges
|
|
builder.add_edge(0, 1, 0.8);
|
|
builder.add_edge(0, 2, 0.3);
|
|
builder.add_edge(1, 3, 0.4);
|
|
builder.add_edge(2, 3, 0.9);
|
|
|
|
let graph = builder.build().unwrap();
|
|
assert_eq!(graph.n_nodes(), 4);
|
|
assert_eq!(graph.n_edges(), 4);
|
|
}
|
|
|
|
#[test]
|
|
fn test_from_connectivity_matrix() {
|
|
let matrix = vec![
|
|
vec![1.0, 0.8, 0.3, 0.1],
|
|
vec![0.8, 1.0, 0.2, 0.4],
|
|
vec![0.3, 0.2, 1.0, 0.9],
|
|
vec![0.1, 0.4, 0.9, 1.0],
|
|
];
|
|
let channel_names = vec!["Fp1", "Fp2", "O1", "O2"];
|
|
|
|
let graph = BrainGraph::from_connectivity_matrix(
|
|
&matrix,
|
|
&channel_names,
|
|
0.5, // threshold
|
|
)
|
|
.unwrap();
|
|
|
|
// Only edges with weight >= 0.5 should be included
|
|
assert_eq!(graph.n_nodes(), 4);
|
|
// Edges: Fp1-Fp2 (0.8), O1-O2 (0.9) - both directions = 4 edges
|
|
assert!(graph.n_edges() >= 2);
|
|
}
|
|
}
|