~/wiki

Flow Matching

Mis à jour le 2026-08-07Confiance : medium
flow-matchinggenerative-modelsoderoboticsdiffusiondeep-learning

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 de v(x, t), liés par v = -∇ 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

See also

keras-tensorflow pytorch world-models