HoumanAHuman/conditional-diffusion-mnist
0
Conditional Diffusion on Modified MNIST
This project trains a class-conditional diffusion model on a modified MNIST dataset where class 1 includes samples from both MNIST digit 1 and FashionMNIST class 'trouser'.
๐ฆ Setup
git clone https://github.com/yourusername/conditional-diffusion-mnist.git
cd conditional-diffusion-mnist
pip install -r requirements.txt๐ Dataset
Prepare the custom dataset:
python prepare_dataset.pyThis generates Dataset/shuffled_mnist_with_trousers.pt, a balanced dataset of MNIST 1s and FashionMNIST trousers labeled as 1.
๐งจ Train the Diffusion Model
python train_diffusion.pyThis trains a conditional DDPM model and saves checkpoints under the DDPM/ directory, including unet_final.pt.
๐จ Sample from the Model
python sample_images.pyThis loads the trained model from DDPM/unet_final.pt and generates class-conditioned samples (by default, class 1).
๐ Folder Structure
conditional-diffusion-mnist/
โโโ DDPM/ # Model checkpoints
โ โโโ unet_final_ema.pt
โ โโโ class_embedder.pt
โโโ Samples/ # Sample generation
โ โโโ class {condition_class}/ # Samples from the original model with respective conditioned class
โ โโโ class {condition_class}-Unlearned with constant lambda/ # Samples from the unlearned (constant lambda) model with respective conditioned class
โ โโโ class {condition_class}-Unlearned with dynamic lambda/ # Samples from the unlearned (dynamic lambda) model with respective conditioned class
โโโ DDPM_Unlearned/ # Model checkpoints after Unlearning with constant lambda
โ โโโ unet_unlearned_ema.pt
โ โโโ class_embedder_unlearned_ema.pt
โโโ DDPM_Unlearned_dynamic/ # Model checkpoints after Unlearning with dynamic lambda
โ โโโ unet_unlearned_dynamic_ema.pt
โ โโโ class_embedder_unlearned_dynamic_ema.pt
โโโ Dataset/ # Custom dataset generation and files
โ โโโ shuffled_augmented_mnist.pt
โ โโโ original_mnist.pt
โ โโโ trousers_subset.pt
โโโ train_diffusion.py # Training script
โโโ preprocess_dataset.py # Dataset creation script
โโโ requirements.txt
โโโ README.md๐งช Requirements
Install all dependencies using:
pip install -r requirements.txt๐ License
MIT License
