Flow Matching
Apprendre un champ de vélocité qui transforme du bruit en données via une trajectoire continue, alternative plus rapide aux modèles de diffusion.
L'idée centrale
Au lieu d'ajouter puis retirer du bruit par étapes discrètes (diffusion), Flow Matching apprend le champ de vitesse v(x, t) qui pousse directement les points d'une distribution source p₀ (bruit gaussien) vers une distribution cible p₁ (données réelles).
ODE qui gouverne le flux
dx(t)/dt = v(x(t), t)
À l'inférence, on intègre cette ODE de t=0 à t=1 pour générer un échantillon.
Entraînement
Paires et interpolation
Pour chaque paire (x₀, x₁) où x₀ ~ N(0, I) et x₁ vient des données :
# Interpolation linéaire
x(t) = (1 - t) * x₀ + t * x₁
# Vélocité de référence (constante)
v_gt = x₁ - x₀
Loss de régression simple
import torch
def flow_matching_loss(model, x0, x1):
t = torch.rand(x0.shape[0], 1) # t ~ Uniform[0,1]
x_t = (1 - t) * x0 + t * x1
v_gt = x1 - x0
v_pred = model(x_t, t)
return ((v_pred - v_gt) ** 2).mean()
Pas de schedule de bruit, pas de KL : juste un MSE sur la vélocité.
Inférence (sampling)
Intégration ODE
from torchdiffeq import odeint
# x0 : bruit initial [batch, dim]
x0 = torch.randn(batch_size, dim)
# Résoudre dx/dt = v_θ(x, t)
def ode_func(t, x):
return model(x, t)
t_span = torch.linspace(0, 1, steps=20)
trajectory = odeint(ode_func, x0, t_span, method='dopri5')
x1 = trajectory[-1] # Échantillon final
Méthode d'Euler simple (10-50 steps suffisent souvent)
x = torch.randn(batch_size, dim)
dt = 1.0 / n_steps
for i in range(n_steps):
t = i * dt
v = model(x, torch.full((batch_size, 1), t))
x = x + dt * v
Comparaison avec diffusion
| Aspect | Diffusion | Flow Matching |
|---|---|---|
| Processus | Ajouter/retirer du bruit | Chemin direct via vélocité |
| Cible d'entraînement | Prédire le bruit ε |
Prédire la vitesse v |
| Steps d'inférence | 50-1000 | 10-50 |
| Déterminisme | Peut être stochastique | Déterministe (ODE) |
| Loss | MSE + éventuellement KL | MSE simple |
Robotique
Pourquoi Flow Matching pour les actions robot
- Rapidité : 10-20 évaluations d'ODE suffisent pour du temps réel
- Lissité : les trajectoires robotiques sont naturellement continues
- Conditionnement : facile d'ajouter état actuel, but, obstacles
Exemple minimal
# v_θ conditionné sur l'état
class ConditionalFlowModel(nn.Module):
def __init__(self, state_dim, action_dim, hidden=256):
super().__init__()
self.net = nn.Sequential(
nn.Linear(action_dim + 1 + state_dim, hidden),
nn.SiLU(),
nn.Linear(hidden, hidden),
nn.SiLU(),
nn.Linear(hidden, action_dim)
)
def forward(self, x, t, state):
inp = torch.cat([x, t, state], dim=-1)
return self.net(inp)
# Génération d'action
state = get_robot_state()
x0 = torch.randn(1, action_dim)
action = odeint(lambda t, x: model(x, t, state), x0, t_span)[-1]
Flow vs gradient : ne pas confondre
| Concept | Définition | Notation |
|---|---|---|
| Gradient | Direction de plus forte variation d'une fonction scalaire | ∇f(x) |
| Flow (vélocité) | Comment un point se déplace dans l'espace-temps | v(x, t) |
Un flow peut être un gradient (gradient flow : dx/dt = -∇f(x)), mais en Flow Matching le champ de vélocité est appris indépendamment.
Techniques liées
- Continuous Normalizing Flows (CNF) : utilisent des ODE mais calculent des log-vraisemblances coûteuses (trace du Jacobien)
- Rectified Flow : variante qui apprend des trajectoires plus droites via itérations successives
- Score-Based Models : apprennent
∇ log p(x)au lieu dev(x, t), liés parv = -∇ log p
Quand l'utiliser
Flow Matching si :
- ✅ Inférence rapide critique (robotique, interactif)
- ✅ Données continues sur variétés lisses (trajectoires, mouvements)
- ✅ Besoin de déterminisme
Diffusion si :
- Images (écosystème mature, modèles pré-entraînés)
- Variation stochastique souhaitée
- La vitesse d'inférence n'est pas bloquante
title: Flow Matching category: cheatsheets created: 2025-01-15 updated: 2025-01-15 tags: [flow-matching, generative-models, deep-learning, diffusion, ode, cnf, pytorch] confidence: high publish: true
Flow matching transforme une distribution simple (bruit gaussien) en distribution complexe (données) via un champ de vitesse déterministe, offrant la qualité de diffusion avec 7-100× moins d'étapes d'inférence et un entraînement plus simple.
L'idée centrale
On apprend un champ de vitesse v_θ(x, t) qui décrit comment transformer du bruit en données :
# ODE à intégrer de t=0 à t=1
dx_t/dt = v_θ(x_t, t)
À l'entraînement, on régresse directement ce champ de vitesse (MSE simple) au lieu de prédire du bruit comme en diffusion. Mathématiquement équivalent à la diffusion gaussienne, mais plus intuitif.
Conditional Flow Matching (CFM)
Le trick qui rend l'entraînement tractable : conditionner sur un point de données x₁.
import torch
import torch.nn as nn
def cfm_loss(model, x1, sigma_min=0.001):
"""
x1: batch de données réelles, shape (B, D)
"""
B, D = x1.shape
# Échantillonner t uniformément
t = torch.rand(B, 1, device=x1.device)
# Chemin conditionnel gaussien
mu_t = t * x1
sigma_t = 1 - t + t * sigma_min
# Échantillonner x_t sur le chemin
x0 = torch.randn_like(x1)
x_t = mu_t + sigma_t * x0
# Vitesse conditionnelle (forme close)
u_t = (x1 - (1 - sigma_min) * x0) / (1 - (1 - sigma_min) * t)
# Prédire et comparer
v_pred = model(x_t, t)
loss = ((v_pred - u_t) ** 2).mean()
return loss
Pas de calcul de posteriors. Pas d'intégration d'ODE à l'entraînement. Juste de la régression.
Optimal Transport coupling (OT-CFM)
Le gain le plus simple : matcher les paires bruit-données avec Sinkhorn au lieu de les coupler aléatoirement.
from torchdyn.core import NeuralODE
import ot # POT library
def ot_cfm_loss(model, x1, reg=0.05):
B = x1.shape[0]
x0 = torch.randn_like(x1)
# Matrice de coût L2
C = torch.cdist(x0, x1) ** 2
# Plan de transport optimal (Sinkhorn)
pi = ot.sinkhorn(torch.ones(B)/B, torch.ones(B)/B,
C.cpu().numpy(), reg)
pi = torch.from_numpy(pi).to(x1.device)
# Rééchantillonner x0 selon le plan
indices = torch.multinomial(pi, 1).squeeze()
x0_coupled = x0[indices]
# CFM loss standard avec couplage OT
t = torch.rand(B, 1, device=x1.device)
x_t = t * x1 + (1 - t) * x0_coupled
u_t = x1 - x0_coupled
v_pred = model(x_t, t)
return ((v_pred - u_t) ** 2).mean()
Réduit la variance, rend les chemins plus droits. Speedup 4.4× en pratique.
Sampling
Intégrer l'ODE avec n'importe quel solveur.
from torchdiffeq import odeint
def sample(model, batch_size, dim, steps=50, method='dopri5'):
# Partir du bruit
x0 = torch.randn(batch_size, dim)
t_span = torch.linspace(0, 1, steps)
# Définir l'ODE
def ode_func(t, x):
t_batch = t.expand(x.shape[0], 1)
return model(x, t_batch)
# Intégrer
trajectory = odeint(ode_func, x0, t_span, method=method)
return trajectory[-1] # x_1
Solveurs courants :
| Méthode | Steps typiques | Qualité |
|---|---|---|
euler |
50-100 | Acceptable |
rk4 |
20-50 | Bon |
dopri5 (adaptatif) |
10-30 | Excellent |
Équivalence avec diffusion
Flow matching et diffusion gaussienne sont mathématiquement identiques (DeepMind 2024).
| Aspect | Diffusion | Flow Matching |
|---|---|---|
| Cible d'entraînement | Bruit ε | Vitesse v |
| Loss | MSE pondérée (SNR) | MSE simple |
| Trajectoire | Stochastique (SDE) | Déterministe (ODE) |
| Mental model | "Nettoyer le bruit" | "Suivre le flux" |
| Convergence | Baseline | 4.4× plus rapide |
Choisir flow matching pour :
- Formulation plus simple
- Chemins plus droits (moins d'étapes)
- Optimal transport naturel
- Contraintes géométriques (SE(3), SO(3))
Rectified Flow
Rendr les chemins encore plus droits via "reflow" : ré-entraîner le modèle sur ses propres générations.
def reflow(model, dataloader, epochs=10):
"""
Génère (x0, x1) via le modèle actuel,
puis ré-entraîne pour apprendre les chemins droits.
"""
synthetic_pairs = []
with torch.no_grad():
for x1_real in dataloader:
x0 = torch.randn_like(x1_real)
x1_gen = sample(model, x0) # via ODE
synthetic_pairs.append((x0, x1_gen))
# Ré-entraîner avec paires synthétiques
for x0, x1 in synthetic_pairs:
t = torch.rand(B, 1)
x_t = t * x1 + (1 - t) * x0
u_t = x1 - x0 # Vitesse droite
loss = ((model(x_t, t) - u_t) ** 2).mean()
# backward, step...
Après 1-2 reflows : génération en 1-5 étapes seulement.
Production : Stable Diffusion 3
SD3 utilise flow matching + DiT (Diffusion Transformer) :
# Architecture simplifiée
class FlowMatchingDiT(nn.Module):
def __init__(self, dim=1024, depth=28, heads=16):
self.pos_embed = ...
self.blocks = nn.ModuleList([
TransformerBlock(dim, heads) for _ in range(depth)
])
self.final = nn.Linear(dim, dim)
def forward(self, x, t, text_embed):
# x: latents, t: timestep, text_embed: CLIP/T5
h = self.pos_embed(x) + self.time_embed(t)
for block in self.blocks:
h = block(h, text_embed) # cross-attention
return self.final(h) # prédire vitesse
Clés du succès :
- Latent space (VAE 8× compression)
- Classifier-free guidance (CFG)
- OT coupling
- 50 steps d'inférence (vs 1000 pour DDPM initial)
TorchCFM
Librairie officielle de référence.
pip install torchdyn torchcfm
from torchcfm.conditional_flow_matching import (
ConditionalFlowMatcher,
TargetConditionalFlowMatcher, # OT version
)
# Setup
fm = TargetConditionalFlowMatcher(sigma=0.0) # sigma=0: chemins droits
# Training loop
for x1 in dataloader:
x0 = torch.randn_like(x1)
t, xt, ut = fm.sample_location_and_conditional_flow(x0, x1)
vt = model(xt, t)
loss = (vt - ut).pow(2).mean()
optimizer.zero_grad()
loss.backward()
optimizer.step()
Cas d'usage
Où flow matching excelle :
- Images : SD3, Flux.1 (12B params)
- Vidéo : Pyramidal Flow (10s, 768p, 24fps)
- Robotique : policies 50Hz, contraintes géométriques
- Protéines : génération SE(3)-équivariante
- Molécules : conformères 3D avec symétrie E(3)
- Parole : synthèse haute fidélité
Limites actuelles :
- Données discrètes (texte) : encore derrière l'autoregressif
- Écosystème moins mature que diffusion
- One-step generation : léger retard sur diffusion distillée
Diagnostic
# Vérifier la qualité du flow
def plot_trajectories(model, x1, steps=50):
x0 = torch.randn_like(x1)
t_span = torch.linspace(0, 1, steps)
trajectory = []
x = x0
for i in range(len(t_span) - 1):
dt = t_span[i+1] - t_span[i]
v = model(x, t_span[i].expand(x.shape[0], 1))
x = x + v * dt
trajectory.append(x.clone())
# Mesurer la courbure
curvature = sum([
torch.norm(trajectory[i+1] - 2*trajectory[i] + trajectory[i-1])
for i in range(1, len(trajectory)-1)
])
print(f"Courbure totale: {curvature:.4f}") # Plus bas = mieux
Chemins droits = moins d'étapes nécessaires.
See also
pytorch keras-tensorflow world-models