Skip to content

Repository files navigation

Neural Diffusion Models (NDM)

Implementation of Neural Diffusion Models with support for both DDPM and NDM training. Built with Hydra for flexible configuration.

Features

  • 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

Installation

# Using pixi (recommended)
pixi install

# Or using pip
pip install torch torchvision hydra-core omegaconf wandb scikit-learn tqdm matplotlib

For CUNet architecture, clone the EDM repo:

git clone https://github.com/NVlabs/edm

Quick Start

# 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

Project Structure

.
├── 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

Configuration

CLI Overrides

# 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

Available Options

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

Outputs

Training outputs are saved to outputs/YYYY-MM-DD/HH-MM-SS/:

  • best_model.pt - Best model checkpoint
  • checkpoint_epoch_N.pt - Periodic checkpoints
  • samples_epoch_N.pt - Generated samples
  • train.log - Hydra logs

References

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages