Skip to content

Flow Matching Objective

Regresses the model's velocity against the ground-truth conditional velocity along an interpolation path.

flow_matching

Flow matching objective for neural ODE flows.

This module provides a single, configurable objective that encompasses: - Standard Flow Matching (Lipman et al., 2023) - Rectified Flow (Liu et al., 2023) - OT-Conditional Flow Matching (Tong et al., 2023)

All variants are achieved through configuration, not inheritance.

FlowMatchingObjective

Bases: Module

Flow-matching objective for NeuralODEFlow.

This objective trains a neural ODE flow by supervised velocity regression along interpolation paths instead of maximum likelihood. Multiple flow matching variants are configured through modular components:

  • Path: How to interpolate between noise and data
  • Time Sampler: Distribution for sampling training times
  • Coupling: How to pair noise and data samples
  • Source: Where trajectories start (noise, or precomputed pairs)

Standard Configurations:

  1. Flow Matching (default):

    FlowMatchingObjective(flow, path='linear', time_sampler='uniform')
    

  2. Rectified Flow:

    FlowMatchingObjective(flow, path='linear', time_sampler='logit_normal')
    

  3. OT-Conditional Flow Matching:

    FlowMatchingObjective(flow, path='ot_conditional', sigma=0.01)
    

Parameters:

Name Type Description Default
flow NeuralODEFlow

Neural ODE flow whose model predicts velocity v(x_t, condition, t)

required
path Union[str, PathInterpolant]

Interpolation path ('linear', 'ot_conditional') or PathInterpolant

'linear'
time_sampler Union[str, TimeSampler]

Time distribution ('uniform', 'logit_normal') or TimeSampler

'uniform'
coupling Union[str, Coupling]

Coupling strategy ('independent', 'minibatch_ot') or Coupling

'independent'
source Union[str, SourceDistribution]

Source distribution ('gaussian', 'batch') or SourceDistribution. Use 'batch' (BatchSource) to train on precomputed (x_0, x_1) pairs, which is what reflow requires.

'gaussian'
sigma float

Noise level for OT-conditional path (default: 0.0)

0.0
target_key Optional[str]

Default batch key for target tensors (default: 'u')

None
condition_key Optional[str]

Default batch key for condition tensors (default: 'f')

None
References
  • Lipman et al., "Flow Matching for Generative Modeling", ICLR 2023
  • Liu et al., "Flow Straight and Fast: Rectified Flow", ICLR 2023
  • Tong et al., "Conditional Flow Matching", NeurIPS 2023
Source code in flowpde/objectives/flow_matching.py
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
class FlowMatchingObjective(nn.Module):
    """
    Flow-matching objective for `NeuralODEFlow`.

    This objective trains a neural ODE flow by supervised velocity regression
    along interpolation paths instead of maximum likelihood. Multiple flow
    matching variants are configured through modular components:

    - **Path**: How to interpolate between noise and data
    - **Time Sampler**: Distribution for sampling training times
    - **Coupling**: How to pair noise and data samples
    - **Source**: Where trajectories start (noise, or precomputed pairs)

    Standard Configurations:

    1. **Flow Matching** (default):
       ```python
       FlowMatchingObjective(flow, path='linear', time_sampler='uniform')
       ```

    2. **Rectified Flow**:
       ```python
       FlowMatchingObjective(flow, path='linear', time_sampler='logit_normal')
       ```

    3. **OT-Conditional Flow Matching**:
       ```python
       FlowMatchingObjective(flow, path='ot_conditional', sigma=0.01)
       ```

    Args:
        flow: Neural ODE flow whose model predicts velocity
            v(x_t, condition, t)
        path: Interpolation path ('linear', 'ot_conditional') or PathInterpolant
        time_sampler: Time distribution ('uniform', 'logit_normal') or TimeSampler
        coupling: Coupling strategy ('independent', 'minibatch_ot') or Coupling
        source: Source distribution ('gaussian', 'batch') or SourceDistribution.
            Use 'batch' (`BatchSource`) to train on precomputed (x_0, x_1)
            pairs, which is what reflow requires.
        sigma: Noise level for OT-conditional path (default: 0.0)
        target_key: Default batch key for target tensors (default: 'u')
        condition_key: Default batch key for condition tensors (default: 'f')

    References:
        - Lipman et al., "Flow Matching for Generative Modeling", ICLR 2023
        - Liu et al., "Flow Straight and Fast: Rectified Flow", ICLR 2023
        - Tong et al., "Conditional Flow Matching", NeurIPS 2023
    """

    def __init__(
        self,
        flow: NeuralODEFlow,
        path: Union[str, PathInterpolant] = "linear",
        time_sampler: Union[str, TimeSampler] = "uniform",
        coupling: Union[str, Coupling] = "independent",
        source: Union[str, SourceDistribution] = "gaussian",
        sigma: float = 0.0,
        target_key: Optional[str] = None,
        condition_key: Optional[str] = None,
    ):
        super().__init__()
        self.flow = flow
        self.model = flow.model
        self.target_key = target_key or flow.target_key
        self.condition_key = condition_key or flow.condition_key

        # Initialize components
        # Pass sigma to OT path if needed
        if isinstance(path, str) and path in ['ot_conditional', 'conditional_optimal_transport', 'ot']:
            self.path = get_path(path, sigma=sigma)
        else:
            self.path = get_path(path)

        self.time_sampler = get_time_sampler(time_sampler)
        self.coupling = get_coupling(coupling)
        self.source = get_source(source)
        self.sigma = sigma

        # Store string names for config
        self._path_name = path if isinstance(path, str) else path.__class__.__name__
        self._time_sampler_name = time_sampler if isinstance(time_sampler, str) else time_sampler.__class__.__name__
        self._coupling_name = coupling if isinstance(coupling, str) else coupling.__class__.__name__
        self._source_name = source if isinstance(source, str) else source.__class__.__name__

    def sample_base_distribution(
        self,
        shape: Tuple[int, ...],
        device: torch.device,
        batch: Optional[Dict[str, Tensor]] = None,
    ) -> Tensor:
        """
        Draw `x_0` from the configured source distribution.

        Args:
            shape: Shape of `x_0`, matching the target.
            device: Device to place the result on.
            batch: Training batch, forwarded so sources such as
                `BatchSource` can read precomputed values from it.
                Omitted at inference, where sources fall back to noise.
        """
        return self.source(shape, device, batch)

    @property
    def model_device(self) -> torch.device:
        return self.flow.model_device

    def _source_defines_pairing(self, batch: Optional[Dict[str, Tensor]]) -> bool:
        """Whether the source already determines which x_0 goes with which x_1.

        Reflow pairs are meaningful only as pairs: the model is retrained on
        exactly the trajectories it generated. Reordering them with a coupling
        would break that correspondence.
        """
        source = self.source
        key = getattr(source, "key", None)
        return key is not None and batch is not None and key in batch

    def compute_loss(
        self,
        batch: Dict[str, Tensor],
        target_key: Optional[str] = None,
        condition_key: Optional[str] = None,
        **kwargs: Any
    ) -> Tensor:
        """
        Compute flow matching loss.

        The loss minimizes the MSE between predicted and target velocities:

        $$\\mathcal{L} = \\mathbb{E}_{t, x_0, x_1}[\\|v_\\theta(x_t, f, t) - v_t\\|^2]$$

        where:
        - $x_t$ is the interpolated point on the path
        - $v_t$ is the target velocity (derivative of path)
        - $f$ is the conditioning information

        Args:
            batch: Dictionary containing target and condition tensors.
            target_key: Batch key for target data. Defaults to this flow's
                configured target key ('u' by default).
            condition_key: Batch key for conditioning data. Defaults to this
                flow's configured condition key ('f' by default).

        Returns:
            MSE loss tensor (scalar)
        """
        # Extract and prepare data
        x_1, condition = self.flow._extract_target_condition(
            batch,
            target_key=target_key or self.target_key,
            condition_key=condition_key or self.condition_key,
        )
        self.flow.set_target_dim(x_1.shape[1])
        batch_size = x_1.shape[0]

        # Draw x_0 from the source. Passing the batch lets BatchSource
        # return the precomputed x_0 that reflow depends on.
        x_0 = self.sample_base_distribution(x_1.shape, self.model_device, batch)

        # Apply coupling strategy. A source that carries its own pairing
        # already fixes which x_0 goes with which x_1, so re-coupling here
        # would destroy it.
        if not self._source_defines_pairing(batch):
            x_0, x_1 = self.coupling(x_0, x_1)

        # Sample time
        t = self.time_sampler(batch_size, self.model_device)

        # Compute path interpolation and target velocity
        x_t, v_target = self.path(x_0, x_1, t)

        # Predict velocity with model
        v_pred = self.model(x_t, condition, t)

        # MSE loss
        loss = F.mse_loss(v_pred, v_target)

        return loss

    def sample(
        self,
        condition: Tensor,
        n_steps: Optional[int] = None,
        solver: Optional[str] = None,
        x_init: Optional[Tensor] = None,
        target_shape: Optional[Union[int, Tuple[int, ...]]] = None,
        return_trajectory: bool = False,
        no_grad: bool = True,
        **solver_kwargs: Any
    ) -> Union[Tensor, Tuple[Tensor, Tensor]]:
        """
        Generate samples by solving the flow ODE.

        Integrates the learned velocity field from t=0 (noise) to t=1 (data):

        $$\\frac{dx}{dt} = v_\\theta(x_t, f, t), \\quad x_0 \\sim \\mathcal{N}(0, I)$$

        The integration itself is `NeuralODEFlow.sample`; what this
        adds is the starting point, drawn from the objective's configured
        `source` so inference matches how the model was trained.

        Args:
            condition: Conditioning tensor (B, *)
            n_steps: Number of integration steps.  Defaults to the flow's
                `ode_n_steps`; required for fixed-step solvers.
            solver: ODE solver ('euler', 'midpoint', 'rk4', 'dopri5').
                Defaults to the flow's `ode_method`.
            x_init: Optional initial noise (default: drawn from `source`)
            target_shape: Shape of generated targets excluding batch, or a
                flattened target dimension.  Only needed when the flow has
                not recorded its target dimension and no `x_init` is given.
            return_trajectory: If True, return full trajectory
            no_grad: Integrate under `torch.no_grad()` (default).  Pass False
                to differentiate through sampling.
            **solver_kwargs: Additional solver arguments (`rtol`, `atol`,
                `adjoint`, `method_options`).  Unknown names raise.

        Returns:
            Generated samples (B, dim)
            If return_trajectory: (samples, trajectory) where trajectory is (n_steps+1, B, dim)
        """
        if x_init is None:
            condition_flat = condition.flatten(start_dim=1).to(self.model_device)
            dim = self.flow._resolve_target_dim(target_shape)
            x_init = self.sample_base_distribution(
                (condition_flat.shape[0], dim), self.model_device
            )

        return self.flow.sample(
            condition=condition,
            n_steps=n_steps,
            solver=solver,
            x_init=x_init,
            return_trajectory=return_trajectory,
            no_grad=no_grad,
            **solver_kwargs,
        )

    def estimate_straightness(
        self,
        batch: Dict[str, Tensor],
        n_time_points: int = 10,
        mode: str = "trajectory",
        n_steps: int = 50,
        solver: str = "euler",
        target_key: Optional[str] = None,
        condition_key: Optional[str] = None,
    ) -> Dict[str, float]:
        """
        Measure how straight the learned transport paths are.

        Straightness follows Liu et al. (2023): a flow is straight when its
        velocity along a trajectory equals the chord connecting the endpoints,

        $$S = \\int_0^1 \\mathbb{E}\\left[\\lVert (Z_1 - Z_0)
              - v_\\theta(Z_t, t) \\rVert^2\\right] dt,$$

        so $S = 0$ exactly when every trajectory is a straight line
        traversed at constant velocity — which is what makes few-step Euler
        sampling accurate, and what reflow is meant to improve.

        Two modes are available:

        - `'trajectory'` (default): integrate the learned ODE and measure
          deviation along the model's **own** trajectories.  This is the
          quantity that predicts few-step sampling quality.
        - `'interpolant'`: measure deviation along the training interpolant
          between sampled `(x_0, x_1)` pairs.  Cheaper (no ODE solve) and it
          reports how far the learned marginal velocity sits from the
          conditional target, but it does *not* describe the sampling paths.

        Args:
            batch: Batch with target and condition tensors.
            n_time_points: Number of time points at which velocity is probed.
            mode: `'trajectory'` or `'interpolant'`.
            n_steps: ODE steps used to build trajectories (`'trajectory'`).
            solver: ODE solver used to build trajectories (`'trajectory'`).
            target_key: Batch key for target data.
            condition_key: Batch key for conditioning data.

        Returns:
            Dictionary with:

            - `'straightness'`: the integral above (0 = perfectly straight).
            - `'normalized_straightness'`: divided by the mean squared chord
                length, making it dimensionless and comparable across datasets
                and normalization choices.
            - `'chord_norm'`: mean chord length, for reference.
        """
        if mode not in {"trajectory", "interpolant"}:
            raise ValueError(
                f"mode must be 'trajectory' or 'interpolant', got '{mode}'"
            )

        was_training = self.training
        self.eval()

        x_1, condition = self.flow._extract_target_condition(
            batch,
            target_key=target_key or self.target_key,
            condition_key=condition_key or self.condition_key,
        )
        batch_size = x_1.shape[0]

        try:
            with torch.no_grad():
                if mode == "trajectory":
                    x_0 = self.sample_base_distribution(x_1.shape, self.model_device)
                    _, trajectory = self.sample(
                        condition=condition,
                        n_steps=n_steps,
                        solver=solver,
                        x_init=x_0,
                        return_trajectory=True,
                    )
                    # trajectory: (n_steps + 1, B, dim)
                    z_0, z_1 = trajectory[0], trajectory[-1]
                    chord = z_1 - z_0

                    # Probe the whole trajectory, endpoints included: the
                    # straightness integral in Liu et al. runs over the
                    # closed interval. (The 'interpolant' mode below uses
                    # [0.01, 0.99] instead, so the two modes are not
                    # numerically comparable -- they are different
                    # estimators, as the docstring says.)
                    indices = torch.linspace(
                        0, trajectory.shape[0] - 1, n_time_points
                    ).round().long()
                    time_grid = torch.linspace(
                        0.0, 1.0, trajectory.shape[0], device=self.model_device
                    )

                    deviations = []
                    for index in indices:
                        z_t = trajectory[index]
                        t = time_grid[index].expand(batch_size, 1)
                        v_pred = self.model(z_t, condition, t)
                        deviations.append((v_pred - chord).pow(2).sum(dim=1))
                else:
                    x_0 = self.sample_base_distribution(x_1.shape, self.model_device)
                    x_0, x_1 = self.coupling(x_0, x_1)
                    chord = x_1 - x_0

                    time_points = torch.linspace(
                        0.01, 0.99, n_time_points, device=self.model_device
                    )
                    deviations = []
                    for t_val in time_points:
                        t = t_val.expand(batch_size, 1)
                        x_t, _ = self.path(x_0, x_1, t)
                        v_pred = self.model(x_t, condition, t)
                        deviations.append((v_pred - chord).pow(2).sum(dim=1))

                # Mean over time (the integral) and over the batch.
                straightness = torch.stack(deviations, dim=0).mean().item()
                chord_sq = chord.pow(2).sum(dim=1).mean()
                normalized = (straightness / chord_sq.clamp(min=1e-12)).item()
                chord_norm = chord.norm(dim=1).mean().item()
        finally:
            if was_training:
                self.train()

        return {
            "straightness": straightness,
            "normalized_straightness": normalized,
            "chord_norm": chord_norm,
        }

    def get_config(self) -> Dict[str, Any]:
        """Return configuration dictionary for serialization."""
        return {
            'objective': 'flow_matching',
            'flow': self.flow.get_config(),
            'path': self._path_name,
            'time_sampler': self._time_sampler_name,
            'coupling': self._coupling_name,
            'source': self.source.get_config(),
            'sigma': self.sigma,
            'target_key': self.target_key,
            'condition_key': self.condition_key,
            'model_type': self.model.__class__.__name__,
        }

    def __repr__(self) -> str:
        return (
            f"FlowMatchingObjective(\n"
            f"  path={self._path_name},\n"
            f"  time_sampler={self._time_sampler_name},\n"
            f"  coupling={self._coupling_name},\n"
            f"  source={self.source!r},\n"
            f"  sigma={self.sigma}\n"
            f")"
        )
sample_base_distribution(shape, device, batch=None)

Draw x_0 from the configured source distribution.

Parameters:

Name Type Description Default
shape Tuple[int, ...]

Shape of x_0, matching the target.

required
device device

Device to place the result on.

required
batch Optional[Dict[str, Tensor]]

Training batch, forwarded so sources such as BatchSource can read precomputed values from it. Omitted at inference, where sources fall back to noise.

None
Source code in flowpde/objectives/flow_matching.py
def sample_base_distribution(
    self,
    shape: Tuple[int, ...],
    device: torch.device,
    batch: Optional[Dict[str, Tensor]] = None,
) -> Tensor:
    """
    Draw `x_0` from the configured source distribution.

    Args:
        shape: Shape of `x_0`, matching the target.
        device: Device to place the result on.
        batch: Training batch, forwarded so sources such as
            `BatchSource` can read precomputed values from it.
            Omitted at inference, where sources fall back to noise.
    """
    return self.source(shape, device, batch)
compute_loss(batch, target_key=None, condition_key=None, **kwargs)

Compute flow matching loss.

The loss minimizes the MSE between predicted and target velocities:

\[\mathcal{L} = \mathbb{E}_{t, x_0, x_1}[\|v_\theta(x_t, f, t) - v_t\|^2]\]

