holehouse.org Blog Machine learning notes

23: Diffusion Models

A note on this chapter

The idea

image noisier noisier noise forward - add noise (fixed, not learned) reverse - learned denoiser, run to generate
Destroy structure in small fixed steps; learn to undo one of them; run it backwards to generate.

The forward process

q(xt∣xt−1)=N(xt;1−βtxt−1,βtI) One forward step: a Gaussian (chapter 15) centred just below the current value, with variance βt. Nothing is learned here — the destruction is by design.
xt=α¯tx0+1−α¯tε,ε∼N(0,I) α¯t=∏s=1t(1−βs) ᾱt is the surviving fraction of the original signal. It starts near 1 and decays towards 0, so xT is (nearly) pure standard Gaussian noise whatever x0 was — every data point is destroyed to the same known distribution, which is exactly what lets generation start from that distribution.

Learning to reverse it

J(θ)=E[‖ε−εθ(xt,t)‖2] The entire training objective is squared error — chapter 02's cost function, with the "right answer" being the noise we ourselves added. No adversary, no likelihood gymnastics; this simplicity is a large part of why diffusion took over.
xt−1=11−βt(xt−βt1−α¯tεθ(xt,t))+βtz,z∼N(0,I) One reverse step. The fresh noise z looks paradoxical — we are trying to remove noise — but it is what makes the reverse process a distribution rather than a single trajectory: run it twice and you get two different samples. At the final step (t = 1) no z is added.

A complete diffusion model in NumPy

import numpy as np

rng = np.random.default_rng(0)

def sample_data(n):                # half near -2, half near +2
    modes = rng.choice([-2.0, 2.0], size=n)
    return modes + 0.3 * rng.standard_normal(n)

T = 50
beta = np.linspace(1e-3, 0.2, T)   # the noise schedule
alpha = 1 - beta
alpha_bar = np.cumprod(alpha)      # alpha_bar[-1] = 0.0045: 0.45% of signal left

x0  = sample_data(10000)
eps = rng.standard_normal(10000)
xT  = np.sqrt(alpha_bar[-1]) * x0 + np.sqrt(1 - alpha_bar[-1]) * eps
xT.mean(), xT.std()                # (0.001, 1.0) - the data is gone;
                                   # every x0 ends as standard normal noise

The closed-form jump verified: after 50 steps the bimodal data is indistinguishable from N(0, 1), which is exactly what the ᾱt equation promised.

H = 64
W1 = rng.standard_normal((H, 2)) * 0.5   # chapter 09: random init, small values
b1 = np.zeros(H)
W2 = rng.standard_normal(H) * 0.5
b2 = 0.0

def predict(x, t_frac):
    a1 = np.tanh(W1 @ np.vstack([x, t_frac]) + b1[:, None])   # hidden layer
    return W2 @ a1 + b2, a1

εθ(xt, t) as sixty-four hidden units. Feeding t in as an input is what lets one network learn to denoise at every noise level at once.

lr = 1e-2
for step in range(4000):                      # mini-batch descent, chapter 17
    x0  = sample_data(256)
    t   = rng.integers(0, T, size=256)        # a random step per example
    eps = rng.standard_normal(256)
    xt  = np.sqrt(alpha_bar[t]) * x0 + np.sqrt(1 - alpha_bar[t]) * eps

    eps_hat, a1 = predict(xt, t / T)
    d = eps_hat - eps                         # the error to send backwards

    gW2 = a1 @ d / len(d)                     # backprop, chapter 09 in miniature:
    gb2 = d.mean()                            # output-layer gradients...
    da1 = np.outer(W2, d) * (1 - a1 ** 2)     # ...delta for the hidden layer...
    gW1 = da1 @ np.vstack([xt, t / T]).T / len(d)
    gb1 = da1.mean(axis=1)                    # ...and its gradients

    W1 -= lr * gW1; b1 -= lr * gb1            # simultaneous update, chapter 02
    W2 -= lr * gW2; b2 -= lr * gb2
# training loss falls from 4.1 to ~0.35 over the 4000 steps (under a second)

The whole training loop. Note what it never does: it never sees a "generated sample is good/bad" signal. It only ever learns to predict added noise, and generation quality follows from that alone.

n = 4000
x = rng.standard_normal(n)                    # x_T: pure noise, no data in sight
for t in range(T - 1, -1, -1):
    eps_hat, _ = predict(x, np.full(n, t / T))
    x = (x - beta[t] / np.sqrt(1 - alpha_bar[t]) * eps_hat) / np.sqrt(alpha[t])
    if t > 0:
        x = x + np.sqrt(beta[t]) * rng.standard_normal(n)   # the fresh z

(x < 0).mean()                     # 0.517          - half in each mode, as in the data
x[x < 0].mean(), x[x >= 0].mean()  # (-1.89, 1.87)  - mode centres (true: -2, +2)
x[x < 0].std(),  x[x >= 0].std()   # (0.43, 0.47)   - mode widths  (true: 0.3)
(np.abs(x) < 1).mean()             # 0.054          - the gap is nearly empty

Four thousand samples of standard normal noise, pushed backwards through the learned denoiser, come out bimodal: two clean modes at ±1.9 with an almost-empty gap between them. Slightly blurrier than the truth (widths 0.45 against 0.3) — about right for a sixty-four-unit network — but this is real generation: a distribution nothing in the sampling loop ever saw, reconstructed from noise.

Conditioning and guidance

What they're used for

Summary