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:
- Framåtprocessen har sluten form — kan beräknas direkt ur utan att iterera.
- Träningsmålet reduceras till enkel MSE på det tillsatta bruset.
- 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:
med ett litet brusschema (typiskt linjärt från till över steg).
Sluten form. Sätt och . Eftersom summan av två oberoende gaussiska variabler är gaussisk med adderade varianser kan stegen slås ihop:
Det här är nyckeln till att träningen är praktisk: man kan hoppa direkt till ett slumpmässigt utan att simulera kedjan.
Bakåtsteget. Den sanna posteriorn är oberäknelig, men betingat på är den gaussisk i sluten form:
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
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å:
Kopplingen till poängbaserade modeller: . 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
- arXiv — Denoising Diffusion Probabilistic Models — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Score-Based Generative Modeling through Stochastic Differential Equations — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Improved Denoising Diffusion Probabilistic Models — arXiv (öppen åtkomst; licens per artikel)