Final-year project on generating synthetic medical images with Denoising Diffusion Probabilistic Models (DDPMs) and using them as training data. It contains:
- a DDPM (three U-Net variants) trained on chest radiographs, with DDPM and DDIM samplers;
- image-quality evaluation of the generated images (FID, KID, Inception Score, PSNR, SSIM, LPIPS);
- a pneumonia classifier (EfficientNet-B2 / ResNet-50 / DenseNet-121 ensemble) that can be trained with real data only or with synthetic images added;
- a controlled augmentation study on BloodMNIST comparing no augmentation, classic augmentation, a conditional GAN and a class-conditional DDPM;
- a small Flask web app for generating images from a trained checkpoint.
Synthetic 256x256 chest X-rays from the DDPM (DDIM, 100 steps).
medsynth/ library code
models/unet.py BasicUNet, UNet, AttentionUNet (+ build_unet)
models/classifier.py ensemble pneumonia classifier, focal loss
diffusion.py noise schedules, training loss, DDPM/DDIM samplers, EMA
data.py datasets and loaders (diffusion + classifier)
checkpoint.py save/load checkpoints; infers the architecture of older checkpoints
metrics.py FID / KID / IS / PSNR / SSIM / LPIPS
sampling.py batched generation and PNG export
scripts/
train_ddpm.py train a diffusion model
generate.py sample images from a checkpoint
evaluate_ddpm.py score a checkpoint against real images
train_classifier.py train + evaluate the pneumonia classifier (optionally with synthetic data)
predict_pneumonia.py classify individual X-rays
check_dataset.py count images per split/class
download_checkpoints.py fetch the published weights from the GitHub release
experiments/
bloodmnist_augmentation_comparison.py none vs. traditional vs. GAN vs. DDPM augmentation
web_ui/ Flask app (generate + gallery)
docs/results/ figures and metrics from the project runs
Python 3.10+ and a CUDA GPU are recommended (everything also runs on CPU, slowly).
python -m venv .venv
.venv\Scripts\activate # Windows (source .venv/bin/activate on Linux/macOS)
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # pick your CUDA version
pip install -r requirements.txtThe scripts work straight from the repository root. To use medsynth from elsewhere run pip install -e ..
The chest X-ray experiments use the Kaggle dataset Chest X-Ray Images (Pneumonia). Extract it so that the layout is:
data/chest_xray/
train/NORMAL/*.jpeg train/PNEUMONIA/*.jpeg
val/NORMAL/... val/PNEUMONIA/...
test/NORMAL/... test/PNEUMONIA/...
python scripts/check_dataset.py data/chest_xray prints the counts. BloodMNIST is downloaded
automatically by medmnist on first use.
# default: 'simple' U-Net, linear schedule, 256x256, EMA on
python scripts/train_ddpm.py --data-root data/chest_xray/train --out-dir outputs/ddpm_simple \
--epochs 50 --batch-size 8
# attention U-Net, cosine schedule, mixed precision, gradient accumulation (fits a 4 GB GPU)
python scripts/train_ddpm.py --data-root data/chest_xray/train --arch attention \
--batch-size 2 --grad-accum 8 --epochs 200 --warmup-epochs 5 --amp --out-dir outputs/ddpm_attentionEvery epoch writes last.pth (resumable with --resume), best.pth (lowest validation loss),
history.csv, training_curves.png and periodic sample grids.
python scripts/generate.py --ckpt outputs/ddpm_simple/best.pth --num 16 --steps 100 --grid samples.png
python scripts/generate.py --ckpt outputs/ddpm_simple/best.pth --num 500 --out-dir data/synthetic/NORMAL--method ddpm runs the full 1000-step sampler; DDIM with 50 to 100 steps is 10 to 20 times faster
with very similar quality.
python scripts/evaluate_ddpm.py --ckpt outputs/ddpm_simple/best.pth --data-root data/chest_xray/test \
--n-real 300 --n-fake 300 --steps 100 --out-dir outputs/evalWrites metrics.json, metrics.png and a sample grid. FID/KID compare distributions and are the
main numbers; PSNR/SSIM/LPIPS are computed against randomly paired real images and only give a rough
indication.
# real data only
python scripts/train_classifier.py --data-root data/chest_xray --out-dir outputs/classifier_real
# real data + synthetic NORMAL images generated in step 2
python scripts/train_classifier.py --data-root data/chest_xray \
--synthetic data/synthetic/NORMAL NORMAL --out-dir outputs/classifier_augmented
# classify new images
python scripts/predict_pneumonia.py --ckpt outputs/classifier_real/best.pth some_xray.jpegTraining uses focal loss, a class-balanced sampler, OneCycle learning-rate schedule, mixed precision
and early stopping. Outputs: best.pth, training_history.png, evaluation.png (confusion matrix,
ROC curve, probability histogram) and metrics.json.
python experiments/bloodmnist_augmentation_comparison.py --out-dir outputs/bloodmnist
python experiments/bloodmnist_augmentation_comparison.py --quick # a few minutes, for a smoke testThe full run trains the conditional DDPM for 1000 epochs and takes many GPU hours.
python web_ui/app.py --checkpoint outputs/ddpm_simple/best.pthOpen http://127.0.0.1:5000. The app offers fast (DDIM 20), balanced (DDIM 100) and full (DDPM 1000) sampling, keeps a gallery of past generations and shows the loaded checkpoint.
Weights are too large for git (120 to 470 MB each) and are published as assets of the v1.0.0 release. Fetch them with:
python scripts/download_checkpoints.py # DDPM used by the web UI (147 MB)
python scripts/download_checkpoints.py --all # everything below (~1 GB)The loader recognises all three U-Net generations from the weight shapes, so no config files are needed.
| Name | Lands in | Architecture | Schedule | Epochs | Notes |
|---|---|---|---|---|---|
ddpm_simple |
checkpoints/ddpm_simple/ddpm_epoch_50.pth |
simple U-Net, 12.8 M params |
linear | 50 | best image quality; default for the web UI |
ddpm_attention |
checkpoints/ddpm_attention/best_model.pth |
attention U-Net, 17.9 M |
cosine | 17 | EMA weights; training was stopped early |
ddpm_basic |
checkpoints/ddpm_basic/ddpm_epoch_50.pth |
basic U-Net, 10.9 M |
linear | 50 | first prototype; timestep not used by the network |
classifier |
checkpoints/classifier/best_production_model.pth |
EfficientNet-B2 + ResNet-50 + DenseNet-121 ensemble | see results below |
Try them right away:
python scripts/generate.py --ckpt checkpoints/ddpm_simple/ddpm_epoch_50.pth --num 9 --steps 100 --grid samples.png
python scripts/predict_pneumonia.py --ckpt checkpoints/classifier/best_production_model.pth some_xray.jpegChest X-ray DDPM. Samples and training curves for the simple model are in
docs/results (ddpm_simple_samples_epoch50.png, ddpm_simple_training_curves.png).
Pneumonia classifier (ensemble, Kaggle test split, 624 images; docs/results/classifier_metrics.json):
| Accuracy | Precision | Recall | F1 | ROC AUC |
|---|---|---|---|---|
| 0.707 | 0.944 | 0.564 | 0.706 | 0.897 |
The classifier is precise but conservative on the pneumonia class; see
docs/results/classifier_evaluation.png for the confusion matrix and ROC curve.
BloodMNIST augmentation study (ResNet-18 test accuracy, 8 classes):
| Training data | Accuracy |
|---|---|
| Original only | 0.930 |
| + traditional augmentation | 0.944 |
| + conditional GAN samples | 0.925 |
| + conditional DDPM samples | 0.947 |
Figures: docs/results/bloodmnist_augmentation_accuracy.png, bloodmnist_ddpm_samples.png,
bloodmnist_tsne.png.
simple(default): U-Net with channel multipliers 1-2-4-4, sinusoidal timestep embedding injected through FiLM-style scale/shift in every block, linear beta schedule.attention: residual U-Net with two residual blocks per level, self-attention at the 8x downsampled level and in the bottleneck, cosine schedule, EMA weights.basic: the first prototype. Its timestep MLP was never wired into the convolutions, so it acts as an unconditional denoiser. It is kept only so the early checkpoints remain loadable.
Released under the MIT License. The chest X-ray data belongs to its original authors (Kermany et al., via Kaggle) and is not redistributed here.