Elucidating the Design Space of Diffusion-Based Generative Models

Tero Karras, Miika Aittala, Timo Aila, Samuli Laine

2022 · NeurIPS

Elucidating the Design Space of Diffusion-Based Generative Models

Problem

Framing

Diffusion models mixed solver choice, noise schedule, stochasticity, and network parameterization into monolithic recipes. This paper factorizes that design space, then replaces each weak choice with a better one: Heun-based sampling, tuned noise schedules, and score-network preconditioning. It reaches FID 1.79 on class-conditional CIFAR-10 and 1.36 on ImageNet-64 with 35 NFE.

Currently Used Methods

Foundational

Proposed Method

Architecture

EDM keeps the backbone family unchanged: DDPM++ for VP, NCSN++ for VE, and ADM for ImageNet. The main architectural change is external preconditioning: the denoiser is expressed as a skip-connected wrapper around a noise-conditioned backbone FθF_\theta.

Dθ(x;σ)=cskip(σ)x+cout(σ)Fθ ⁣(cin(σ)x;cnoise(σ))D_\theta(\mathbf{x}; \sigma) = c_{\mathrm{skip}}(\sigma)\,\mathbf{x} + c_{\mathrm{out}}(\sigma)\,F_\theta\!\left(c_{\mathrm{in}}(\sigma)\,\mathbf{x}; c_{\mathrm{noise}}(\sigma)\right)

Loss / Objective

The objective trains the preconditioned denoiser over noise levels sampled from a log-normal distribution.

L(θ)=Ey,n,σ[λ(σ)Dθ(y+n;σ)y2],ypdata,  nN(0,σ2I)\mathcal{L}(\theta) = \mathbb{E}_{\mathbf{y},\mathbf{n},\sigma}\left[\lambda(\sigma)\,\left\| D_\theta(\mathbf{y}+\mathbf{n};\sigma)-\mathbf{y} \right\|^2\right],\quad \mathbf{y}\sim p_{\mathrm{data}},\; \mathbf{n}\sim \mathcal{N}(\mathbf{0},\sigma^2\mathbf{I}) cskip(σ)=σdata2σ2+σdata2,cout(σ)=σσdataσ2+σdata2,cin(σ)=1σ2+σdata2,λ(σ)=1cout(σ)2=σ2+σdata2σ2σdata2c_{\mathrm{skip}}(\sigma)=\frac{\sigma_{\mathrm{data}}^2}{\sigma^2+\sigma_{\mathrm{data}}^2},\quad c_{\mathrm{out}}(\sigma)=\frac{\sigma\,\sigma_{\mathrm{data}}}{\sqrt{\sigma^2+\sigma_{\mathrm{data}}^2}},\quad c_{\mathrm{in}}(\sigma)=\frac{1}{\sqrt{\sigma^2+\sigma_{\mathrm{data}}^2}},\quad \lambda(\sigma)=\frac{1}{c_{\mathrm{out}}(\sigma)^2}=\frac{\sigma^2+\sigma_{\mathrm{data}}^2}{\sigma^2\sigma_{\mathrm{data}}^2}

Sampling Rule / Algorithm

For deterministic sampling, EDM integrates the probability-flow ODE with Heun's second-order method on a noise schedule {σi}\{\sigma_i\}.

x˙=σ˙(t)σ(t)(xDθ(x;σ(t)))\dot{\mathbf{x}} = \frac{\dot{\sigma}(t)}{\sigma(t)}\left(\mathbf{x}-D_\theta(\mathbf{x};\sigma(t))\right) x^i+1=xi+(ti+1ti)di,di=xiDθ(xi;ti)ti,\hat{\mathbf{x}}_{i+1}=\mathbf{x}_i+(t_{i+1}-t_i)\,d_i, \qquad d_i=\frac{\mathbf{x}_i-D_\theta(\mathbf{x}_i;t_i)}{t_i}, di=x^i+1Dθ(x^i+1;ti+1)ti+1,xi+1=xi+(ti+1ti)(di+di2)d'_i=\frac{\hat{\mathbf{x}}_{i+1}-D_\theta(\hat{\mathbf{x}}_{i+1};t_{i+1})}{t_{i+1}}, \qquad \mathbf{x}_{i+1}=\mathbf{x}_i+(t_{i+1}-t_i)\left(\frac{d_i+d'_i}{2}\right)

EDM then adds optional stochastic "churn" between selected noise levels to improve low-NFE sampling.

Training Procedure

Evaluation

Datasets

Metrics

Headline results

Results plot: three FID-vs-NFE panels compare deterministic, stochastic, and prior samplers on unconditional CIFAR-10 and class-conditional ImageNet-64.

Ablations

Method Strengths and Weaknesses

Strengths

Weaknesses

Suggestions from the authors

Links

Prior Papers

Further Papers