| name | connectome-constrained-neural-network |
| description | Connectome-Constrained Neural Network (CCNN) methodology for brain-inspired AI. Integrates biological structural connectivity (connectome) into artificial neural network architectures to improve generalization and biological plausibility. Activation: connectome constraint, structural connectivity, brain-inspired architecture, connectome-based AI, wiring cost, brain network prior, diffusion MRI connectivity. |
| tags | ["connectome","brain-networks","structural-connectivity","wiring-cost","brain-inspired","DWI","diffusion-mri","network-constraint"] |
Connectome-Constrained Neural Networks
Overview
Connectome-Constrained Neural Networks (CCNNs) integrate biological brain connectivity patterns into artificial neural network architectures. The core insight is that the brain's wiring diagram (connectome) has been optimized through evolution for efficient computation, and these structural constraints can improve AI systems.
Core Concepts
Why Connectome Constraints?
Traditional Neural Network:
- Fully connected or simple local patterns
- No wiring cost consideration
- Prone to overfitting
- Biologically implausible
Connectome-Constrained Network:
- Structured by biological connectivity
- Wiring efficiency built-in
- Better generalization
- Biologically grounded
Types of Connectome Constraints
- Structural Connectivity: Physical wiring from diffusion MRI
- Functional Connectivity: Correlation-based connections from fMRI/EEG
- Wiring Cost: Metabolic/physical costs of connections
- Modular Organization: Community structure of brain networks
Implementation
Loading Brain Connectome Data
import numpy as np
import scipy.io as sio
from nilearn import datasets
def load_human_connectome(atlas='aal', n_regions=116):
"""
Load human structural connectivity matrix.
Args:
atlas: 'aal', 'desikan', 'destrieux', 'harvard_oxford'
n_regions: Number of brain regions
Returns:
connectivity: (n_regions, n_regions) connectivity matrix
region_labels: List of region names
coordinates: 3D coordinates of regions
"""
if atlas == 'aal':
atlas_data = datasets.fetch_atlas_aal()
elif atlas == 'desikan':
atlas_data = datasets.fetch_atlas_destrieux_2009()
elif atlas == 'harvard_oxford':
atlas_data = datasets.fetch_atlas_harvard_oxford('cort-maxprob-thr25-2mm')
connectivity = load_diffusion_connectivity(atlas_data)
return connectivity, atlas_data.labels, atlas_data.coordinates
def load_diffusion_connectivity(atlas_data, dataset='hcp'):
"""
Load structural connectivity from diffusion MRI.
Args:
atlas_data: Atlas information
dataset: 'hcp', 'ukb' (Human Connectome Project, UK Biobank)
Returns:
SC: Structural connectivity matrix (streamline counts)
"""
if dataset == 'hcp':
SC = fetch_hcp_connectome(atlas_data)
SC = SC / SC.max()
SC[SC < ] =
SC
Connectome-Constrained Layer
import torch
import torch.nn as nn
import torch.nn.functional as F
class ConnectomeLinear(nn.Module):
"""
Linear layer with connectome-inspired sparse connectivity.
Implements constrained weight matrix where connections follow
biological connectivity patterns.
"""
def __init__(self, in_features, out_features, connectivity_matrix,
constraint_type='hard', sparsity_target=0.1):
"""
Args:
in_features: Input dimension
out_features: Output dimension
connectivity_matrix: Binary or weighted connectivity (n_out, n_in)
constraint_type: 'hard' (fixed), 'soft' (regularized), 'init' (initialization only)
sparsity_target: Target connection density
"""
super().__init__()
self.in_features = in_features
self.out_features = out_features
self.constraint_type = constraint_type
if connectivity_matrix is not None:
connectivity = self._resize_connectivity(
connectivity_matrix,
(out_features, in_features)
)
self.register_buffer('connectivity_mask',
torch.tensor(connectivity, dtype=torch.float32))
else:
self.connectivity_mask = torch.rand(out_features, in_features) < sparsity_target
self.connectivity_mask = .connectivity_mask.()
constraint_type == :
n_connections = (.connectivity_mask.().item())
.weight_values = nn.Parameter(torch.randn(n_connections) * )
.register_buffer(, ._create_sparse_indices())
:
.weight = nn.Parameter(torch.randn(out_features, in_features) * )
.bias = nn.Parameter(torch.zeros(out_features))
constraint_type == :
.wiring_cost = ._compute_wiring_cost()
():
scipy.ndimage zoom
connectivity.shape == target_shape:
connectivity
zoom_factors = (target_shape[] / connectivity.shape[],
target_shape[] / connectivity.shape[])
resized = zoom(connectivity, zoom_factors, order=)
(resized > np.percentile(resized, )).astype()
():
rows, cols = torch.where(.connectivity_mask > )
torch.stack([rows, cols], dim=)
():
coordinates :
coords_out = torch.randn(.out_features, )
coords_in = torch.randn(.in_features, )
:
coords_out = coordinates[:.out_features]
coords_in = coordinates[:.in_features]
distances = torch.cdist(coords_out, coords_in)
distances / distances.()
():
.constraint_type == :
weight = torch.zeros(.out_features, .in_features,
device=x.device)
weight[.weight_indices[], .weight_indices[]] = .weight_values
:
weight = .weight
.constraint_type == :
weight = weight * ( + * .connectivity_mask)
.constraint_type == :
F.linear(x, weight, .bias)
():
.constraint_type != :
cost = (.weight.() * .wiring_cost.to(.weight.device)).()
lambda_wiring * cost
():
(
)
Full Connectome-Constrained Network
class ConnectomeConstrainedNN(nn.Module):
"""
Full neural network with connectome constraints at multiple layers.
Architecture inspired by cortical hierarchy:
- Early layers: Sensory-like (local connectivity)
- Middle layers: Association (long-range connectivity)
- Late layers: Output (task-specific)
"""
def __init__(self, input_dim, output_dim, hidden_dims=[512, 256, 128],
connectome_data=None, constraint_layers=[1, 2]):
"""
Args:
input_dim: Input feature dimension
output_dim: Output dimension (classes)
hidden_dims: List of hidden layer dimensions
connectome_data: Dictionary with connectivity matrices
constraint_layers: Which layers to apply connectome constraints
"""
super().__init__()
self.layers = nn.ModuleList()
dims = [input_dim] + hidden_dims + [output_dim]
for i in range(len(dims) - 1):
if i in constraint_layers and connectome_data is not None:
layer = ConnectomeLinear(
dims[i], dims[i+1],
connectivity_matrix=connectome_data.get(f'layer_{i}'),
constraint_type='soft'
)
else:
layer = nn.Linear(dims[i], dims[i+1])
self.layers.append(layer)
i < (dims) - :
.layers.append(nn.ReLU())
.layers.append(nn.Dropout())
():
layer .layers:
x = layer(x)
x
():
total_cost =
layer .layers:
(layer, ConnectomeLinear):
total_cost += layer.get_wiring_cost_loss()
total_cost
Connectome-Based Graph Neural Networks
import torch_geometric as pyg
from torch_geometric.nn import GCNConv, GATConv
class ConnectomeGNN(nn.Module):
"""
Graph Neural Network using brain connectome as graph structure.
Nodes = Brain regions
Edges = Structural connectivity
Node features = Neural activity or ROI features
"""
def __init__(self, n_regions, feature_dim, hidden_dim=64, output_dim=10,
connectome_edge_index=None, connectome_weights=None):
"""
Args:
n_regions: Number of brain regions (nodes)
feature_dim: Dimension of node features
hidden_dim: Hidden layer dimension
output_dim: Output classes
connectome_edge_index: (2, n_edges) connectivity
connectome_weights: Edge weights from connectome
"""
super().__init__()
self.n_regions = n_regions
if connectome_edge_index is None:
self.edge_index = self._connectome_to_edges(n_regions)
else:
self.register_buffer('edge_index', connectome_edge_index)
if connectome_weights is not None:
self.register_buffer('edge_weights', connectome_weights)
else:
self.edge_weights = None
.conv1 = GCNConv(feature_dim, hidden_dim)
.conv2 = GCNConv(hidden_dim, hidden_dim)
.conv3 = GCNConv(hidden_dim, hidden_dim)
.global_pool = nn.AdaptiveAvgPool1d()
.classifier = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim // ),
nn.ReLU(),
nn.Linear(hidden_dim // , output_dim)
)
():
rows, cols = np.where(connectome_matrix > )
edge_index = torch.tensor(np.stack([rows, cols]), dtype=torch.long)
edge_index
():
x = .conv1(x, .edge_index, .edge_weights)
x = F.relu(x)
x = F.dropout(x, p=, training=.training)
x = .conv2(x, .edge_index, .edge_weights)
x = F.relu(x)
x = .conv3(x, .edge_index, .edge_weights)
batch :
x = x.mean(dim=)
:
x = pyg.nn.global_mean_pool(x, batch)
out = .classifier(x)
out
Wiring Cost Optimization
class WiringCostOptimizer:
"""
Optimize network connectivity for minimal wiring cost
while maintaining performance.
"""
def __init__(self, network, coordinates, lambda_wiring=0.001):
"""
Args:
network: Neural network to optimize
coordinates: 3D coordinates of neurons/regions
lambda_wiring: Weight of wiring cost in loss
"""
self.network = network
self.coordinates = coordinates
self.lambda_wiring = lambda_wiring
self.distances = self._compute_all_distances()
def _compute_all_distances(self):
"""Compute pairwise distances between all neurons."""
coords = torch.tensor(self.coordinates)
return torch.cdist(coords, coords)
def compute_wiring_loss(self):
"""
Compute total wiring cost as weighted sum of connection distances.
"""
total_cost = 0.0
for layer in self.network.modules():
if isinstance(layer, nn.Linear):
weights = layer.weight.abs()
n_out, n_in = weights.shape
dist_sample = self._sample_distances(n_out, n_in)
cost = (weights * dist_sample.to(weights.device)).sum()
total_cost += cost
.lambda_wiring * total_cost
():
indices_out = torch.randint(, (.coordinates), (n_out,))
indices_in = torch.randint(, (.coordinates), (n_in,))
.distances[indices_out][:, indices_in]
():
pruned_network = copy.deepcopy(.network)
layer pruned_network.modules():
(layer, nn.Linear):
weights = layer.weight.data
n_out, n_in = weights.shape
dist_sample = ._sample_distances(n_out, n_in)
costs = weights.() * dist_sample
threshold = torch.quantile(costs.flatten(), - pruning_ratio)
mask = costs < threshold
weights *= mask.().to(weights.device)
pruned_network
Applications
1. Brain Age Prediction
def brain_age_prediction(connectomes, ages, train_idx, test_idx):
"""
Predict biological age from structural connectome.
Uses connectome-constrained GNN to learn age-related connectivity changes.
"""
data = prepare_connectome_data(connectomes, ages)
model = ConnectomeGNN(
n_regions=connectomes.shape[1],
feature_dim=connectomes.shape[0],
hidden_dim=128,
output_dim=1,
connectome_edge_index=data.edge_index,
connectome_weights=data.edge_weights
)
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
criterion = nn.MSELoss()
for epoch in range(100):
model.train()
optimizer.zero_grad()
pred = model(data.x, data.batch)
loss = criterion(pred[train_idx], data.y[train_idx])
loss.backward()
optimizer.step()
model.eval()
with torch.no_grad():
pred_age = model(data.x, data.batch)
mae = (pred_age[test_idx] - data.y[test_idx]).abs().mean()
return model, mae
2. Disease Classification
def disease_classification(connectomes, labels, disease='Alzheimer'):
"""
Classify neurological disease from connectome alterations.
Uses connectome constraints to focus on biologically plausible patterns.
"""
model = ConnectomeConstrainedNN(
input_dim=connectomes.shape[1],
output_dim=2,
hidden_dims=[256, 128],
connectome_data={'layer_1': connectomes.mean(axis=0)},
constraint_layers=[1]
)
optimizer = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-5)
criterion = nn.CrossEntropyLoss()
for epoch in range(50):
optimizer.zero_grad()
logits = model(connectomes)
class_loss = criterion(logits, labels)
wiring_loss = model.get_wiring_cost()
total_loss = class_loss + wiring_loss
total_loss.backward()
optimizer.step()
return model
3. Transfer Learning with Connectome Priors
def connectome_transfer_learning(source_connectomes, source_labels,
target_connectomes, target_labels,
n_finetune_regions=10):
"""
Transfer learning using connectome structure as prior.
Source: Large dataset (e.g., HCP)
Target: Small disease dataset
"""
source_model = ConnectomeGNN(
n_regions=source_connectomes.shape[1],
feature_dim=source_connectomes.shape[2],
connectome_edge_index=source_connectomes.edge_index
)
train_model(source_model, source_connectomes, source_labels)
target_model = copy.deepcopy(source_model)
affected_regions = identify_altered_regions(
source_connectomes, target_connectomes
)
freeze_non_disease_regions(target_model, affected_regions)
train_model(target_model, target_connectomes, target_labels)
return target_model
Evaluation Metrics
def evaluate_connectome_constraints(model, test_data):
"""
Evaluate the effect of connectome constraints.
Metrics:
1. Performance (accuracy, etc.)
2. Efficiency (sparsity, wiring cost)
3. Biological plausibility
"""
results = {}
model.eval()
with torch.no_grad():
predictions = model(test_data.x)
results['accuracy'] = compute_accuracy(predictions, test_data.y)
total_params = sum(p.numel() for p in model.parameters())
results['total_parameters'] = total_params
nonzero = 0
for layer in model.modules():
if isinstance(layer, (nn.Linear, ConnectomeLinear)):
nonzero += (layer.weight.abs() > 0.01).sum().item()
results['nonzero_connections'] = nonzero
results['sparsity'] = nonzero / total_params
wiring_cost = 0.0
if hasattr(model, 'get_wiring_cost'):
wiring_cost = model.get_wiring_cost().item()
results['wiring_cost'] = wiring_cost
if hasattr(model, 'extract_connectivity'):
learned_conn = model.extract_connectivity()
biological_conn = test_data.connectome
results['connectome_correlation'] = compute_connectivity_correlation(
learned_conn, biological_conn
)
results
References
- Bettens, D., et al. (2024). Connectome-constrained deep learning improves prediction accuracy and reveals whole-brain dynamics. bioRxiv.
- Brünn, A., et al. (2024). Brain-inspired learning in artificial neural networks with neuroanatomical connectomes. bioRxiv.
- Oldham, S., et al. (2022). Connectome smoothing via low-dimensional structural embeddings. NeuroImage.
- Sarwar, T., et al. (2021). Connectome-based prediction of brain network response to targeted stimulation. PNAS.
- Shafiei, G., et al. (2020). Spatiotemporal dynamics of functional connectivity in the human brain. NeuroImage.
Related Skills
brain-graph-neural: GNN methods for brain connectivity
functional-connectivity-graph-neural-networks: Functional connectivity analysis
structural-functional-brain-gnn: Structural-functional fusion
geometric-brain-dynamics-mapping: Geometric methods for brain dynamics
Activation Keywords
- connectome constraint
- structural connectivity
- wiring cost
- brain network prior
- connectome-based AI
- diffusion MRI connectivity
- brain-inspired architecture
- connectome constrained neural network
- wiring efficiency
- brain graph neural network