where: - \(x_t\) is the interpolated point on the path - \(v_t\) is the target velocity (derivative of path) - \(f\) is the conditioning information

Parameters:

Name Type Description Default
batch Dict[str, Tensor]

Dictionary containing target and condition tensors.

required
target_key Optional[str]

Batch key for target data. Defaults to this flow's configured target key ('u' by default).

None
condition_key Optional[str]

Batch key for conditioning data. Defaults to this flow's configured condition key ('f' by default).

None

Returns:

Type Description
Tensor

MSE loss tensor (scalar)

Source code in flowpde/objectives/flow_matching.py
def compute_loss(
    self,
    batch: Dict[str, Tensor],
    target_key: Optional[str] = None,
    condition_key: Optional[str] = None,
    **kwargs: Any
) -> Tensor:
    """
    Compute flow matching loss.

    The loss minimizes the MSE between predicted and target velocities:

    $$\\mathcal{L} = \\mathbb{E}_{t, x_0, x_1}[\\|v_\\theta(x_t, f, t) - v_t\\|^2]$$

    where:
    - $x_t$ is the interpolated point on the path
    - $v_t$ is the target velocity (derivative of path)
    - $f$ is the conditioning information

    Args:
        batch: Dictionary containing target and condition tensors.
        target_key: Batch key for target data. Defaults to this flow's
            configured target key ('u' by default).
        condition_key: Batch key for conditioning data. Defaults to this
            flow's configured condition key ('f' by default).

    Returns:
        MSE loss tensor (scalar)
    """
    # Extract and prepare data
    x_1, condition = self.flow._extract_target_condition(
        batch,
        target_key=target_key or self.target_key,
        condition_key=condition_key or self.condition_key,
    )
    self.flow.set_target_dim(x_1.shape[1])
    batch_size = x_1.shape[0]

    # Draw x_0 from the source. Passing the batch lets BatchSource
    # return the precomputed x_0 that reflow depends on.
    x_0 = self.sample_base_distribution(x_1.shape, self.model_device, batch)

    # Apply coupling strategy. A source that carries its own pairing
    # already fixes which x_0 goes with which x_1, so re-coupling here
    # would destroy it.
    if not self._source_defines_pairing(batch):
        x_0, x_1 = self.coupling(x_0, x_1)

    # Sample time
    t = self.time_sampler(batch_size, self.model_device)

    # Compute path interpolation and target velocity
    x_t, v_target = self.path(x_0, x_1, t)

    # Predict velocity with model
    v_pred = self.model(x_t, condition, t)

    # MSE loss
    loss = F.mse_loss(v_pred, v_target)

    return loss
