Files
rustytorch/crates/specialized/rtx-neuro-gnn/src/lib.rs
T
2026-03-04 00:08:42 +00:00

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);
}
}