| name | neuroaps-net-alzheimer-point-cloud |
| description | Neuro-Anatomically Aware Point Cloud Representation (NeuroAPS-Net) for efficient Alzheimer's disease classification from MRI. Converts T1-weighted MRI into anatomically-informed 2D point clouds with region-aware feature encoding. Activation triggers: Alzheimer's classification, neuroanatomical point cloud, MRI analysis, geometric deep learning. |
NeuroAPS-Net: Neuro-Anatomically Aware Point Cloud Representation for Alzheimer's Disease Classification
A lightweight geometric deep learning model that converts T1-weighted MRI into anatomically-informed 2D point clouds for efficient and interpretable Alzheimer's disease classification.
Metadata
- Source: arXiv:2604.22883v1
- Authors: Towhidul Islam, Mufti Mahmud
- Published: 2026-04-24
- Category: Neuroimaging, Geometric Deep Learning, Alzheimer's Disease
Core Methodology
Problem Statement
Alzheimer's disease (AD) classification from structural MRI faces challenges:
- Computational cost - 3D CNNs are resource-intensive
- Limited deployment - Difficult to deploy in resource-constrained settings
- Memory requirements - 3D convolutions consume significant GPU memory
- Interpretability - Voxel-based methods lack anatomical interpretability
Key Innovations
1. Anatomical Priority Sampling (APS)
Converts T1-weighted MRI into neuroanatomically-labeled 2D point clouds:
- Prioritizes sampling from AD-relevant brain regions
- Preserves anatomical structure in point cloud representation
- Creates ADNI-2DPC: first neuroanatomically labeled MRI-derived point cloud dataset
2. NeuroAPS-Net Architecture
Lightweight geometric deep learning model with:
- Region-aware feature encoding
- ROI token aggregation
- Anatomical prior integration
System Pipeline
T1-weighted MRI
↓
[Preprocessing: Skull Stripping, Registration]
↓
[Anatomical Segmentation: AAL or Destrieux Atlas]
↓
[Anatomical Priority Sampling (APS)]
↓
Neuroanatomical Point Cloud (ADNI-2DPC)
↓
[NeuroAPS-Net: Geometric Deep Learning]
↓
AD Classification (CN/MCI/AD)
Anatomical Priority Sampling (APS)
AD-Relevant Brain Regions:
- Hippocampus (medial temporal lobe)
- Amygdala
- Entorhinal cortex
- Posterior cingulate cortex
- Precuneus
- Lateral temporal cortex
- Parietal association cortex
Sampling Strategy:
Traditional Uniform Sampling:
┌──────────────────────────────┐
│ • • • • • │ ← Equal density everywhere
│ • • • • • │
│ • • • • • │
└──────────────────────────────┘
Anatomical Priority Sampling:
┌──────────────────────────────┐
│ ••• (hippocampus) │ ← Higher density in AD regions
│ • ••••• • │
│ • (precuneus) • │
│ • • • • • │ ← Lower density elsewhere
└──────────────────────────────┘
NeuroAPS-Net Architecture
Input Point Cloud [N_points, 3(xyz) + C_features + R_roi_id]
↓
┌───────────────────────┐
│ Point Feature Encoder│
│ - MLP for local feat │
└───────────┬───────────┘
↓
┌───────────────────────┐
│ Region-Aware Encoding│
│ - ROI-specific layers│
│ - Anatomical priors │
└───────────┬───────────┘
↓
┌───────────────────────┐
│ ROI Token Aggregation│
│ - Pool by anatomical │
│ region │
└───────────┬───────────┘
↓
┌───────────────────────┐
│ Classification Head │
│ - MLP + Softmax │
└───────────┬───────────┘
↓
AD / MCI / CN
Implementation Guide
Prerequisites
numpy
scipy
torch
torch-geometric
nibabel
scikit-learn
ants
freesurfer
Anatomical Priority Sampling
import numpy as np
import nibabel as nib
from scipy.spatial import cKDTree
class AnatomicalPrioritySampler:
"""
Convert T1-weighted MRI to anatomically-informed point cloud.
"""
def __init__(self, ad_relevant_regions=None, base_samples=2048,
priority_ratio=0.6):
"""
Args:
ad_relevant_regions: List of ROI IDs for AD-relevant regions
base_samples: Total number of points to sample
priority_ratio: Fraction of samples allocated to priority regions
"""
self.ad_regions = ad_relevant_regions or [
37, 38,
39, 40,
89, 90,
85, 86,
67, 68,
]
self.base_samples = base_samples
self.priority_ratio = priority_ratio
def load_mri_and_atlas(self, mri_path, atlas_path):
"""Load T1 MRI and anatomical atlas."""
mri_img = nib.load(mri_path)
atlas_img = nib.load(atlas_path)
mri_data = mri_img.get_fdata()
atlas_data = atlas_img.get_fdata()
coords = np.argwhere(mri_data > )
mri_data, atlas_data, coords
():
priority_points = []
priority_labels = []
n_priority_samples = (.base_samples * .priority_ratio)
roi_id .ad_regions:
roi_mask = atlas_data == roi_id
roi_coords = np.argwhere(roi_mask)
(roi_coords) == :
n_samples_per_roi = n_priority_samples // (.ad_regions)
(roi_coords) > n_samples_per_roi:
idx = np.random.choice((roi_coords), n_samples_per_roi, replace=)
sampled = roi_coords[idx]
:
sampled = roi_coords
intensities = mri_data[sampled[:, ], sampled[:, ], sampled[:, ]]
i, coord (sampled):
priority_points.append([coord[], coord[], coord[], intensities[i]])
priority_labels.append(roi_id)
np.array(priority_points), np.array(priority_labels)
():
background_mask = ~np.isin(atlas_data, .ad_regions) & (mri_data > )
background_coords = np.argwhere(background_mask)
(background_coords) > n_samples:
idx = np.random.choice((background_coords), n_samples, replace=)
sampled = background_coords[idx]
:
sampled = background_coords
intensities = mri_data[sampled[:, ], sampled[:, ], sampled[:, ]]
points = []
labels = []
i, coord (sampled):
points.append([coord[], coord[], coord[], intensities[i]])
labels.append()
np.array(points), np.array(labels)
():
projection_plane == :
points_2d = np.column_stack([
points_3d[:, ],
points_3d[:, ],
points_3d[:, ]
])
projection_plane == :
points_2d = np.column_stack([
points_3d[:, ],
points_3d[:, ],
points_3d[:, ]
])
:
points_2d = np.column_stack([
points_3d[:, ],
points_3d[:, ],
points_3d[:, ]
])
points_2d
():
mri_data, atlas_data, _ = .load_mri_and_atlas(mri_path, atlas_path)
priority_points, priority_labels = .sample_priority_regions(
mri_data, atlas_data,
)
n_background = .base_samples - (priority_points)
background_points, background_labels = .sample_background(
mri_data, atlas_data, n_background
)
all_points = np.vstack([priority_points, background_points])
all_labels = np.concatenate([priority_labels, background_labels])
point_cloud_2d = .convert_to_2d(all_points, projection)
point_cloud_2d, all_labels
NeuroAPS-Net Model
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import MessagePassing, global_mean_pool
class PointFeatureEncoder(nn.Module):
"""
Encode local point features using MLP.
"""
def __init__(self, in_channels=3, hidden_dim=64):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(in_channels, hidden_dim),
nn.ReLU(),
nn.BatchNorm1d(hidden_dim),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.BatchNorm1d(hidden_dim)
)
def forward(self, x):
return self.mlp(x)
class RegionAwareEncoding(nn.Module):
"""
Region-aware feature encoding with anatomical priors.
"""
def __init__(self, num_rois=116, feature_dim=64, embed_dim=32):
super().__init__()
self.num_rois = num_rois
self.roi_embedding = nn.Embedding(num_rois + 1, embed_dim)
self.feature_transform = nn.Sequential(
nn.Linear(feature_dim + embed_dim, feature_dim),
nn.ReLU(),
nn.Linear(feature_dim, feature_dim)
)
():
roi_embeds = .roi_embedding(roi_labels.long())
combined = torch.cat([features, roi_embeds], dim=-)
output = .feature_transform(combined)
output
(nn.Module):
():
().__init__()
.feature_dim = feature_dim
.num_rois = num_rois
():
roi_tokens = []
roi_id (, .num_rois + ):
mask = roi_labels == roi_id
mask.() > :
roi_feat = features[mask].mean(dim=)
:
roi_feat = torch.zeros(.feature_dim, device=features.device)
roi_tokens.append(roi_feat)
torch.stack(roi_tokens)
(nn.Module):
():
().__init__()
.point_encoder = PointFeatureEncoder(in_channels, hidden_dim)
.region_encoder = RegionAwareEncoding(num_rois, hidden_dim)
.roi_aggregator = ROITokenAggregation(hidden_dim, num_rois)
.classifier = nn.Sequential(
nn.Linear(hidden_dim * num_rois, ),
nn.ReLU(),
nn.Dropout(),
nn.Linear(, ),
nn.ReLU(),
nn.Dropout(),
nn.Linear(, num_classes)
)
():
features = .point_encoder(point_cloud)
features = .region_encoder(features, roi_labels)
roi_tokens = .roi_aggregator(features, roi_labels)
roi_tokens_flat = roi_tokens.view(, -)
logits = .classifier(roi_tokens_flat)
logits
Training Pipeline
def train_neuroaps_net(model, train_loader, val_loader, epochs=100, lr=1e-3):
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
criterion = nn.CrossEntropyLoss()
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=10)
best_val_acc = 0
for epoch in range(epochs):
model.train()
train_loss = 0
for point_cloud, roi_labels, labels in train_loader:
optimizer.zero_grad()
logits = model(point_cloud, roi_labels)
loss = criterion(logits, labels)
loss.backward()
optimizer.step()
train_loss += loss.item()
model.eval()
val_correct = 0
val_total = 0
with torch.no_grad():
for point_cloud, roi_labels, labels in val_loader:
logits = model(point_cloud, roi_labels)
_, predicted = torch.max(logits, 1)
val_correct += (predicted == labels).sum().item()
val_total += labels.size(0)
val_acc = val_correct / val_total
scheduler.step(val_acc)
if val_acc > best_val_acc:
best_val_acc = val_acc
torch.save(model.state_dict(), 'best_neuroaps_net.pth')
print(f"Epoch {epoch+1}: Train Loss={train_loss/len(train_loader):.4f}, "
f"Val Acc={val_acc:f}")
Applications
- Early AD Detection - Screen for mild cognitive impairment
- Clinical Decision Support - Assist radiologists in diagnosis
- Longitudinal Tracking - Monitor disease progression
- Research Studies - Large-scale AD analysis
- Resource-Constrained Settings - Deploy in clinics with limited GPU resources
Key Metrics
- Accuracy: Competitive with state-of-the-art 3D CNNs
- Efficiency: Significantly reduced inference latency
- Memory: Lower GPU memory requirements
- Interpretability: ROI-level predictions explain which brain regions contribute
Pitfalls
- Atlas Dependency - Requires accurate anatomical segmentation
- Sampling Variability - Random sampling may affect reproducibility
- 2D Projection - Some 3D spatial information is lost
- ROI Selection - AD-relevant regions are dataset-dependent
- Point Cloud Size - Trade-off between detail and computational cost
Related Skills
- alzheimer-pet-suvr-network-models - Spatio-temporal AD models
- multimodal-brain-connectivity-gnn - Multi-modal brain analysis
- brain-graph-neural - Graph-based brain network analysis
References
@article{islam2026neuroaps,
title={NeuroAPS-Net: Neuro-Anatomically Aware Point Cloud Representation for Efficient Alzheimer's Disease Classification},
author={Islam, Towhidul and Mahmud, Mufti},
journal={arXiv preprint arXiv:2604.22883},
year={2026}
}