sample(condition, n_steps=None, solver=None, x_init=None, target_shape=None, return_trajectory=False, no_grad=True, **solver_kwargs)

Generate samples by solving the flow ODE.

Integrates the learned velocity field from t=0 (noise) to t=1 (data):

\[\frac{dx}{dt} = v_\theta(x_t, f, t), \quad x_0 \sim \mathcal{N}(0, I)\]

The integration itself is NeuralODEFlow.sample; what this adds is the starting point, drawn from the objective's configured source so inference matches how the model was trained.

Parameters:

Name Type Description Default
condition Tensor

Conditioning tensor (B, *)

required
n_steps Optional[int]

Number of integration steps. Defaults to the flow's ode_n_steps; required for fixed-step solvers.

None
solver Optional[str]

ODE solver ('euler', 'midpoint', 'rk4', 'dopri5'). Defaults to the flow's ode_method.

None
x_init Optional[Tensor]

Optional initial noise (default: drawn from source)

None
target_shape Optional[Union[int, Tuple[int, ...]]]

Shape of generated targets excluding batch, or a flattened target dimension. Only needed when the flow has not recorded its target dimension and no x_init is given.

None
return_trajectory bool

If True, return full trajectory

False
no_grad bool

Integrate under torch.no_grad() (default). Pass False to differentiate through sampling.

