Implementation of Neural Diffusion Models with support for both DDPM and NDM training. Built with Hydra for flexible configuration.
- DDPM & NDM: Both denoising diffusion and neural diffusion models
- Multiple architectures: MLP (2D toys), SimpleMLP, CUNet (images)
- Multiple datasets: Checkerboard, moons, circles, swiss roll, MNIST, Fashion-MNIST
- Hydra config: Easy experiment management with CLI overrides
- WandB integration: Optional experiment tracking with visualizations
# Using pixi (recommended)
pixi install
# Or using pip
pip install torch torchvision hydra-core omegaconf wandb scikit-learn tqdm matplotlibFor CUNet architecture, clone the EDM repo:
git clone https://github.com/NVlabs/edm# Train DDPM on checkerboard (default)
python train.py
# Train on MNIST with CUNet
python train.py dataset=mnist architecture=cunet
# Train NDM model
python train.py model=ndm
# Enable WandB logging
python train.py logging.wandb=true.
├── train.py # Main training script (Hydra-based)
├── diffusion_model.py # DDPM and NDM implementations
├── visualization.py # Dataset-specific visualizations
├── models/
│ ├── mlp_2d.py # MLP for 2D toy datasets
│ ├── simple_mlp.py # Configurable simple MLP
│ ├── cunet.py # Conditional UNet for images
│ └── edm.py # EDM integration wrapper
├── datasets/
│ ├── toy.py # 2D datasets (checkerboard, moons, etc.)
│ └── images.py # Image datasets (MNIST, Fashion-MNIST)
└── configs/
├── config.yaml # Main config with defaults
├── model/ # ddpm.yaml, ndm.yaml
├── architecture/ # mlp_2d.yaml, simple_mlp.yaml, cunet.yaml
├── dataset/ # checkerboard.yaml, mnist.yaml, etc.
└── scheduler/ # cosine.yaml, linear.yaml, step.yaml
# Change hyperparameters
python train.py training.batch_size=512 training.num_epochs=500
# Change optimizer
python train.py optimizer.lr=1e-3 optimizer.name=adam
# Change LR scheduler
python train.py scheduler=linear
# Change sampling settings
python train.py eval.num_steps=200 eval.use_sde=false| Category | Options |
|---|---|
| Models | ddpm, ndm |
| Architectures | mlp_2d, simple_mlp, cunet |
| Datasets | checkerboard, moons, circles, swiss_roll, mnist, fashion_mnist |
| Schedulers | cosine, linear, step |
| Optimizers | adamw, adam, sgd |
Training outputs are saved to outputs/YYYY-MM-DD/HH-MM-SS/:
best_model.pt- Best model checkpointcheckpoint_epoch_N.pt- Periodic checkpointssamples_epoch_N.pt- Generated samplestrain.log- Hydra logs