A multimodal conditional diffusion framework for synthetic prostate MRI generation using:
- T2-weighted MRI
- ADC MRI
- High b-value diffusion MRI (HBV)
- Tumor segmentation masks
- Clinical metadata (Age, PSA, ISUP Grade)
This repository accompanies the dissertation work on:
Conditional diffusion-based prostate MRI synthesis integrating spatial anatomical priors and non-spatial clinical metadata.
This project implements a conditional diffusion model based on a modified U-Net architecture for generating realistic prostate MRI slices.
The framework integrates:
-
Spatial conditioning
- Tumor segmentation masks
-
Temporal conditioning
- Diffusion timestep embeddings
-
Clinical conditioning
- Age
- PSA
- ISUP Grade Group
The generated outputs are evaluated using:
- Fréchet inception distance (FID)
- Qualitative visual assessment
- Anatomical consistency evaluation
.
├── data/ # Dataset directory (symlinked externally)
├── evaluation/
│ ├── fake_images/ # Generated synthetic MRI slices
│ └── real_images/ # Real MRI slices for comparison
├── models/ # Saved model checkpoints (symlinked externally)
├── results/ # Dissertation result figures
├── scripts/
│ ├── find_tumors.py
│ ├── generate_conditional.py
│ ├── mass_evaluation.py
│ ├── train_picai_monai.py
│ └── user_test.py
├── slurm_scripts/ # HPC SLURM job scripts
├── src/
│ ├── datasets/
│ │ ├── picai_dataset.py
│ │ └── preprocess_picai_2d.py
│ └── training/
│ └── multimodal_unet.py
└── README.md
Recommended:
- NVIDIA GPU
- CUDA 11.8+
- ≥16GB VRAM preferred
Tested on:
- NVIDIA A100 GPUs
- Stanage HPC Cluster
git clone [https://github.com/aca23sa/Prostate_Gen_AI.git]
cd Prostate_Gen_AIpython3 -m venv venvActivate the environment:
source venv/bin/activatevenv\Scripts\activatepip install --upgrade pippip install -r requirements.txtThis project uses the PI-CAI prostate MRI dataset.
Dataset resources:
After obtaining access to the PI-CAI dataset:
Create the required directories:
mkdir -p data/raw/images
mkdir -p data/raw/labelsPlace:
- MRI volumes into
data/raw/images - Segmentation masks into
data/raw/labels
data/
└── raw/
├── picai/
│ ├── 10000/
│ │ ├── 10000_1000000_adc.mha
│ │ ├── 10000_1000000_hbv.mha
│ │ ├── 10000_1000000_t2w.mha
│ │ ├── 10000_1000000_cor.mha
│ │ └── 10000_1000000_sag.mha
│ ├── 10001/
│ ├── 10002/
│ └── ...
└── picai_labels/
├── anatomical_delineations/
├── clinical_information/
├── csPCa_lesion_delineations/
├── additional_resources/
├── LICENSE
└── README.md
Each patient folder inside picai/ contains multiple MRI modalities:
- adc → Apparent Diffusion Coefficient
- hbv → High b-value diffusion imaging
- t2w → T2-weighted MRI
- cor → Coronal view
- sag → Sagittal view
The picai_labels/ directory contains:
- anatomical_delineations/ → Anatomical prostate segmentation masks
- csPCa_lesion_delineations/ → Clinically significant prostate cancer lesion masks
- clinical_information/ → Patient and clinical metadata
- additional_resources/ → Supporting challenge resources
The preprocessing scripts automatically load the .mha MRI volumes and associated label masks for training and tumour-conditioned generation.
Run:
python src/datasets/preprocess_picai_2d.pyThis preprocessing pipeline performs:
- MRI normalization
- Slice extraction
- Modality alignment
- Tensor formatting
- Tumor mask integration
Run:
python scripts/train_picai_monai.pyThe UNet architecture is implemented in:
src/training/multimodal_unet.py
Submit the training job:
sbatch slurm_scripts/train_gpu.shExample SLURM script:
#!/bin/bash
#SBATCH --job-name=gen_ai_prostate
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=8
#SBATCH --mem=32G
#SBATCH --time=24:00:00
source venv/bin/activate
python scripts/train_picai_monai.pyRun:
python scripts/generate_conditional.pyGenerated images will be saved in:
results/
To evaluate generated images against real MRI scans:
python scripts/mass_evaluation.pyEvaluation images are stored in:
evaluation/
├── fake_images/
└── real_images/
Run FID evaluation:
python -m pytorch_fid evaluation/real_images evaluation/fake_imagesTo identify and extract slices containing tumour regions:
python scripts/find_tumors.pyThis script scans the validation dataset and outputs all MRI slice indices that contain tumour masks.
The output can be used to:
- Identify clinically significant tumour slices
- Select validation examples for qualitative analysis
- Compare tumour-conditioned generations against real MRI scans
- Support visual evaluation experiments
The generated slice indices can then be used with:
python scripts/user_test.pyuser_test.py allows interactive visual comparison between:
- Real prostate MRI slices
- AI-generated prostate MRI slices
- Tumour-conditioned synthetic outputs Users can input slice numbers identified by find_tumors.py to directly compare generated MRI outputs against the corresponding real MRI slices from the validation dataset.
The diffusion model integrates:
- Age
- PSA
- ISUP Grade Group
through:
- MLP embedding layers
- Cross-attention conditioning
- Multimodal feature fusion
inside the conditional U-Net bottleneck and decoder blocks.
The architecture is based on:
- Conditional U-Net
- DDPM-style diffusion training
- Sinusoidal timestep embeddings
- Cross-attention conditioning
Input tensor:
x_t ∈ R^(4 × 256 × 256)
Input channels:
- T2 MRI
- ADC MRI
- HBV MRI
- Tumor mask
git clone [https://github.com/aca23sa/Prostate_Gen_AI.git]python3 -m venv venv
source venv/bin/activatepip install -r requirements.txtPlace MRI scans and masks into:
data/raw/
python datasets/preprocess_picai_2d.pypython scripts/train_picai_monai.pypython scripts/generate_conditional.pypython scripts/mass_evaluation.pyFID Evaluation The dissertation uses Fréchet Inception Distance (FID) to quantitatively compare generated MRI slices against real prostate MRI images. After running:
python scripts/mass_evaluation.pythe generated and real images are placed into:
evaluation/fake_images/
evaluation/real_images/
FID can then be computed using the Python FID package. Install the package:
pip install pytorch-fidRun FID evaluation:
python -m pytorch_fid evaluation/real_images evaluation/fake_imagesLower FID scores indicate that the generated images are more similar to the real MRI distribution.
Compare newly generated outputs against the figures inside:
results/
Evaluation outputs:
evaluation/
├── fake_images/
├── real_images/
Saved model checkpoints:
models/conditional
- Training from scratch may require several hours to days depending on GPU hardware.
- A100 GPUs are recommended for full-resolution diffusion training.
- Mixed precision training can significantly reduce VRAM usage.
- Running on HPC clusters is recommended for reproducibility.
If you use this repository, please cite:
@misc{gen_ai_prostate,
title={Conditional Diffusion Models for Synthetic Prostate MRI Generation},
author={Shayaan Ather Hashmi},
year={2026},
url={https://github.com/aca23sa/Prostate_Gen_AI.git}
}- PI-CAI Challenge Dataset
- MONAI Framework
- PyTorch
- Hugging Face Diffusers