True
**solver_kwargs Any

Additional solver arguments (rtol, atol, adjoint, method_options). Unknown names raise.

{}

Returns:

Type Description
Union[Tensor, Tuple[Tensor, Tensor]]

Generated samples (B, dim)

Union[Tensor, Tuple[Tensor, Tensor]]

If return_trajectory: (samples, trajectory) where trajectory is (n_steps+1, B, dim)

Source code in flowpde/objectives/flow_matching.py
def sample(
    self,
    condition: Tensor,
    n_steps: Optional[int] = None,
    solver: Optional[str] = None,
    x_init: Optional[Tensor] = None,
    target_shape: Optional[Union[int, Tuple[int, ...]]] = None,
    return_trajectory: bool = False,
    no_grad: bool = True,
    **solver_kwargs: Any
) -> Union[Tensor, Tuple[Tensor, Tensor]]:
    """
    Generate samples by solving the flow ODE.

    Integrates the learned velocity field from t=0 (noise) to t=1 (data):

    $$\\frac{dx}{dt} = v_\\theta(x_t, f, t), \\quad x_0 \\sim \\mathcal{N}(0, I)$$

    The integration itself is `NeuralODEFlow.sample`; what this
    adds is the starting point, drawn from the objective's configured
    `source` so inference matches how the model was trained.

    Args:
        condition: Conditioning tensor (B, *)
        n_steps: Number of integration steps.  Defaults to the flow's
            `ode_n_steps`; required for fixed-step solvers.
        solver: ODE solver ('euler', 'midpoint', 'rk4', 'dopri5').
            Defaults to the flow's `ode_method`.
        x_init: Optional initial noise (default: drawn from `source`)
        target_shape: Shape of generated targets excluding batch, or a
            flattened target dimension.  Only needed when the flow has
            not recorded its target dimension and no `x_init` is given.
        return_trajectory: If True, return full trajectory
        no_grad: Integrate under `torch.no_grad()` (default).  Pass False
            to differentiate through sampling.
        **solver_kwargs: Additional solver arguments (`rtol`, `atol`,
            `adjoint`, `method_options`).  Unknown names raise.

    Returns:
        Generated samples (B, dim)
        If return_trajectory: (samples, trajectory) where trajectory is (n_steps+1, B, dim)
    """
    if x_init is None:
        condition_flat = condition.flatten(start_dim=1).to(self.model_device)
        dim = self.flow._resolve_target_dim(target_shape)
        x_init = self.sample_base_distribution(
            (condition_flat.shape[0], dim), self.model_device
        )

    return self.flow.sample(
        condition=condition,
        n_steps=n_steps,
        solver=solver,
        x_init=x_init,
        return_trajectory=return_trajectory,
        no_grad=no_grad,
        **solver_kwargs,
    )
estimate_straightness(batch, n_time_points=10, mode='trajectory', n_steps=50, solver='euler', target_key=None, condition_key=None)

Measure how straight the learned transport paths are.

Straightness follows Liu et al. (2023): a flow is straight when its velocity along a trajectory equals the chord connecting the endpoints,

\[S = \int_0^1 \mathbb{E}\left[\lVert (Z_1 - Z_0) - v_\theta(Z_t, t) \rVert^2\right] dt,\]

so \(S = 0\) exactly when every trajectory is a straight line traversed at constant velocity — which is what makes few-step Euler sampling accurate, and what reflow is meant to improve.

