Diffusionsmodeller
Kunna förklara framåt- och bakåtprocessen, träna en liten diffusionsmodell och generera bilder.
Förkunskaper
- EFaltningsnät (CNN)underlättar
- EGenerativa modeller — översiktkrävs
Intuition
Framåtprocessen förstör en bild: lägg till lite gaussiskt brus, om och om igen, i T steg tills bara brus återstår. Den är fast och kräver ingen inlärning.
Bakåtprocessen är det modellen lär sig: givet en brusig bild vid steg t, förutsäg bruset som lades till. Kan man det, kan man ta ett steg tillbaka mot en renare bild. Upprepa från rent brus och en bild växer fram.
Träningen är förvånansvärt enkel: dra en bild, dra ett slumpmässigt t, lägg på motsvarande brus, låt nätet gissa bruset, och mät med MSE. Ingen diskriminator, ingen instabilitet — det är därför diffusion slog GAN.
Formellt
Framåt: . Med och kan man hoppa direkt till valfritt steg:
Träningsmål (Ho m.fl. 2020) — förvånansvärt nog bara en MSE:
Sampling: starta i och iterera bakåt. DDPM använder alla T steg (ofta 1 000); DDIM gör processen deterministisk och klarar 20–50 steg med nästan samma kvalitet.
Classifier-free guidance: träna modellen både med och utan textvillkor (villkoret slumpas bort i ~10 % av fallen) och extrapolera vid sampling: Högre ger starkare följsamhet mot prompten men mindre variation — den parameter användare känner som «guidance scale».
Latent diffusion (Stable Diffusion) kör hela processen i ett komprimerat latent rum från en autoencoder, vilket sänker beräkningen ~50× och är varför bildgenerering går på konsumenthårdvara.
Kod
import torch, torch.nn as nn
T = 1000
betas = torch.linspace(1e-4, 0.02, T)
alphas = 1.0 - betas
abar = torch.cumprod(alphas, dim=0)
def brusa(x0, t, eps=None):
"""Hoppa direkt till steg t."""
eps = torch.randn_like(x0) if eps is None else eps
a = abar[t].view(-1, 1, 1, 1)
return a.sqrt() * x0 + (1 - a).sqrt() * eps, eps
def traningssteg(modell, x0, opt):
t = torch.randint(0, T, (x0.size(0),))
xt, eps = brusa(x0, t)
loss = nn.functional.mse_loss(modell(xt, t), eps) # förutsäg bruset
loss.backward(); opt.step(); opt.zero_grad()
return loss.item()
@torch.no_grad()
def sampla(modell, form, steg=50):
x = torch.randn(form)
for t in reversed(range(0, T, T // steg)): # DDIM-liknande hopp
tt = torch.full((form[0],), t, dtype=torch.long)
eps = modell(x, tt)
a = abar[t]
x0_hat = (x - (1 - a).sqrt() * eps) / a.sqrt()
a_prev = abar[max(t - T // steg, 0)]
x = a_prev.sqrt() * x0_hat + (1 - a_prev).sqrt() * eps
return x
Behärskning innebär
- Förklarar framåt- och bakåtprocessen
- Beskriver träningsmålet (förutsäg bruset)
- Förstår styrning och samplingssteg
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 — High-Resolution Image Synthesis with Latent Diffusion Models — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Classifier-Free Diffusion Guidance — arXiv (öppen åtkomst; licens per artikel)