Flow matching
Kunna jämföra flow matching med diffusion och läsa aktuell litteratur.
Förkunskaper
Intuition
Diffusion definierar en stokastisk väg från data till brus och lär sig vända den. Flow matching definierar i stället en deterministisk väg och lär sig hastigheten längs den.
Tänk på det som skillnaden mellan att beskriva hur ett dammkorn slumpvandrar och att beskriva ett vindfält: det andra säger direkt vart varje punkt är på väg.
Den enklaste vägen är en rak linje:
där är brus och är data. Hastigheten längs den vägen är trivial:
Och det är hela träningsmålet: lär modellen att förutsäga givet och .
| Diffusion (DDPM) | Flow matching | |
|---|---|---|
| Väg | stokastisk, definierad av brusschema | deterministisk, ofta rak |
| Modellen lär sig | bruset | hastighetsfältet |
| Sampling | lös SDE, många steg | lös ODE, få steg |
| Schemaval | kritiskt | mindre känsligt |
| Matematik | krävande härledning | förvånansvärt enkel |
Härledning
Målet är att transportera en enkel fördelning (brus) till datafördelningen via ett tidsberoende hastighetsfält :
Det ideala fältet är det som genererar rätt fördelning vid varje — men det är oberäkneligt, eftersom det kräver marginalisering över all data.
Genombrottet (Lipman m.fl. 2023) är att man kan träna mot det betingade fältet i stället. För ett fixt par är den optimala vägen känd, och det visar sig att
Gradienterna är identiska. Man kan alltså optimera det enkla, betingade målet och få samma lösning som det oberäkneliga marginella.
Med den räta vägen blir målet:
Tre rader kod. Ingen ELBO, inget brusschema, ingen variansberäkning.
Rektifierade flöden. Räta vägar mellan slumpmässiga par korsar varandra, och det inlärda fältet blir krökt. Reflow rättar det: generera par med den tränade modellen, träna om på dem, och upprepa. Vägarna rätas ut, och efter ett par iterationer kan sampling göras i ett enda steg med bibehållen kvalitet.
Kopplingen till diffusion. Diffusionens probability flow ODE är ett specialfall: med ett visst val av interpolationsväg (variansbevarande i stället för rätlinjig) återfår man DDPM:s deterministiska sampler. Flow matching generaliserar alltså diffusion snarare än ersätter den — och det förklarar varför de flesta samplers går att använda på båda.
Var det används i dag: Stable Diffusion 3 och flera andra ledande bildmodeller bygger på rektifierade flöden, och metoden har spridit sig till ljud, video och molekylgenerering.
Kod
import torch, torch.nn as nn, torch.nn.functional as F
# Hela träningsmålet — tre rader
def flow_forlust(modell, x1):
"""x1: verklig data. x0: brus. Modellen ska förutsäga x1 - x0."""
x0 = torch.randn_like(x1)
t = torch.rand(x1.size(0), device=x1.device).view(-1, *([1] * (x1.dim() - 1)))
xt = (1 - t) * x0 + t * x1
return F.mse_loss(modell(xt, t.flatten()), x1 - x0)
# Sampling: lös ODE:n framåt. Euler räcker långt.
@torch.no_grad()
def sampla(modell, form, steg=20, device="cpu"):
x = torch.randn(form, device=device)
dt = 1.0 / steg
for i in range(steg):
t = torch.full((form[0],), i * dt, device=device)
x = x + modell(x, t) * dt
return x
# Midpoint-solver: halva antalet steg för samma kvalitet
@torch.no_grad()
def sampla_midpoint(modell, form, steg=10, device="cpu"):
x = torch.randn(form, device=device)
dt = 1.0 / steg
for i in range(steg):
t = torch.full((form[0],), i * dt, device=device)
v1 = modell(x, t)
v2 = modell(x + v1 * dt / 2, t + dt / 2)
x = x + v2 * dt
return x
# Reflow: räta ut vägarna genom att träna om på egna par
@torch.no_grad()
def generera_par(modell, n, form, steg=50, device="cpu"):
x0 = torch.randn(n, *form, device=device)
x = x0.clone()
dt = 1.0 / steg
for i in range(steg):
t = torch.full((n,), i * dt, device=device)
x = x + modell(x, t) * dt
return x0, x # kopplade par: brus och den bild det ledde till
def reflow_forlust(modell, x0, x1):
"""Samma mål, men med KOPPLADE par i stället för slumpmässiga."""
t = torch.rand(x1.size(0), device=x1.device).view(-1, *([1] * (x1.dim() - 1)))
xt = (1 - t) * x0 + t * x1
return F.mse_loss(modell(xt, t.flatten()), x1 - x0)
# Mät hur raka vägarna är — det avgör hur få steg som räcker
@torch.no_grad()
def rakhet(modell, form, steg=50, device="cpu"):
x0 = torch.randn(1, *form, device=device)
x, banlangd = x0.clone(), 0.0
dt = 1.0 / steg
for i in range(steg):
t = torch.full((1,), i * dt, device=device)
v = modell(x, t)
banlangd += float(v.norm()) * dt
x = x + v * dt
rak_langd = float((x - x0).norm())
return {"banlangd": round(banlangd, 3), "rakt_avstand": round(rak_langd, 3),
"rakhet": round(rak_langd / max(banlangd, 1e-9), 4)}
# rakhet nära 1.0 → vägarna är raka → ett fåtal steg räcker
rakhet är måttet som förklarar varför reflow fungerar. Före reflow ligger den kanske på 0,6; efter ett par iterationer nära 1,0 — och då ger ett enda Euler-steg nästan samma resultat som femtio.
Behärskning innebär
- Förklarar flow matching-målet
- Jämför med diffusion
- Vet vad rektifierade flöden ger
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Flow Matching for Generative Modeling — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Flow Straight and Fast: Learning to Generate and Transfer Data with Rectified Flow — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Scaling Rectified Flow Transformers for High-Resolution Image Synthesis — arXiv (öppen åtkomst; licens per artikel)