Two modes are available:

  • 'trajectory' (default): integrate the learned ODE and measure deviation along the model's own trajectories. This is the quantity that predicts few-step sampling quality.
  • 'interpolant': measure deviation along the training interpolant between sampled (x_0, x_1) pairs. Cheaper (no ODE solve) and it reports how far the learned marginal velocity sits from the conditional target, but it does not describe the sampling paths.

Parameters:

Name Type Description Default
batch Dict[str, Tensor]

Batch with target and condition tensors.

required
n_time_points int

Number of time points at which velocity is probed.

10
mode str

'trajectory' or 'interpolant'.

'trajectory'
n_steps int

ODE steps used to build trajectories ('trajectory').

50
solver str

ODE solver used to build trajectories ('trajectory').

'euler'
target_key Optional[str]

Batch key for target data.

None
condition_key Optional[str]

Batch key for conditioning data.

None

Returns:

Type Description
Dict[str, float]

Dictionary with:

Dict[str, float]
  • 'straightness': the integral above (0 = perfectly straight).
Dict[str, float]
  • 'normalized_straightness': divided by the mean squared chord length, making it dimensionless and comparable across datasets and normalization choices.
Dict[str, float]
  • 'chord_norm': mean chord length, for reference.
Source code in flowpde/objectives/flow_matching.py
def estimate_straightness(
    self,
    batch: Dict[str, Tensor],
    n_time_points: int = 10,
    mode: str = "trajectory",
    n_steps: int = 50,
    solver: str = "euler",
    target_key: Optional[str] = None,
    condition_key: Optional[str] = None,
) -> Dict[str, float]:
    """
    Measure how straight the learned transport paths are.

    Straightness follows Liu et al. (2023): a flow is straight when its
    velocity along a trajectory equals the chord connecting the endpoints,

    $$S = \\int_0^1 \\mathbb{E}\\left[\\lVert (Z_1 - Z_0)
          - v_\\theta(Z_t, t) \\rVert^2\\right] dt,$$

    so $S = 0$ exactly when every trajectory is a straight line
    traversed at constant velocity — which is what makes few-step Euler
    sampling accurate, and what reflow is meant to improve.

    Two modes are available:

    - `'trajectory'` (default): integrate the learned ODE and measure
      deviation along the model's **own** trajectories.  This is the
      quantity that predicts few-step sampling quality.
    - `'interpolant'`: measure deviation along the training interpolant
      between sampled `(x_0, x_1)` pairs.  Cheaper (no ODE solve) and it
      reports how far the learned marginal velocity sits from the
      conditional target, but it does *not* describe the sampling paths.

    Args:
        batch: Batch with target and condition tensors.
        n_time_points: Number of time points at which velocity is probed.
        mode: `'trajectory'` or `'interpolant'`.
        n_steps: ODE steps used to build trajectories (`'trajectory'`).
        solver: ODE solver used to build trajectories (`'trajectory'`).
        target_key: Batch key for target data.
        condition_key: Batch key for conditioning data.

    Returns:
        Dictionary with:

        - `'straightness'`: the integral above (0 = perfectly straight).
        - `'normalized_straightness'`: divided by the mean squared chord
            length, making it dimensionless and comparable across datasets
            and normalization choices.
        - `'chord_norm'`: mean chord length, for reference.
    """
    if mode not in {"trajectory", "interpolant"}:
        raise ValueError(
            f"mode must be 'trajectory' or 'interpolant', got '{mode}'"
        )

    was_training = self.training
    self.eval()

    x_1, condition = self.flow._extract_target_condition(
        batch,
        target_key=target_key or self.target_key,
        condition_key=condition_key or self.condition_key,
    )
    batch_size = x_1.shape[0]

    try:
        with torch.no_grad():
            if mode == "trajectory":
                x_0 = self.sample_base_distribution(x_1.shape, self.model_device)
                _, trajectory = self.sample(
                    condition=condition,
                    n_steps=n_steps,
                    solver=solver,
                    x_init=x_0,
                    return_trajectory=True,
                )
                # trajectory: (n_steps + 1, B, dim)
                z_0, z_1 = trajectory[0], trajectory[-1]
                chord = z_1 - z_0

                # Probe the whole trajectory, endpoints included: the
                # straightness integral in Liu et al. runs over the
                # closed interval. (The 'interpolant' mode below uses
                # [0.01, 0.99] instead, so the two modes are not
                # numerically comparable -- they are different
                # estimators, as the docstring says.)
                indices = torch.linspace(
                    0, trajectory.shape[0] - 1, n_time_points
                ).round().long()
                time_grid = torch.linspace(
                    0.0, 1.0, trajectory.shape[0], device=self.model_device
                )

                deviations = []
                for index in indices:
                    z_t = trajectory[index]
                    t = time_grid[index].expand(batch_size, 1)
                    v_pred = self.model(z_t, condition, t)
                    deviations.append((v_pred - chord).pow(2).sum(dim=1))
            else:
                x_0 = self.sample_base_distribution(x_1.shape, self.model_device)
                x_0, x_1 = self.coupling(x_0, x_1)
                chord = x_1 - x_0

                time_points = torch.linspace(
                    0.01, 0.99, n_time_points, device=self.model_device
                )
                deviations = []
                for t_val in time_points:
                    t = t_val.expand(batch_size, 1)
                    x_t, _ = self.path(x_0, x_1, t)
                    v_pred = self.model(x_t, condition, t)
                    deviations.append((v_pred - chord).pow(2).sum(dim=1))

            # Mean over time (the integral) and over the batch.
            straightness = torch.stack(deviations, dim=0).mean().item()
            chord_sq = chord.pow(2).sum(dim=1).mean()
            normalized = (straightness / chord_sq.clamp(min=1e-12)).item()
            chord_norm = chord.norm(dim=1).mean().item()
    finally:
        if was_training:
            self.train()

    return {
        "straightness": straightness,
        "normalized_straightness": normalized,
        "chord_norm": chord_norm,
    }
get_config()

Return configuration dictionary for serialization.

Source code in flowpde/objectives/flow_matching.py
def get_config(self) -> Dict[str, Any]:
    """Return configuration dictionary for serialization."""
    return {
        'objective': 'flow_matching',
        'flow': self.flow.get_config(),
        'path': self._path_name,
        'time_sampler': self._time_sampler_name,
        'coupling': self._coupling_name,
        'source': self.source.get_config(),
        'sigma': self.sigma,
        'target_key': self.target_key,
        'condition_key': self.condition_key,
        'model_type': self.model.__class__.__name__,
    }

create_flow_matching(flow, variant='standard', **kwargs)

Create a flow matching objective with preset configurations.

Parameters:

Name Type Description Default
flow NeuralODEFlow

Neural ODE flow

required
variant str

Preset name: - 'standard': Standard flow matching (linear, uniform) - 'rectified': Rectified flow (linear, logit-normal) - 'ot_cfm': OT-Conditional FM (ot_conditional, uniform)

'standard'
**kwargs Any

Override any default parameters

{}

Returns:

Type Description
FlowMatchingObjective

Configured flow matching objective

Source code in flowpde/objectives/flow_matching.py
def create_flow_matching(
    flow: NeuralODEFlow,
    variant: str = "standard",
    **kwargs: Any
) -> FlowMatchingObjective:
    """
    Create a flow matching objective with preset configurations.

    Args:
        flow: Neural ODE flow
        variant: Preset name:
            - 'standard': Standard flow matching (linear, uniform)
            - 'rectified': Rectified flow (linear, logit-normal)
            - 'ot_cfm': OT-Conditional FM (ot_conditional, uniform)
        **kwargs: Override any default parameters

    Returns:
        Configured flow matching objective
    """
    presets = {
        'standard': {
            'path': 'linear',
            'time_sampler': 'uniform',
            'coupling': 'independent',
            'sigma': 0.0,
        },
        'rectified': {
            'path': 'linear',
            'time_sampler': 'logit_normal',
            'coupling': 'independent',
            'sigma': 0.0,
        },
        'ot_cfm': {
            'path': 'ot_conditional',
            'time_sampler': 'uniform',
            'coupling': 'independent',
            'sigma': 0.01,
        },
        'ot_cfm_coupled': {
            'path': 'ot_conditional',
            'time_sampler': 'uniform',
            'coupling': 'minibatch_ot',
            'sigma': 0.01,
        },
    }

    if variant not in presets:
        raise ValueError(f"Unknown variant: '{variant}'. Available: {list(presets.keys())}")

    config = presets[variant].copy()
    config.update(kwargs)

    return FlowMatchingObjective(flow, **config)