Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

MedSynth: synthetic chest X-rays with diffusion models

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 chest X-rays sampled with DDIM (100 steps)
Synthetic 256x256 chest X-rays from the DDPM (DDIM, 100 steps).

Repository layout

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

Setup

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.txt

The scripts work straight from the repository root. To use medsynth from elsewhere run pip install -e ..

Data

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.

Usage

1. Train the diffusion model

# 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_attention

Every epoch writes last.pth (resumable with --resume), best.pth (lowest validation loss), history.csv, training_curves.png and periodic sample grids.

2. Generate images

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.

3. Evaluate generation 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/eval

Writes 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.

4. Train the pneumonia classifier

# 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.jpeg

Training 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.

5. Augmentation study on BloodMNIST

python experiments/bloodmnist_augmentation_comparison.py --out-dir outputs/bloodmnist
python experiments/bloodmnist_augmentation_comparison.py --quick    # a few minutes, for a smoke test

The full run trains the conditional DDPM for 1000 epochs and takes many GPU hours.

6. Web UI

python web_ui/app.py --checkpoint outputs/ddpm_simple/best.pth

Open 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.

Trained checkpoints

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.jpeg

Results

Chest 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.

Notes on the model variants

  • 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.

License

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.

About

Synthetic chest X-ray generation with DDPMs, pneumonia classification, and a GAN/DDPM augmentation study (final-year project)

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages