Hoppa till innehållet
AI-grafen
G· Frontier Labgenerativa-modeller· ca 120 min· volatil — kontrolleras ofta· verifierad 2026-09-21

Diffusionsmodeller — matematiken

Kunna härleda brusschemat och träningsmålet för DDPM.

Förkunskaper

Intuition

Diffusion vilar på en obalans: att förstöra en bild är trivialt, att återskapa den är svårt. Så man definierar förstörelsen exakt och lär modellen att vända den, ett litet steg i taget.

x₀ ──brus──→ x₁ ──brus──→ ... ──brus──→ x_T ≈ rent brus
   ←─avbrusa──   ←─avbrusa──    ←─avbrusa──

Tre saker gör konstruktionen praktisk:

  1. Framåtprocessen har sluten form — xtx_t kan beräknas direkt ur x0x_0 utan att iterera.
  2. Träningsmålet reduceras till enkel MSE på det tillsatta bruset.
  3. Varje bakåtsteg är litet, så en gaussisk approximation räcker.

Det är det andra som är förvånande: efter en ganska lång härledning står man med mse_loss(modellens_gissning, det_faktiska_bruset) — och inget mer.

Härledning

Framåtprocessen definieras som en Markovkedja:

q(xt∣xt−1)=N ⁣(xt; 1−βt xt−1, βtI)q(x_t \mid x_{t-1}) = \mathcal{N}\!\left(x_t;\ \sqrt{1-\beta_t}\,x_{t-1},\ \beta_t I\right)

med ett litet brusschema β1<⋯<βT\beta_1 < \dots < \beta_T (typiskt linjärt från 10−410^{-4} till 0,020{,}02 över T=1000T = 1000 steg).

Sluten form. Sätt αt=1−βt\alpha_t = 1-\beta_t och αˉt=∏s≤tαs\bar\alpha_t = \prod_{s\le t}\alpha_s. Eftersom summan av två oberoende gaussiska variabler är gaussisk med adderade varianser kan stegen slås ihop:

q(xt∣x0)=N ⁣(xt; αˉt x0, (1−αˉt)I)q(x_t \mid x_0) = \mathcal{N}\!\left(x_t;\ \sqrt{\bar\alpha_t}\,x_0,\ (1-\bar\alpha_t)I\right) ⟹xt=αˉt x0+1−αˉt ε,ε∼N(0,I)\Longrightarrow\quad x_t = \sqrt{\bar\alpha_t}\,x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon, \qquad \varepsilon \sim \mathcal{N}(0, I)

Det här är nyckeln till att träningen är praktisk: man kan hoppa direkt till ett slumpmässigt tt utan att simulera kedjan.

Bakåtsteget. Den sanna posteriorn är oberäknelig, men betingat på x0x_0 är den gaussisk i sluten form:

q(xt−1∣xt,x0)=N ⁣(xt−1; μ~t(xt,x0), β~tI)q(x_{t-1}\mid x_t, x_0) = \mathcal{N}\!\left(x_{t-1};\ \tilde\mu_t(x_t, x_0),\ \tilde\beta_t I\right)

μ~t=αˉt−1βt1−αˉtx0+αt(1−αˉt−1)1−αˉtxt,β~t=1−αˉt−11−αˉtβt\tilde\mu_t = \frac{\sqrt{\bar\alpha_{t-1}}\beta_t}{1-\bar\alpha_t}x_0 + \frac{\sqrt{\alpha_t}(1-\bar\alpha_{t-1})}{1-\bar\alpha_t}x_t, \qquad \tilde\beta_t = \frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}\beta_t

Förenklingen. Variationsgränsen (ELBO) ger en summa av KL-termer mellan gaussiska fördelningar. Ho m.fl. (2020) visade att om man parametriserar modellen så att den förutsäger bruset i stället för medelvärdet, och sedan släpper viktningsfaktorerna, reduceras hela målet till

Lsimple=Ex0,ε,t[∥ε−εθ ⁣(αˉtx0+1−αˉtε, t)∥2]\boxed{\mathcal{L}_{\text{simple}} = \mathbb{E}_{x_0,\varepsilon,t}\left[\left\|\varepsilon - \varepsilon_\theta\!\left(\sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\varepsilon,\ t\right)\right\|^2\right]}

