File size: 5,967 Bytes
fdeb032 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 | ---
license: other
license_name: celeba-non-commercial-research
license_link: https://mmlab.ie.cuhk.edu.hk/projects/CelebA.html
tags:
- diffusion
- ddpm
- ddim
- unconditional-image-generation
- from-scratch
- pytorch
- mps
datasets:
- flwrlabs/celeba
pipeline_tag: unconditional-image-generation
library_name: pytorch
---
# mini-diffusion β an 18.5M-parameter DDPM trained overnight on a laptop
A denoising diffusion model written from scratch in plain PyTorch. The U-Net, the noise schedule,
four samplers and the training loop are all in the repository β `diffusers` is never imported.
**Code: https://github.com/vous99/mini-diffusion**
Trained on one MacBook Pro (M5 Pro, 24 GB, MPS) in a single overnight run. Every number below was
measured on that run, not quoted from a paper.
## Results
| Metric | Value |
|---|---|
| **FID-tv** | **30.12** (10,000 samples, DDIM-50, EMA weights) |
| Steps | 22,098 |
| Images seen | 2.83M (17.4 epochs of CelebA train) |
| Best val loss | 0.0352 (`L_simple`, MSE on epsilon) |
| Training time | ~9 hours of compute at 87 images/s |
| Parameters | 18,538,371 |
FID over training: 73.35 (step 2,500) β 50.58 β 43.10 β 39.37 β 37.69 β 35.57 β 34.67 β 35.17
(step 20,000). Most of the gain lands in the first third of the run; the curve flattens after
roughly step 15,000.
**FID-tv is not the FID of the literature.** It uses torchvision's Inception-v3 weights rather than
the TensorFlow-ported weights every published FID is built on. The numbers here are internally
consistent β valid for comparing checkpoints, samplers and guidance scales β but not directly
comparable to a paper's "FID 3.5".
## What the model does
- Unconditional face generation at 64Γ64.
- Conditional generation on six CelebA attributes: `Male`, `Smiling`, `Young`, `Eyeglasses`,
`Blond_Hair`, `Bangs`.
- Classifier-free guidance from the same weights β trained with 15% condition dropout against a
learned null embedding, so one checkpoint serves both the conditional and unconditional branch.
- Deterministic DDIM with an `eta` parameter that reproduces ancestral DDPM exactly at `eta=1`.
- Karras sigma schedule with Heun's method.
- DDIM inversion, and therefore spherical interpolation between two real photographs.
## Usage
```bash
git clone https://github.com/vous99/mini-diffusion && cd mini-diffusion
pip install -r requirements.txt
python -c "
from huggingface_hub import hf_hub_download
import shutil, os
os.makedirs('ckpt', exist_ok=True)
shutil.copy(hf_hub_download('vous99/mini-diffusion', 'best.pt'), 'ckpt/best.pt')
"
python sample.py # an 8x8 grid
python sample.py --attrs "Male=1,Eyeglasses=1" # conditional
python sample.py --guidance-sweep # the guidance-scale figure
python sample.py --sampler-comparison # DDPM / DDIM / Heun
```
Sampling needs no dataset. Retraining does β `prepare_data.py` rebuilds it from the HuggingFace
CelebA shards.
## Architecture
U-Net over four resolutions, `base_ch=64`, `channel_mults=(1,2,2,4)` β 64@64Β² 128@32Β² 128@16Β²
256@8Β². Two ResBlocks per level down, three up, self-attention at 16Β² and 8Β² and in the middle
block, `head_dim=32`, GroupNorm with 32 groups.
Conditioning is a `Linear(6, 256)` projection added to the sinusoidal timestep embedding, plus a
learned null vector. Stable Diffusion instead cross-attends to a *sequence* of text tokens; six
fixed flags are not a sequence, so addition is the honest analogue β the same mechanism, a simpler
carrier.
The output convolution is zero-initialised, so the loss at step 0 is exactly `E||eps||Β² = 1.0000`.
That single number confirms the target is epsilon, the data is scaled to [-1,1], and the reduction
is a mean.
## Training details
| | |
|---|---|
| Schedule | cosine (Nichol & Dhariwal), T=1000 |
| Prediction target | epsilon |
| Timestep sampling | stratified over the batch, not i.i.d. uniform |
| Optimizer | AdamW, lr 2e-4, betas (0.9, 0.999), weight decay 0, grad clip 1.0 |
| LR schedule | linear warmup 500 steps β cosine decay to 2e-5 |
| Batch | 64 Γ 2 gradient accumulation = 128 effective |
| EMA | 0.999 with warmup |
| Dropout | 0.0 |
| Augmentation | horizontal flip, p=0.5 |
| Precision | bfloat16 autocast (GroupNorm stays fp32) |
## Limitations
- **64Γ64 only.** Faces at higher resolution are out of scope for this budget.
- **No text conditioning.** The condition is six binary flags, not a prompt.
- **Aligned frontal portraits only.** CelebA is centred faces; profiles, multiple people and full
bodies are out of distribution.
- **Roughly 15β20% of samples collapse into noise.** This is the dominant remaining defect.
- **Colour casts.** Even the cosine schedule ends at `abar_T β 2.4e-9` rather than 0, so the model
never trains on pure noise but is handed pure noise at sampling time. With epsilon-prediction
this shows up as a per-image brightness and colour bias. It receded substantially over training
but has not disappeared; v-prediction with zero terminal SNR is the proper fix.
- **This is 4.4% of the compute budget of the DDIM paper's CelebA-64 run.** The register is a good
2016 GAN, not "indistinguishable from a photograph".
## Licence and intended use
The weights inherit the **CelebA licence: non-commercial research use only**. This is a study
artifact for understanding how diffusion models work, not a product component.
It generates synthetic faces of people who do not exist. Do not use it to impersonate real people
or to produce material presented as a genuine photograph of anyone.
## Checkpoint contents
`best.pt` is a `torch.save` dict with `model` (raw weights), `ema` (averaged weights β what you
should sample from), `cfg` (the `UNetConfig`), `diffusion` (the `DiffusionConfig`), `attributes`
(the six conditioning names, in order), `iter` and `val_loss`. Load it with `checkpoint.py` from
the GitHub repository.
|