Flow matching
Be able to compare flow matching with diffusion and read the current literature.
Prerequisites
- GDiffusion models — the mathematicsrequired
Intuition
Diffusion defines a stochastic path from data to noise and learns to reverse it. Flow matching instead defines a deterministic path and learns the velocity along it.
Think of it as the difference between describing how a speck of dust random-walks and describing a wind field: the second says directly where every point is heading.
The simplest path is a straight line:
where is noise and is data. The velocity along that path is trivial:
And that is the whole training objective: teach the model to predict given and .
| Diffusion (DDPM) | Flow matching | |
|---|---|---|
| The path | stochastic, defined by a noise schedule | deterministic, often straight |
| What the model learns | the noise | the velocity field |
| Sampling | solve an SDE, many steps | solve an ODE, few steps |
| The choice of schedule | critical | less sensitive |
| The mathematics | a demanding derivation | surprisingly simple |
Derivation
The goal is to transport a simple distribution (noise) to the data distribution via a time-dependent velocity field :
The ideal field is the one that generates the right distribution at every — but it is intractable, since it requires marginalising over all the data.
The breakthrough (Lipman et al. 2023) is that you can train against the conditional field instead. For a fixed pair the optimal path is known, and it turns out that
The gradients are identical. You can therefore optimise the simple, conditional objective and get the same solution as the intractable marginal one.
With the straight path the objective becomes:
Three lines of code. No ELBO, no noise schedule, no variance computation.
Rectified flows. Straight paths between random pairs cross each other, and the learnt field becomes curved. Reflow corrects that: generate pairs with the trained model, retrain on them, and repeat. The paths straighten out, and after a couple of iterations sampling can be done in a single step with the quality retained.
The connection to diffusion. Diffusion's probability flow ODE is a special case: with a particular choice of interpolation path (variance-preserving instead of straight-line) you recover DDPM's deterministic sampler. Flow matching therefore generalises diffusion rather than replacing it — and that explains why most samplers can be used on both.
Where it is used today: Stable Diffusion 3 and several other leading image models are built on rectified flows, and the method has spread to audio, video and molecule generation.
Code
import torch, torch.nn as nn, torch.nn.functional as F
# The whole training objective — three lines
def flow_loss(model, x1):
"""x1: real data. x0: noise. The model is to predict 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(model(xt, t.flatten()), x1 - x0)
# Sampling: solve the ODE forwards. Euler goes a long way.
@torch.no_grad()
def sample(model, shape, steps=20, device="cpu"):
x = torch.randn(shape, device=device)
dt = 1.0 / steps
for i in range(steps):
t = torch.full((shape[0],), i * dt, device=device)
x = x + model(x, t) * dt
return x
# A midpoint solver: half the number of steps for the same quality
@torch.no_grad()
def sample_midpoint(model, shape, steps=10, device="cpu"):
x = torch.randn(shape, device=device)
dt = 1.0 / steps
for i in range(steps):
t = torch.full((shape[0],), i * dt, device=device)
v1 = model(x, t)
v2 = model(x + v1 * dt / 2, t + dt / 2)
x = x + v2 * dt
return x
# Reflow: straighten the paths out by retraining on your own pairs
@torch.no_grad()
def generate_pairs(model, n, shape, steps=50, device="cpu"):
x0 = torch.randn(n, *shape, device=device)
x = x0.clone()
dt = 1.0 / steps
for i in range(steps):
t = torch.full((n,), i * dt, device=device)
x = x + model(x, t) * dt
return x0, x # coupled pairs: the noise and the image it led to
def reflow_loss(model, x0, x1):
"""The same objective, but with COUPLED pairs instead of random ones."""
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(model(xt, t.flatten()), x1 - x0)
# Measure how straight the paths are — that decides how few steps suffice
@torch.no_grad()
def straightness(model, shape, steps=50, device="cpu"):
x0 = torch.randn(1, *shape, device=device)
x, path_length = x0.clone(), 0.0
dt = 1.0 / steps
for i in range(steps):
t = torch.full((1,), i * dt, device=device)
v = model(x, t)
path_length += float(v.norm()) * dt
x = x + v * dt
straight_length = float((x - x0).norm())
return {"path_length": round(path_length, 3), "straight_distance": round(straight_length, 3),
"straightness": round(straight_length / max(path_length, 1e-9), 4)}
# a straightness near 1.0 → the paths are straight → a handful of steps suffice
straightness is the measure that explains why reflow works. Before reflow it might be around 0.6; after a couple of iterations near 1.0 — and then a single Euler step gives almost the same result as fifty.
Mastery means
- Explains the flow matching objective
- Compares it with diffusion
- Knows what rectified flows give
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Flow Matching for Generative Modeling — arXiv (open access; licence per article)
- arXiv — Flow Straight and Fast: Learning to Generate and Transfer Data with Rectified Flow — arXiv (open access; licence per article)
- arXiv — Scaling Rectified Flow Transformers for High-Resolution Image Synthesis — arXiv (open access; licence per article)