Group robustness-aware MRI training pipeline for Alzheimer's disease (AD) prediction, combining 3D MRI with tabular clinical/demographic features via multiple fusion methods.
DEAL supports:
- Targets: binary AD classification (
ad), MCI-to-AD transition (ad_transition), and time-to-AD Cox regression (ad_time_cox) - MRI backbones: MONAI ResNet, EfficientNet, Vision Transformer (ViT-T/S/B), with optional 2D sagittal view
- Fusion methods: feature concatenation, FiLM, FiLM+demographics, DAFT, tabular concatenation (with/without linear projection), tabular attention, and tabular-only baselines
- Group Robustness: Group DRO, SUBG, and age-based classifier decoupling
uv venv --python=3.10
source .venv/bin/activate
uv pip install torch==2.1.2 torchvision==0.16.2 torchaudio==2.1.2 --index-url https://download.pytorch.org/whl/cu121
uv pip install monai --no-deps
uv pip install numpy==1.26.4
uv pip install scikit-learn scikit-survival pandas torchio wandbWeights of pretrained models trained with AD/CN clssification are provided in https://drive.google.com/drive/folders/1VgpcKJDJJza9II9LzZVwy7OgQ96m8zel?usp=sharing.
Train (example: FiLM + demographics, AD transition, pretrained ResNet):
python main.py \
--method film_demo \
--film-only-last \
--film-aux-net mlp \
--film-aux-net-act silu \
--tabular-input mmse cdrsb adas11 age_dummy gender educat faq \
--target ad_transition \
--mri-arch monai_resnet18 \
--resolution 1mm \
--pretrained \
--normalize \
--epochs 30 \
--batch-size 4 \
--lr 1e-5 \
--mri-modelpath /path/to/pretrained_mri.pt \
--save-model \
--date my_experiment| Category | Arguments |
|---|---|
| Data | --data-dir, --resolution (1mm / 2mm), --target (ad / ad_transition / ad_time_cox), --normalize, --merge-train-val, --sagittal |
| Backbone | --mri-arch (efficientnet, vit_t, vit_s, vit_b, monai_resnet18), --resnet-backbone, --efficientnet-backbone |
| Method | --method (scratch, feature_concat, film, film_demo, daft, tab_concat, tab_concat_lin_proj, tab_attention, tabular_only) |
| FiLM/DAFT | --film-aux-net, --film-aux-net-act, --film-only-last, --tabular-input, --daft-input |
| Optimization | --epochs, --batch-size, --lr, --optimizer, --scheduler, --weight-decay |
| Group Robustness | --group-dro, --mri-subg, --decouple, --group-weight-lr |
| Evaluation | --test-set-id, --seed, --last-eval-only, --no-best-th, --save-model |
See argument.py for the full list.
scripts/run_DEAL.sh– example runs for DEALscripts/run_deal_arch_abl.sh– architecture ablations (EfficientNet, ViT-B, 2D sagittal) forad_transitionscripts/run_baselines.sh– baseline configurations
- Checkpoints:
./trained_models/<date>/mri/<target>/ - Results/logs:
./results/<date>/mri/<target>/
See repository for license terms.