CIFAR-10 DDPM (18.7M params, from scratch)

A small denoising diffusion probabilistic model (DDPM, Ho et al. 2020) trained on CIFAR-10, written from scratch in ~250 lines of plain PyTorch.

Try it in your browser — the ONNX export runs client-side on WebGPU, generating an image in ~2 seconds on your own GPU.

samples

Uncurated 8×8 grid after 300 epochs.

Model details

  • U-Net, 18.7M parameters: three resolution levels (32→16→8), residual blocks with sinusoidal timestep embeddings, self-attention at 16×16 and 8×8, dropout 0.1
  • Linear beta schedule, T=1000, trained with the simplified noise-prediction MSE objective
  • EMA of weights (decay 0.9995) used for sampling; DDIM sampling with 50 steps
  • Unconditional — no class or text control
  • Trained 300 epochs (~117k steps, batch 128) in ~4 hours on a single RTX 3090; final loss ≈ 0.029

Files

File Description
checkpoint.pt Full PyTorch checkpoint: raw weights, EMA weights, optimizer state (resumable)
unet.onnx EMA weights exported to ONNX (fp32), verified to 4e-6 against PyTorch
train_diffusion.py Complete training script
sample.py Generate image grids from the checkpoint
export_onnx.py Reproduce the ONNX export

Usage

pip install torch torchvision
python sample.py --n 64 --seed 42          # sample a grid from checkpoint.pt
python train_diffusion.py --base 128 --epochs 300 --out out_big --sample-every 10 --ema-decay 0.9995   # retrain

Source and local web app: github.com/dannysheesh/cifar10-ddpm

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train PeterRabbit/cifar10-ddpm

Space using PeterRabbit/cifar10-ddpm 1

Paper for PeterRabbit/cifar10-ddpm