De bortsläppta vikterna gör att målet inte längre är en exakt ELBO — men i praktiken tränar det bättre, eftersom viktningen annars ger de mest brusiga stegen oproportionerligt stor vikt.

Samplingssteget blir då:

xt−1=1αt(xt−βt1−αˉtεθ(xt,t))+σtz,z∼N(0,I)x_{t-1} = \frac{1}{\sqrt{\alpha_t}}\left(x_t - \frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\varepsilon_\theta(x_t,t)\right) + \sigma_t z, \qquad z\sim\mathcal{N}(0,I)

Kopplingen till poängbaserade modeller: ∇xlog⁡q(xt)=−εθ(xt,t)/1−αˉt\nabla_{x}\log q(x_t) = -\varepsilon_\theta(x_t,t)/\sqrt{1-\bar\alpha_t}. Att förutsäga bruset är alltså samma sak som att skatta score-funktionen, och DDPM och score matching är två vyer av samma modell. Song m.fl. (2021) visade att båda är diskretiseringar av en stokastisk differentialekvation, vilket förklarar varför olika samplers kan bytas fritt efter träning.

Kod

import torch, torch.nn.functional as F

T = 1000
beta = torch.linspace(1e-4, 0.02, T)
alfa = 1.0 - beta
alfa_bar = torch.cumprod(alfa, dim=0)

def q_sample(x0, t, brus=None):
    """Sluten form: hoppa direkt till steg t utan att iterera."""
    brus = torch.randn_like(x0) if brus is None else brus
    a = alfa_bar[t].view(-1, 1, 1, 1)
    return a.sqrt() * x0 + (1 - a).sqrt() * brus, brus

def forlust(modell, x0):
    t = torch.randint(0, T, (x0.size(0),), device=x0.device)
    xt, brus = q_sample(x0, t)
    return F.mse_loss(modell(xt, t), brus)        # hela träningsmålet

@torch.no_grad()
def sampla(modell, form, device="cpu"):
    x = torch.randn(form, device=device)
    for t in reversed(range(T)):
        tt = torch.full((form[0],), t, device=device, dtype=torch.long)
        eps = modell(x, tt)
        a, ab = alfa[t], alfa_bar[t]
        medel = (x - beta[t] / (1 - ab).sqrt() * eps) / a.sqrt()
        x = medel + beta[t].sqrt() * torch.randn_like(x) if t > 0 else medel
    return x

# Kontrollera den slutna formen numeriskt mot den iterativa definitionen
torch.manual_seed(0)
x0 = torch.randn(1, 1, 8, 8)
t_mal = 200

# Iterativt
x = x0.clone()
for t in range(t_mal + 1):
    x = alfa[t].sqrt() * x + beta[t].sqrt() * torch.randn_like(x)
var_iterativ = float(x.var())

# Sluten form — 4000 stickprov för att jämföra fördelningarna
prover = torch.stack([q_sample(x0, torch.tensor([t_mal]))[0] for _ in range(4000)])
print(f"teoretisk varians: {float(1 - alfa_bar[t_mal]):.4f}")
print(f"empirisk varians:  {float((prover - alfa_bar[t_mal].sqrt() * x0).var()):.4f}")

# Brusschemat: hur mycket signal återstår vid varje steg
for t in (0, 100, 300, 500, 800, 999):
    print(f"  t={t:>4}: signal {float(alfa_bar[t].sqrt()):.4f}  "
          f"brus {float((1 - alfa_bar[t]).sqrt()):.4f}")
#   t=   0: signal 0.9999  brus 0.0100
#   t= 500: signal 0.3403  brus 0.9403
#   t= 999: signal 0.0063  brus 1.0000    ← praktiskt taget rent brus

# Kosinusschema (Nichol & Dhariwal): behåller mer signal längre
def alfa_bar_kosinus(t, T=1000, s=0.008):
    import math
    f = lambda u: math.cos((u / T + s) / (1 + s) * math.pi / 2) ** 2
    return f(t) / f(0)

for t in (0, 300, 500, 800):
    print(f"  t={t:>3}: linjärt {float(alfa_bar[t]):.4f}  "
          f"kosinus {alfa_bar_kosinus(t):.4f}")

Behärskning innebär

  • Härleder den slutna formen för framåtprocessen
  • Förklarar varför träningsmålet blir enkel MSE
  • Tolkar samplingssteget

Logga in för att göra övningarna och bygga upp din behärskning.

Källor

Alla källor och licenser