Adni Project
ð¯ Project Overview A production-ready deep learning system implementing a multi-modal ensemble learning paradigm for Alzheimer's Disease classification using DTI and sMRI neuroimaging data from th
About the Project
ð¯ Project Overview
A production-ready deep learning system implementing a multi-modal ensemble learning paradigm for Alzheimer's Disease classification using DTI and sMRI neuroimaging data from the ADNI dataset.
Key Innovation
Combines structural (sMRI) and diffusion tensor imaging (DTI) brain information using committee machine-based ensemble learning to provide more reliable and accurate classification than conventional single-CNN approaches.
Classification Task
CN / MCI / EMCI / AD - Four-class classification of Alzheimer's Disease stages
â Requirements Compliance
This implementation fully satisfies all project requirements:
â Multi-Modal DTI and sMRI Data Processing
- DTI Modalities: Mean Diffusivity (MD), Fractional Anisotropy (FA)
- sMRI Modalities: Grey Matter (GM), White Matter (WM) tissue maps (preferred)
- sMRI Fallback: T1-weighted structural MRI when tissue maps unavailable
- Multi-Modal Fusion: Combines DTI and sMRI of same patient with intelligent modality selection
- Preference Hierarchy: GM/WM tissue maps â T1-weighted fallback
- Real-time Reporting: Shows exactly which modalities used per patient during training
â Deep CNN Models
Implements multiple state-of-the-art architectures:
- Xception
- MobileNetV3 (Small & Large)
- InceptionV3
- InceptionResNetV2
- ResNet50
- DenseNet121
- EfficientNet (B0, B3)
â Ensemble Learning (Committee Machines)
Three ensemble methods for improved reliability:
- Voting Ensemble (Hard/Soft voting)
- Weighted Averaging (Optimized weights)
- Stacking Ensemble (Meta-learner)
â Comprehensive Evaluation Metrics
- Accuracy
- Sensitivity (Recall)
- Specificity
- Precision
- F1-Score
- Confusion Matrix
- Area Under ROC Curve (AUC)
- Per-class metrics for all categories
â ADNI Dataset Support
- Automated dataset parsing and organization
- Support for DTI (MD, FA) and MRI (T1) modalities
- Handles CN/MCI/EMCI/AD classifications
- Required combinations: MCI + DTI, MCI + MRI + DTI â
ð Enhanced Features
Advanced Training Pipeline (train_enhanced.py)
- Comprehensive Metrics Tracking: Tracks accuracy, loss, F1, precision, recall, AUC, and learning rate over epochs
- Data Imbalance Handling:
- Enhanced Class Weights Mode (default): Applies weighted loss with 3x scaling for minority classes
- Balanced Mode (
--balanced): Oversamples minority class to create balanced dataset - Advanced Balancing Strategies: SMOTE oversampling, random undersampling, combined SMOTE+ENN
- Optimized Performance: Enhanced configuration targeting accuracy > 0.93 and AUC 0.7-0.9
- AdamW Optimizer with weight decay and gradient clipping
- Deeper Architecture: 1024â512â256 dense layers with batch normalization
- Enhanced Regularization: Optimized dropout (0.4) and L2 regularization (0.0005)
- Better Learning Rate Schedule: Increased patience (15 epochs) with 0.5 reduction factor
- More Training Epochs: 100 epochs with early stopping for better convergence
- Larger Batch Size: 16 for more stable gradient estimates
- Individual Model Tracking: Separate results directory for each model with complete metrics
- Comprehensive Visualizations for ALL Models (including Ensemble):
- Confusion Matrices: Enhanced with mako colormap and rotated labels (45°)
- ROC Curves: Multi-class ROC curves with AUC scores for each class
- Per-Class Metrics: Bar charts showing F1, Precision, and Recall for each class
- Classification Reports: Heatmap visualization of precision, recall, and F1 per class
- Training Metrics: 2x3 grid plots showing accuracy, loss, learning rate, F1, precision, recall over epochs
- Model Comparison: Accuracy bar charts and performance heatmaps comparing all models
- Grad-CAM Visualizations: Explainable AI heatmaps showing which brain regions influence model predictions
- High-resolution: All plots saved at 300 DPI for publication quality
Interactive Jupyter Notebook
- Step-by-step training workflow with comprehensive documentation
- Real-time visualization of all metrics during training
- Interactive parameter tuning for experimentation
- Educational explanations of each step
- Located in
notebooks/adni_enhanced_training.ipynb
Organized Project Structure
src/: All Python source code organized in one placenotebooks/: Interactive Jupyter notebooks for training and analysisresults/: Structured output with individual model directoriesmodels/: Centralized model storage with timestampsdocs/: Essential documentation only
Explainable AI (XAI) with Grad-CAM
- Automated Grad-CAM Generation: Runs automatically after model evaluation
- Model Interpretability: Visualizes which brain regions influence predictions
- Multi-Model Support: Generates Grad-CAM for all trained models (Xception, ResNet50, EfficientNet, MobileNet)
- Sample Selection: Automatically selects representative samples from each class
- High-Quality Outputs: Side-by-side original and Grad-CAM overlay images at 300 DPI
- Organized Results: Separate directories for each model's Grad-CAM visualizations
ð Project Structure & File Responsibilities
.
âââ main.py # Main entry point - orchestrates full pipeline â
âââ main.ipynb # Jupyter notebook version of main.py (for Colab) â
âââ evaluate_saved_models.py # Comprehensive model evaluation script
âââ gradcam_visualization.py # Grad-CAM explainability visualization script â
âââ COLAB_USAGE.md # Google Colab usage guide â
â
âââ src/ # Source code directory
â âââ config.py # Global configuration settings and hyperparameters
â âââ parse_adni_dataset.py # ADNI dataset parser - scans and organizes raw data
â âââ adni_loader.py # Multi-modal data loader - loads and preprocesses data
â âââ data_preprocessing.py # DTI (MD/FA) and sMRI (GM/WM) preprocessing pipelines
â âââ cnn_models.py # CNN architecture implementations (8+ models)
â âââ ensemble_model.py # Committee machine ensemble learning framework
â âââ train_adni_ensemble.py # Basic training pipeline
â âââ train_enhanced.py # Enhanced training pipeline with comprehensive metrics â
â âââ evaluation.py # Comprehensive evaluation metrics calculator
â âââ visualization.py # Result visualization and plotting utilities
â âââ utils.py # Utility functions and helper methods
â
âââ notebooks/ # Jupyter notebooks for interactive training
â âââ adni_enhanced_training.ipynb # Enhanced training notebook with visualizations â
â
âââ data/ # ADNI dataset storage
â âââ ADNI/ # Raw ADNI neuroimaging data
â âââ 27aug25mci,dti,mri_metadata/ # XML metadata files
â âââ processed/ # Preprocessed and cached data
â
âââ models/ # Saved trained models (central storage)
â âââ MobileNetV3Small_TIMESTAMP.keras # Individual CNN models
â âââ EfficientNetB0_TIMESTAMP.keras # with timestamps
â âââ Xception_TIMESTAMP.keras # for version control
â âââ ensemble_TIMESTAMP.pkl # Ensemble models
â
âââ results/ # Experiment results and outputs
â âââ adni_ensemble_*/ # Basic training results
â âââ adni_ensemble_enhanced_*/ # Enhanced training results â
â âââ MobileNetV3Small_model/ # Individual model results
â â âââ model.keras # Trained model
â â âââ MobileNetV3Small_results.json # Metrics JSON
â â âââ MobileNetV3Small_comprehensive_metrics.png # 2x3 grid: Acc/Loss/LR/F1/Recall/Precision
â â âââ MobileNetV3Small_confusion_matrix.png # Mako colormap, rotated labels
â â âââ MobileNetV3Small_roc_curves.png # Multi-class ROC with AUC
â â âââ MobileNetV3Small_per_class_metrics.png # F1/Precision/Recall bar charts
â â âââ MobileNetV3Small_classification_report.png # Heatmap of metrics per class
â â âââ MobileNetV3Small_training_log.csv # Epoch-by-epoch training log
â âââ EfficientNetB0_model/ # Same structure as above
â âââ Xception_model/ # Same structure as above
â âââ ResNet50_model/ # Same structure as above (if trained)
â âââ Ensemble/ # Ensemble results with ALL visualizations
â â âââ ensemble_model.pkl # Saved ensemble
â â âââ ensemble_results.json # Ensemble metrics
â â âââ Ensemble_confusion_matrix.png # Mako colormap
â â âââ Ensemble_roc_curves.png # Multi-class ROC
â â âââ Ensemble_per_class_metrics.png # Per-class bar charts
â â âââ Ensemble_classification_report.png # Classification heatmap
â âââ comparison_plots/ # Model comparison visualizations
â â âââ accuracy_comparison.png # Bar chart of all models
â â âââ performance_heatmap.png # Heatmap: models vs metrics
â â âââ model_comparison_summary.csv # Detailed comparison table
â âââ summary_table.csv # Overall performance summary (all models + ensemble)
â
âââ evaluation_results_*/ # Comprehensive evaluation results (generated by evaluate_saved_models.py)
â âââ MobileNetV3Small_evaluation/ # Per-model evaluation
â âââ EfficientNetB0_evaluation/
â âââ Xception_evaluation/
â âââ ResNet50_evaluation/
â âââ Ensemble_evaluation/
â âââ comparison_plots/ # Cross-model comparison
â âââ model_comparison_summary.csv
â âââ gradcam_visualizations/ # Grad-CAM explainability visualizations â
â âââ MobileNetV3Small/ # Grad-CAM for each model
â â âââ gradcam_CN_0.png # Sample from CN class
â â âââ gradcam_MCI_0.png # Sample from MCI class
â â âââ gradcam_EMCI_0.png # Sample from EMCI class
â â âââ gradcam_AD_0.png # Sample from AD class
â âââ EfficientNetB0/
â âââ Xception/
â âââ ResNet50/
â
âââ docs/ # Documentation
â âââ AD experiment.docx # Experiment documentation
â âââ Metrics_Table1.docx # Metrics tables
â âââ Summary.md # Requirements compliance verification
â
âââ requirements.txt # Python dependencies
âââ adni_dataset_complete.csv # Parsed dataset information (generated)
âââ README.md # This comprehensive documentation
ð§ Core File Functions
1. main.py - Main Entry Point â
- Purpose: Orchestrates the complete training and evaluation pipeline
- Responsibilities:
- Step 1: Executes dataset parsing (
parse_adni_dataset.py) - Step 2: Runs enhanced training (
train_enhanced.py) - Step 3: Evaluates saved models (
evaluate_saved_models.py) - Step 4: Generates Grad-CAM visualizations (
gradcam_visualization.py) - Provides command-line options (
--skip-parse,--basic,--balanced) - Handles workflow automation and error management
- Displays progress and completion status
- Simplifies execution for end users
- Step 1: Executes dataset parsing (
2. src/config.py - System Configuration
- Purpose: Centralized configuration management
- Responsibilities:
- Training hyperparameters (batch size, epochs, learning rate)
- Model selection and ensemble configuration
- Data preprocessing parameters
- File paths and directory structure
- Class definitions for CN/MCI/EMCI/AD
3. src/parse_adni_dataset.py - Dataset Parser
- Purpose: ADNI dataset discovery and organization
- Responsibilities:
- Scans ADNI directory structure recursively
- Extracts patient metadata from XML files
- Identifies available modalities (DTI MD, FA, AD, RD; MRI T1)
- Maps diagnosis information (CN/MCI/EMCI/AD)
- Creates comprehensive dataset summary CSV
- Handles ADNI-specific folder naming conventions
4. src/adni_loader.py - Multi-Modal Data Loader
- Purpose: Data loading and intelligent multi-modal fusion
- Responsibilities:
- Loads NIfTI neuroimaging files (.nii, .nii.gz)
- Implements hierarchical multi-modal data fusion (DTI+MRI)
- Preference System: GM/WM tissue maps â T1-weighted fallback
- Handles missing modalities gracefully with automatic fallback
- Converts 3D volumes to 2D slices for CNN input
- Performs train/validation/test data splitting
- Real-time Modality Tracking: Reports which modalities used per patient
- Manages patient-level data organization with modality usage statistics
4. src/data_preprocessing.py - Neuroimaging Preprocessing
- Purpose: DTI and sMRI data preprocessing pipelines
- Responsibilities:
- DTI Processing:
- Loads DWI, bval, bvec files using DIPY
- Fits diffusion tensor models
- Extracts MD (Mean Diffusivity) and FA (Fractional Anisotropy)
- sMRI Processing:
- Loads T1-weighted MRI volumes
- Performs skull stripping and brain extraction
- Tissue segmentation (GM/WM separation)
- Volume normalization and standardization
- Spatial resizing and resampling
- DTI Processing:
5. src/cnn_models.py - CNN Architecture Library
- Purpose: Deep CNN model implementations
- Responsibilities:
- Implements 8+ state-of-the-art CNN architectures
- Models: Xception, MobileNetV3 (Small/Large), InceptionV3, InceptionResNetV2, ResNet50, DenseNet121, EfficientNet (B0/B3)
- Handles transfer learning and fine-tuning
- Configurable input shapes for different modalities
- Model compilation with appropriate optimizers and losses
6. src/ensemble_model.py - Committee Machine Framework
- Purpose: Ensemble learning implementation
- Responsibilities:
- Voting Ensemble: Hard and soft voting mechanisms
- Weighted Averaging: Optimized weight calculation
- Stacking Ensemble: Meta-learner training
- Committee machine paradigm implementation
- Prediction aggregation and consensus building
- Model combination strategies
7. src/train_adni_ensemble.py - Basic Training Pipeline
- Purpose: Basic training pipeline for quick experiments
- Responsibilities:
- Simple training workflow
- Data loading and preprocessing
- Multiple CNN model training
- Basic ensemble creation
- Model saving and evaluation
8. src/train_enhanced.py - Enhanced Training Pipeline
- Purpose: Advanced training with comprehensive metrics tracking
- Responsibilities:
- Comprehensive Metrics: Accuracy, Loss, F1, Precision, Recall, AUC, Learning Rate
- Data Imbalance Handling:
- Class weights mode (default)
- Balanced mode with oversampling (
--balancedflag) - Advanced balancing strategies: SMOTE, undersampling, combined
- Enhanced Visualizations: 2x3 grid plots with all metrics over epochs
- Individual Model Tracking: Separate results for each model
- Confusion Matrices: Mako colormap with rotated labels
- Model Comparison: Side-by-side performance analysis
- Grad-CAM Generation: Automatic explainability visualizations for all models
- Target Performance: Accuracy > 0.93, AUC 0.7-0.9
9. src/evaluation.py - Metrics Calculator
- Purpose: Comprehensive model evaluation
- Responsibilities:
- Accuracy: Overall classification accuracy
- Sensitivity/Recall: True positive rate per class
- Specificity: True negative rate per class
- Precision: Positive predictive value
- F1-Score: Harmonic mean of precision and recall
- Confusion Matrix: Detailed classification breakdown
- AUC: Area under ROC curve calculation
- Grad-CAM Visualization: Generate explainability heatmaps for model predictions
- Per-class and macro/weighted averaging
- Statistical significance testing
10. src/visualization.py - Result Visualization
- Purpose: Generate plots and visual outputs
- Responsibilities:
- Confusion matrix heatmaps (mako colormap, rotated labels)
- ROC curves for all classes
- Training history plots (2x3 grid: accuracy, loss, LR, F1, precision, recall)
- Per-class performance bar charts
- Model comparison visualizations
- Grad-CAM Heatmaps: Compute and overlay gradient-weighted class activation maps
- Explainability Visualizations: Show which brain regions influence predictions
- Publication-ready figure generation
11. src/utils.py - Utility Functions
- Purpose: Common helper functions and utilities
- Responsibilities:
- File I/O operations
- Data format conversions
- Logging and debugging utilities
- Path management functions
- Error handling helpers
- Memory management utilities
12. evaluate_saved_models.py - Comprehensive Model Evaluation
- Purpose: Evaluate all trained models on test data
- Responsibilities:
- Loads latest trained models from results directory
- Loads test data from processed directory
- Evaluates individual models and ensemble
- Generates comprehensive metrics and visualizations
- Creates comparison reports and summary tables
- Saves results to evaluation_results_* directory
13. gradcam_visualization.py - Grad-CAM Explainability â
- Purpose: Generate Grad-CAM visualizations for model interpretability
- Responsibilities:
- Automatic Layer Detection: Finds last convolutional layer for each model architecture
- Multi-Model Support: Works with Xception, ResNet50, EfficientNet, MobileNet, etc.
- Sample Selection: Automatically selects representative images from each class
- Gradient Computation: Computes class activation maps using gradient-weighted features
- Heatmap Overlay: Superimposes heatmaps on original brain images
- High-Quality Output: Saves side-by-side comparisons at 300 DPI
- Organized Results: Creates separate directories for each model's visualizations
- Preprocessing: Applies model-specific preprocessing (EfficientNet, Xception, ResNet, MobileNet)
14. notebooks/adni_enhanced_training.ipynb - Interactive Training Notebook â
- Purpose: Interactive training and experimentation
- Responsibilities:
- Automatic dataset parsing: Creates CSV if not exists
- Step-by-step training workflow (16 cells)
- Real-time visualization of metrics
- Interactive parameter tuning
- Comprehensive documentation and explanations
- Easy experimentation and prototyping
ð Quick Start
1. Installation
# Install dependencies pip install -r requirements.txt
2. Run Complete Pipeline
You can run the pipeline using either Python script or Jupyter Notebook:
Option A: Python Script (Recommended for Local)
# Run the complete pipeline (parse + train + evaluate + Grad-CAM) python main.py
Option B: Jupyter Notebook (Recommended for Colab)
# Open the notebook jupyter notebook main.ipynb # Or upload to Google Colab and run there # See COLAB_USAGE.md for detailed instructions
For Google Colab users: See COLAB_USAGE.md for detailed setup instructions.
What it does:
- Step 1: Parses ADNI dataset and creates CSV summary
- Scans
data/ADNI/directory - Extracts patient information from XML metadata
- Identifies DTI (MD, FA, AD, RD) and MRI (T1) files
- Creates
adni_dataset_complete.csvwith complete patient inventory
- Scans
- Step 2: Trains enhanced ensemble model
- Comprehensive metrics tracking
- Data imbalance handling
- Enhanced visualizations
- Individual model results
- Step 3: Evaluates all saved models
- Loads trained models and test data
- Generates comprehensive evaluation metrics
- Creates comparison reports and visualizations
- Step 4: Generates Grad-CAM visualizations
- Creates explainability heatmaps for all models
- Shows which brain regions influence predictions
- Saves high-quality visualizations for interpretation
Options:
# Skip parsing if CSV already exists python main.py --skip-parse # Use basic training instead of enhanced python main.py --basic # Train with balanced dataset (oversampling minority class) python main.py --balanced # Skip parsing + balanced training python main.py --skip-parse --balanced
Balanced vs Imbalanced Training:
-
Default (Imbalanced): Uses class weights in loss function
- Faster training (no data duplication)
- Good for moderate imbalance (e.g., 70/30 split)
- Example: Class 0: 132 samples, Class 1: 58 samples
-
Balanced (
--balanced): Oversamples minority class- More training data (duplicates minority samples)
- Better for severe imbalance or when AUC is stuck at 0.5
- Example: Class 0: 132 samples, Class 1: 132 samples (58 original + 74 duplicated)
- Recommended when model only predicts majority class
3. Alternative: Step-by-Step Execution
Option A: Enhanced Training (Recommended)
# Step 1: Parse dataset cd src python parse_adni_dataset.py # Step 2: Train with comprehensive metrics tracking python train_enhanced.py # OR: Train with balanced dataset (oversampling) python train_enhanced.py --balanced
What it does:
- Loads multi-modal data with intelligent modality selection
- Prefers: DTI (MD, FA) + MRI (GM, WM tissue maps)
- Falls back to: DTI (MD, FA) + MRI (T1-weighted) when tissue maps unavailable
- Further fallback: MRI-only (GM/WM or T1) if DTI unavailable
- Trains 3 CNN models (MobileNetV3Small, EfficientNetB0, Xception)
- Tracks comprehensive metrics: Accuracy, Loss, F1, Precision, Recall, AUC, Learning Rate
- Handles data imbalance:
- Default: Computes and applies class weights
--balanced: Oversamples minority class to balance dataset
- Creates ensemble classifier with soft voting (target accuracy > 0.93, AUC 0.7-0.9)
- Comprehensive Visualizations for ALL Models (including Ensemble):
- Training Metrics: 2x3 grid plots (Accuracy, Loss, Learning Rate, F1, Precision, Recall)
- Confusion Matrices: Mako colormap with 45° rotated labels
- ROC Curves: Multi-class ROC with AUC scores for each class
- Per-Class Metrics: Bar charts showing F1, Precision, Recall for each class
- Classification Reports: Heatmap visualization of all metrics per class
- Model Comparison: Accuracy bar charts and performance heatmaps
- Grad-CAM Visualizations: Explainability heatmaps showing which brain regions influence predictions
- Individual model results: Separate directories with complete metrics and all visualizations
- Outputs everything to
results/adni_ensemble_enhanced_TIMESTAMP/
Option B: Interactive Notebook Training
# Open Jupyter notebook jupyter notebook notebooks/adni_enhanced_training.ipynb
What it provides:
- Automatic dataset parsing (creates CSV if not exists)
- Step-by-step interactive training
- Real-time visualization of all metrics
- Comprehensive documentation and explanations
- Easy parameter tuning and experimentation
Option C: Basic Training
# Step 1: Parse dataset cd src python parse_adni_dataset.py # Step 2: Quick basic training python train_adni_ensemble.py
What it does:
- Basic training pipeline without advanced metrics
- Faster for quick experiments
- Outputs to
results/adni_ensemble_TIMESTAMP/
4. Run on Google Colab
For running on Google Colab with GPU acceleration:
- Upload project to Google Drive
- Open
main.ipynbin Google Colab - Follow the instructions in COLAB_USAGE.md
Key changes for Colab:
- Mount Google Drive at the beginning
- Install neuroimaging packages:
!pip install nibabel nilearn dipy imbalanced-learn - Enable GPU runtime (Runtime â Change runtime type â GPU)
- Adjust batch size if needed (reduce to 8 or 16 for free Colab)
See COLAB_USAGE.md for complete setup instructions and troubleshooting.
5. Run Evaluation and Grad-CAM Separately
If you've already trained models and want to run evaluation or Grad-CAM separately:
# Evaluate all saved models python evaluate_saved_models.py # Generate Grad-CAM visualizations for all models (default) python gradcam_visualization.py # OR: Generate Grad-CAM only for ensemble model python gradcam_visualization.py --ensemble-only
What it does:
-
evaluate_saved_models.py:
- Finds latest results directory automatically
- Loads all trained models
- Evaluates on test data
- Generates comprehensive metrics and visualizations
- Creates comparison reports
-
gradcam_visualization.py:
- Finds latest trained models automatically
- Selects sample images from each class (2 per class)
- Computes Grad-CAM heatmaps for each model
- Overlays heatmaps on original brain images
- Saves high-quality visualizations (300 DPI)
- Organizes results by model in separate directories
- Options:
- Default: Generates Grad-CAM for all individual models
--ensemble-only: Generates Grad-CAM only for the ensemble model (faster)
Example Enhanced Training Output:
Loading DTI+MRI data...
Loading patient 027_S_4919... OK (DTI_MD, DTI_FA, MRI_GM, MRI_WM)
Loading patient 123_S_4567... OK (MRI_T1)
Modality Usage Summary:
DTI_MD+DTI_FA+MRI_GM+MRI_WM: 5 patients
MRI_T1: 20 patients
DTI_MD+DTI_FA+MRI_T1: 3 patients
Class weights for imbalanced data:
Class 0: 1.2500
Class 1: 0.8750
Training MobileNetV3Small...
Epoch 1/50 - f1: 0.8234 - precision: 0.8456 - recall: 0.8123 - lr: 1.00e-04
...
ð Dataset Information
Current ADNI Data Structure
data/27aug25mci,dti,mri/ADNI/
âââ 027_S_4919/ # Patient ID
â âââ Native_Space_Axial_Diffusivity_Image_*/ # DTI AD
â âââ Native_Space_Radial_Diffusivity_Image_*/ # DTI RD
â âââ corrected_FA_image_*/ # DTI FA
â âââ corrected_MD_image_*/ # DTI MD
â âââ MT1__N3__Scaled_*/ # T1 MRI
âââ [other patients...]
data/27aug25mci,dti,mri_metadata/ADNI/
âââ ADNI_027_S_4919_corrected_AD_image_*.xml # Diagnosis metadata
âââ [other XML files...]
Data Statistics
- Total Patients: 34
- Diagnosis Distribution:
- MCI: 31 patients
- EMCI: 3 patients
- Modalities Available:
- DTI (MD + FA): Limited patients with complete data
- MRI (T1): 34 patients
- Multi-modal (DTI+MRI): Subset of patients
ð Usage Examples
Basic Training
cd src python train_adni_ensemble.py
Enhanced Training with Comprehensive Metrics
cd src python train_enhanced.py
Interactive Notebook Training
jupyter notebook notebooks/adni_enhanced_training.ipynb
Custom Configuration
Edit src/config.py:
# Enhanced Training Parameters for Better Performance BATCH_SIZE = 16 # Increased for better gradient estimates EPOCHS = 100 # More epochs with early stopping LEARNING_RATE = 0.0005 # Optimized initial learning rate # Model selection CNN_MODELS = ['MobileNetV3Small', 'EfficientNetB0', 'Xception'] # Ensemble method ENSEMBLE_METHOD = 'voting' # 'voting', 'averaging', 'stacking' VOTING_TYPE = 'soft' # 'hard' or 'soft' # Enhanced features for better performance USE_CLASS_WEIGHTS = True # Handle data imbalance ENHANCED_CLASS_WEIGHTS = True # 3x scaling for minority classes DROPOUT_RATE = 0.4 # Optimized dropout L2_REGULARIZATION = 0.0005 # Optimized L2 regularization EARLY_STOPPING_PATIENCE = 15 # Increased patience REDUCE_LR_PATIENCE = 5 # LR reduction patience LR_REDUCTION_FACTOR = 0.5 # Reduce LR by half
Programmatic Usage
import sys sys.path.insert(0, 'src') from adni_loader import ADNIDataLoader from cnn_models import CNNModelBuilder from ensemble_model import EnsembleClassifier from evaluation import ModelEvaluator # Load multi-modal data loader = ADNIDataLoader(target_size=(128, 128, 128)) X, y, patient_ids = loader.load_dataset( modality='DTI+MRI', # or 'MRI' or 'DTI' use_2d=True, num_slices=32 ) # Split data X_train, X_val, X_test, y_train, y_val, y_test = loader.split_data(X, y) # Build and train models builder = CNNModelBuilder(input_shape=(128, 128, 1), num_classes=4) models = [] for model_name in ['Xception', 'MobileNetV3Small', 'EfficientNetB0']: model = builder.build_model(model_name) model.fit(X_train, y_train, validation_data=(X_val, y_val), epochs=50) models.append(model) # Create ensemble ensemble = EnsembleClassifier(models, method='voting', voting='soft') predictions = ensemble.predict(X_test) # Evaluate evaluator = ModelEvaluator(class_names=['CN', 'MCI', 'EMCI', 'AD']) metrics = evaluator.calculate_metrics(y_test, predictions)
ð Results and Evaluation
Generated Outputs
Enhanced Training Results
Each enhanced training run creates a comprehensive directory structure with complete visualizations for ALL models:
results/adni_ensemble_enhanced_TIMESTAMP/
âââ MobileNetV3Small_model/ # Individual model with ALL visualizations
â âââ model.keras # Trained model
â âââ MobileNetV3Small_results.json # Complete metrics
â âââ MobileNetV3Small_comprehensive_metrics.png # 2x3 grid: Acc/Loss/LR/F1/Recall/Precision
â âââ MobileNetV3Small_confusion_matrix.png # Mako colormap, 45° rotated labels
â âââ MobileNetV3Small_roc_curves.png # Multi-class ROC with AUC per class
â âââ MobileNetV3Small_per_class_metrics.png # F1/Precision/Recall bar charts
â âââ MobileNetV3Small_classification_report.png # Heatmap of metrics per class
â âââ MobileNetV3Small_training_log.csv # Epoch-by-epoch training log
âââ EfficientNetB0_model/ # Same complete visualization set
âââ Xception_model/ # Same complete visualization set
âââ ResNet50_model/ # Same complete visualization set (if trained)
âââ Ensemble/ # Ensemble with ALL visualizations
â âââ ensemble_model.pkl # Saved ensemble
â âââ ensemble_results.json # Complete metrics
â âââ Ensemble_confusion_matrix.png # Mako colormap, rotated labels
â âââ Ensemble_roc_curves.png # Multi-class ROC curves
â âââ Ensemble_per_class_metrics.png # Per-class bar charts
â âââ Ensemble_classification_report.png # Classification heatmap
âââ comparison_plots/ # Model comparison visualizations
â âââ accuracy_comparison.png # Bar chart comparing all models
â âââ performance_heatmap.png # Heatmap: models vs metrics
â âââ model_comparison_summary.csv # Detailed comparison table
âââ summary_table.csv # Overall performance summary (all models + ensemble)
Visualization Details
1. Training Metrics (2x3 Grid)
- Top Row: Accuracy, Learning Rate, Loss
- Bottom Row: F1 Score, Recall, Precision
- Shows training and validation curves over epochs
- Learning rate on log scale for better visibility
2. Confusion Matrix
- Colormap: Mako (as requested)
- Labels: 45° rotation for readability
- Annotations: Count values in each cell
- High Resolution: 300 DPI for publication quality
3. ROC Curves
- Multi-class: One curve per class with individual AUC scores
- Comparison: Shows performance vs random classifier
- Grid: Semi-transparent grid for easier reading
4. Per-Class Metrics
- Three Bar Charts: F1, Precision, Recall
- Value Labels: Exact values displayed on bars
- Color Coded: Different colors for each metric
5. Classification Report Heatmap
- Metrics: Precision, Recall, F1-Score per class
- Colormap: Blues gradient
- Format: 3 decimal places for precision
6. Model Comparison
- Accuracy Bar Chart: All models sorted by performance
- Performance Heatmap: All metrics across all models
- Summary Table: CSV with complete metrics for all models
Basic Training Results
results/adni_ensemble_TIMESTAMP/
âââ MobileNetV3Small.keras # Trained CNN model 1
âââ EfficientNetB0.keras # Trained CNN model 2
âââ Xception.keras # Trained CNN model 3
âââ ensemble_model.pkl # Ensemble model
âââ metrics.json # Complete performance metrics
âââ confusion_matrix.png # Confusion matrix visualization
âââ summary.json # Run summary
Evaluation Metrics Computed
Per-Epoch Metrics (Enhanced Training)
- Accuracy (training and validation)
- Loss (training and validation)
- F1 Score (validation)
- Precision (validation)
- Recall (validation)
- Learning Rate (tracked over epochs)
Final Evaluation Metrics
- Accuracy: Overall classification accuracy
- Sensitivity (Recall): True positive rate per class
- Specificity: True negative rate per class
- Precision: Positive predictive value per class
- F1-Score: Harmonic mean of precision and recall
- Confusion Matrix: 4x4 matrix for CN/MCI/EMCI/AD (with mako colormap)
- AUC: Area under ROC curve (macro and weighted)
- Per-class metrics: Individual performance for each diagnosis
Visualization Features
Enhanced Training Visualizations
- 2x3 Grid Plot: Comprehensive metrics over epochs
- Top row: Accuracy, Learning Rate (log scale), Loss
- Bottom row: F1 Score, Precision, Recall
- Enhanced Confusion Matrix: Mako colormap with rotated labels
- Model Comparison: Side-by-side performance analysis
- Grad-CAM Explainability: Visual explanations of model predictions
- High-resolution outputs: 300 DPI for publication quality
Grad-CAM Explainability
What is Grad-CAM?
Gradient-weighted Class Activation Mapping (Grad-CAM) is an explainable AI technique that visualizes which regions of the brain images are most important for the model's predictions. It helps clinicians understand and trust the model's decisions.
How It Works
- Gradient Computation: Computes gradients of the predicted class with respect to the final convolutional layer
- Importance Weighting: Weights each feature map by its importance to the prediction
- Heatmap Generation: Creates a heatmap showing which brain regions influenced the decision
- Overlay Visualization: Superimposes the heatmap on the original brain scan
Grad-CAM Outputs
For each trained model, Grad-CAM generates:
- Sample Visualizations: 3 samples per class showing original image + heatmap overlay
- Color-coded Heatmaps: Red/yellow regions indicate high importance, blue regions indicate low importance
- Class-specific Activations: Shows which brain regions are important for each diagnosis (CN/MCI/EMCI/AD)
- Saved in:
results/MODEL_NAME/gradcam/directory
Example Grad-CAM Output
results/adni_ensemble_enhanced_TIMESTAMP/
âââ MobileNetV3Small_model/
â âââ gradcam/
â âââ MobileNetV3Small_gradcam_MCI_0.png
â âââ MobileNetV3Small_gradcam_MCI_1.png
â âââ MobileNetV3Small_gradcam_EMCI_0.png
â âââ ...
âââ EfficientNetB0_model/
â âââ gradcam/
â âââ [similar structure]
âââ Xception_model/
âââ gradcam/
âââ [similar structure]
Clinical Interpretation
- Hippocampus Activation: Often highlighted in AD/MCI predictions (memory-related region)
- Cortical Regions: May show importance in structural changes
- White Matter Tracts: DTI-based models may highlight fiber integrity changes
- Validation Tool: Helps verify that models focus on clinically relevant brain regions
ð¬ Technical Implementation Details
DTI Preprocessing Pipeline (data_preprocessing.py)
- Load DWI Data: Reads diffusion-weighted images, b-values, b-vectors
- Tensor Fitting: Uses DIPY's tensor model fitting
- Parameter Extraction: Computes MD (mean diffusivity) and FA (fractional anisotropy)
- Normalization: Z-score normalization and intensity scaling
- Spatial Processing: Resizing and resampling to target dimensions
sMRI Preprocessing Pipeline (data_preprocessing.py)
- T1 Loading: Reads T1-weighted structural MRI
- Skull Stripping: Brain extraction using FSL/ANTs tools
- Tissue Segmentation: Separates grey matter (GM) and white matter (WM)
- Normalization: Intensity normalization and standardization
- Spatial Processing: Resizing and resampling
Multi-Modal Fusion (adni_loader.py)
- Channel Combination: Creates 3-channel input [MD, FA, T1]
- Slice Extraction: Converts 3D volumes to 2D slices
- Missing Data Handling: Graceful fallback when modalities unavailable
- Patient Alignment: Ensures DTI and MRI from same patient/session
Ensemble Methods (ensemble_model.py)
- Voting Ensemble:
- Soft voting: Averages predicted probabilities
- Hard voting: Majority class vote
- Weighted Averaging: Optimizes weights based on validation performance
- Stacking Ensemble: Meta-learner trained on base model predictions
� System Requirements
Software Dependencies
- Python: 3.8+
- TensorFlow: 2.10+
- CUDA: 11.2+ (for GPU acceleration)
- Neuroimaging: nibabel, dipy, nilearn
- Scientific: numpy, pandas, scikit-learn
- Visualization: matplotlib, seaborn
Hardware Recommendations
- Minimum: 8GB RAM, CPU-only
- Recommended: 16GB+ RAM, NVIDIA GPU with 6GB+ VRAM
- Storage: 10GB+ for dataset and results
ð Citation
@software{ad_multimodal_ensemble_2025, title={Multi-Modal Deep CNN Ensemble for Alzheimer's Disease Classification}, author={ADNI Research Team}, year={2025}, note={ADNI Dataset Implementation with Committee Machine Ensemble Learning} }
ð Acknowledgments
- ADNI (Alzheimer's Disease Neuroimaging Initiative) for the comprehensive dataset
- DIPY for advanced DTI processing capabilities
- TensorFlow/Keras for deep learning framework
- Nilearn for neuroimaging analysis utilities
- scikit-learn for machine learning evaluation tools
ð§ Support & Troubleshooting
For issues or questions:
- Requirements Verification: Check
Summary.mdfor compliance details - Training Logs: Review logs in
results/directory - Dataset Issues: Verify with
python parse_adni_dataset.py - Configuration: Adjust parameters in
config.py
Common Issues:
- Memory errors: Reduce batch size in
config.py - GPU issues: Ensure CUDA compatibility
- Data loading: Check file paths and permissions
- Missing modalities: System automatically falls back to available data
ð¯ Performance Target
Target Accuracy: > 0.75 (Enhanced from baseline) Target AUC: > 0.85 (Enhanced from baseline) Ensemble Method: Soft Voting with optimized architecture Visualization: Comprehensive 2x3 grid plots with mako colormap confusion matrices
Performance Improvements
- Enhanced Architecture: Deeper dense layers (1024â512â256) with batch normalization
- Better Optimizer: AdamW with weight decay (0.0001) and gradient clipping (1.0)
- Optimized Regularization: Dropout (0.4) and L2 (0.0005) for better generalization
- Improved Training: 100 epochs, batch size 16, learning rate 0.0005
- Enhanced Class Weighting: 3x scaling for minority classes to handle imbalance
- Better LR Schedule: Patience 15, reduction factor 0.5, min LR 1e-7
Utility Scripts
src/copy_chosen_dataset.py: Copy chosen patient data to separate folder for reference
Createspython src/copy_chosen_dataset.pychosen_dataset/folder with all patient data used in training
Status: â Production Ready | Enhanced Performance | All Requirements Met | Committee Machine Ensemble Learning
Project Timeline
Technologies
External Links
Related Projects
Projects built with similar technologies.
Online Html Editor And Viewer
The Online HTML Editor and Viewer is a simple web application built with Flask that allows users to write and preview HTML code in real-time.
Qrgen
A premium, feature-rich QR Code Generator engineered with Python (Flask) and a pristine Glassmorphism frontend.
Examina Ai
Using AI, It transforms raw study materials into structured, verified examination sets with support for institutional export formats like Moodle XML.