Skip to content

Flow Components

Modular components that FlowMatchingObjective composes, so flow-matching variants are configuration rather than subclasses.

Every getter (get_path, get_time_sampler, get_coupling, get_source) accepts either a registry name or an already-constructed instance.

Paths

The interpolation between noise and data.

Invariant

A path's velocity() must be the exact time derivative of its interpolate(). If they disagree, training silently regresses a target that does not match the interpolation the model is shown.

paths

Path Interpolation Strategies for Flow Matching.

Defines how to interpolate between noise (x_0) and data (x_1) and compute the corresponding target velocity fields.

PathInterpolant

Bases: ABC

Abstract base class for interpolation paths in flow matching.

A path defines: 1. How to interpolate: x_t = interpolate(x_0, x_1, t) 2. The target velocity: v_t = velocity(x_0, x_1, t)

The velocity is the derivative dx_t/dt that the model learns to predict.

Source code in flowpde/flows/components/paths.py
class PathInterpolant(ABC):
    """
    Abstract base class for interpolation paths in flow matching.

    A path defines:
    1. How to interpolate: x_t = interpolate(x_0, x_1, t)
    2. The target velocity: v_t = velocity(x_0, x_1, t)

    The velocity is the derivative dx_t/dt that the model learns to predict.
    """

    @abstractmethod
    def interpolate(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
        """
        Compute interpolated point x_t on the path.

        Args:
            x_0: Starting point (noise), shape (B, *)
            x_1: Ending point (data), shape (B, *)
            t: Time in [0, 1], shape (B, 1)

        Returns:
            x_t: Interpolated point, shape (B, *)
        """
        pass

    @abstractmethod
    def velocity(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
        """
        Compute target velocity v_t = dx_t/dt.

        Args:
            x_0: Starting point (noise), shape (B, *)
            x_1: Ending point (data), shape (B, *)
            t: Time in [0, 1], shape (B, 1)

        Returns:
            v_t: Target velocity, shape (B, *)
        """
        pass

    def __call__(
        self,
        x_0: Tensor,
        x_1: Tensor,
        t: Tensor
    ) -> Tuple[Tensor, Tensor]:
        """
        Compute both interpolation and velocity.

        Returns:
            (x_t, v_t): Interpolated point and target velocity
        """
        x_t = self.interpolate(x_0, x_1, t)
        v_t = self.velocity(x_0, x_1, t)
        return x_t, v_t
interpolate(x_0, x_1, t) abstractmethod

Compute interpolated point x_t on the path.

Parameters:

Name Type Description Default
x_0 Tensor

Starting point (noise), shape (B, *)

required
x_1 Tensor

Ending point (data), shape (B, *)

required
t Tensor

Time in [0, 1], shape (B, 1)

required

Returns:

Name Type Description
x_t Tensor

Interpolated point, shape (B, *)

Source code in flowpde/flows/components/paths.py
@abstractmethod
def interpolate(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
    """
    Compute interpolated point x_t on the path.

    Args:
        x_0: Starting point (noise), shape (B, *)
        x_1: Ending point (data), shape (B, *)
        t: Time in [0, 1], shape (B, 1)

    Returns:
        x_t: Interpolated point, shape (B, *)
    """
    pass
velocity(x_0, x_1, t) abstractmethod

Compute target velocity v_t = dx_t/dt.

Parameters:

Name Type Description Default
x_0 Tensor

Starting point (noise), shape (B, *)

required
x_1 Tensor

Ending point (data), shape (B, *)

required
t Tensor

Time in [0, 1], shape (B, 1)

required

Returns:

Name Type Description
v_t Tensor

Target velocity, shape (B, *)

Source code in flowpde/flows/components/paths.py
@abstractmethod
def velocity(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
    """
    Compute target velocity v_t = dx_t/dt.

    Args:
        x_0: Starting point (noise), shape (B, *)
        x_1: Ending point (data), shape (B, *)
        t: Time in [0, 1], shape (B, 1)

    Returns:
        v_t: Target velocity, shape (B, *)
    """
    pass
__call__(x_0, x_1, t)

Compute both interpolation and velocity.

Returns:

Type Description
(x_t, v_t)

Interpolated point and target velocity

Source code in flowpde/flows/components/paths.py
def __call__(
    self,
    x_0: Tensor,
    x_1: Tensor,
    t: Tensor
) -> Tuple[Tensor, Tensor]:
    """
    Compute both interpolation and velocity.

    Returns:
        (x_t, v_t): Interpolated point and target velocity
    """
    x_t = self.interpolate(x_0, x_1, t)
    v_t = self.velocity(x_0, x_1, t)
    return x_t, v_t

LinearPath

Bases: PathInterpolant

Linear interpolation path (straight lines).

Path: x_t = (1 - t) * x_0 + t * x_1 Velocity: v_t = x_1 - x_0 (constant along path)

This is the standard path used in: - Flow Matching (Lipman et al., 2023) - Rectified Flow (Liu et al., 2023)

Properties: - Straight lines from noise to data - Constant velocity (simplest to learn) - x_0 at t=0, x_1 at t=1

Source code in flowpde/flows/components/paths.py
class LinearPath(PathInterpolant):
    """
    Linear interpolation path (straight lines).

    Path: x_t = (1 - t) * x_0 + t * x_1
    Velocity: v_t = x_1 - x_0 (constant along path)

    This is the standard path used in:
    - Flow Matching (Lipman et al., 2023)
    - Rectified Flow (Liu et al., 2023)

    Properties:
    - Straight lines from noise to data
    - Constant velocity (simplest to learn)
    - x_0 at t=0, x_1 at t=1
    """

    def interpolate(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
        """x_t = (1 - t) * x_0 + t * x_1"""
        t_expanded = self._expand_t(t, x_0.dim())
        return (1 - t_expanded) * x_0 + t_expanded * x_1

    def velocity(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
        """v_t = x_1 - x_0 (constant)"""
        return x_1 - x_0

    def _expand_t(self, t: Tensor, ndim: int) -> Tensor:
        """Expand t for broadcasting: (B, 1) -> (B, 1, 1, ..., 1)"""
        return t.view(-1, *([1] * (ndim - 1)))
interpolate(x_0, x_1, t)

x_t = (1 - t) * x_0 + t * x_1

Source code in flowpde/flows/components/paths.py
def interpolate(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
    """x_t = (1 - t) * x_0 + t * x_1"""
    t_expanded = self._expand_t(t, x_0.dim())
    return (1 - t_expanded) * x_0 + t_expanded * x_1
velocity(x_0, x_1, t)

v_t = x_1 - x_0 (constant)

Source code in flowpde/flows/components/paths.py
def velocity(self, x_0: Tensor, x_1: Tensor, t: Tensor) -> Tensor:
    """v_t = x_1 - x_0 (constant)"""
    return x_1 - x_0

OTConditionalPath

Bases: PathInterpolant

Optimal-transport conditional flow path (Tong et al., 2023).

Two noise schedules are available, and both regress the exact derivative of the path the model is shown.

schedule='constant' (default) — the OT-CFM path of Tong et al.:

x_t = t * x_1 + (1 - t) * x_0 + sigma * eps
v_t = x_1 - x_0

The noise term does not depend on t, so the chord is the conditional velocity. This is what makes sigma > 0 usable: the regression target stays bounded everywhere.

schedule='bridge' — a Brownian-bridge tube that pins both endpoints:

x_t = t * x_1 + (1 - t) * x_0 + sigma * sqrt(t(1-t)) * eps
v_t = x_1 - x_0 + sigma * (1 - 2t) / (2 * sqrt(t(1-t))) * eps

Here the tube width is time dependent, so the chord alone is not the derivative: the second term is required, or the model regresses a target that does not match its own input. That term diverges as t approaches 0 or 1 (it is clamped, but stays large), which is why the constant schedule is the default.

With sigma = 0 both schedules reduce exactly to LinearPath.

Parameters:

Name Type Description Default
sigma float

Noise scale (default: 0.0).

0.0
schedule str

'constant' (default) or 'bridge'.

'constant'
Source code in flowpde/flows/components/paths.py
class OTConditionalPath(PathInterpolant):
    r"""
    Optimal-transport conditional flow path (Tong et al., 2023).

    Two noise schedules are available, and both regress the *exact* derivative
    of the path the model is shown.

    ``schedule='constant'`` (default) — the OT-CFM path of Tong et al.:

        x_t = t * x_1 + (1 - t) * x_0 + sigma * eps
        v_t = x_1 - x_0

    The noise term does not depend on t, so the chord *is* the conditional
    velocity.  This is what makes sigma > 0 usable: the regression target
    stays bounded everywhere.

    ``schedule='bridge'`` — a Brownian-bridge tube that pins both endpoints:

        x_t = t * x_1 + (1 - t) * x_0 + sigma * sqrt(t(1-t)) * eps
        v_t = x_1 - x_0 + sigma * (1 - 2t) / (2 * sqrt(t(1-t))) * eps

    Here the tube width is time dependent, so the chord alone is *not* the
    derivative: the second term is required, or the model regresses a target
    that does not match its own input.  That term diverges as t approaches 0
    or 1 (it is clamped, but stays large), which is why the constant schedule
    is the default.

    With sigma = 0 both schedules reduce exactly to `LinearPath`.

    Args:
        sigma: Noise scale (default: 0.0).
        schedule: `'constant'` (default) or `'bridge'`.
    """

    _SCHEDULES = ("constant", "bridge")

    # Floor on t(1-t) inside the bridge derivative, so the endpoints give a
    # large-but-finite target instead of an infinity.
    _BRIDGE_FLOOR = 1e-6

    def __init__(self, sigma: float = 0.0, schedule: str = "constant"):
        if schedule not in self._SCHEDULES:
            raise ValueError(
                f"Unknown schedule: '{schedule}'. "
                f"Available: {list(self._SCHEDULES)}"
            )
        self.sigma = sigma
        self.schedule = schedule

    def _expand_t(self, t: Tensor, ndim: int) -> Tensor:
        """Expand t for broadcasting."""
        return t.view(-1, *([1] * (ndim - 1)))

    def _tube_width(self, t_expanded: Tensor) -> Union[float, Tensor]:
        r"""$\sigma_t$, the width of the conditional tube at time t."""
        if self.schedule == "constant":
            return self.sigma
        return self.sigma * torch.sqrt(
            (t_expanded * (1 - t_expanded)).clamp(min=0.0)
        )

    def _tube_width_derivative(self, t_expanded: Tensor) -> Union[float, Tensor]:
        r"""$d\sigma_t/dt$, the term the chord alone leaves out."""
        if self.schedule == "constant":
            return 0.0
        denominator = 2 * torch.sqrt(
            (t_expanded * (1 - t_expanded)).clamp(min=self._BRIDGE_FLOOR)
        )
        return self.sigma * (1 - 2 * t_expanded) / denominator

    def interpolate(
        self,
        x_0: Tensor,
        x_1: Tensor,
        t: Tensor,
        noise: Optional[Tensor] = None,
    ) -> Tensor:
        """x_t on the conditional path.

        Args:
            noise: The epsilon to use.  Pass the same one to `velocity()` so
                the target is the derivative of *this* sample path; `__call__`
                does that for you.
        """
        t_expanded = self._expand_t(t, x_0.dim())
        x_t = t_expanded * x_1 + (1 - t_expanded) * x_0

        if self.sigma > 0:
            noise = torch.randn_like(x_0) if noise is None else noise
            x_t = x_t + self._tube_width(t_expanded) * noise

        return x_t

    def velocity(
        self,
        x_0: Tensor,
        x_1: Tensor,
        t: Tensor,
        noise: Optional[Tensor] = None,
    ) -> Tensor:
        """Exact derivative dx_t/dt of the path defined by `interpolate`."""
        v_t = x_1 - x_0

        if self.sigma > 0 and self.schedule != "constant":
            t_expanded = self._expand_t(t, x_0.dim())
            noise = torch.randn_like(x_0) if noise is None else noise
            v_t = v_t + self._tube_width_derivative(t_expanded) * noise

        return v_t

    def __call__(
        self,
        x_0: Tensor,
        x_1: Tensor,
        t: Tensor,
    ) -> Tuple[Tensor, Tensor]:
        """Interpolated point and its target velocity, from a single epsilon.

        Drawing epsilon twice would hand the model an input from one sample
        path and the derivative of a different one.
        """
        noise = torch.randn_like(x_0) if self.sigma > 0 else None
        return (
            self.interpolate(x_0, x_1, t, noise=noise),
            self.velocity(x_0, x_1, t, noise=noise),
        )
interpolate(x_0, x_1, t, noise=None)

x_t on the conditional path.

Parameters:

Name Type Description Default
noise Optional[Tensor]

The epsilon to use. Pass the same one to velocity() so the target is the derivative of this sample path; __call__ does that for you.

None
Source code in flowpde/flows/components/paths.py
def interpolate(
    self,
    x_0: Tensor,
    x_1: Tensor,
    t: Tensor,
    noise: Optional[Tensor] = None,
) -> Tensor:
    """x_t on the conditional path.

    Args:
        noise: The epsilon to use.  Pass the same one to `velocity()` so
            the target is the derivative of *this* sample path; `__call__`
            does that for you.
    """
    t_expanded = self._expand_t(t, x_0.dim())
    x_t = t_expanded * x_1 + (1 - t_expanded) * x_0

    if self.sigma > 0:
        noise = torch.randn_like(x_0) if noise is None else noise
        x_t = x_t + self._tube_width(t_expanded) * noise

    return x_t
velocity(x_0, x_1, t, noise=None)

Exact derivative dx_t/dt of the path defined by interpolate.

Source code in flowpde/flows/components/paths.py
def velocity(
    self,
    x_0: Tensor,
    x_1: Tensor,
    t: Tensor,
    noise: Optional[Tensor] = None,
) -> Tensor:
    """Exact derivative dx_t/dt of the path defined by `interpolate`."""
    v_t = x_1 - x_0

    if self.sigma > 0 and self.schedule != "constant":
        t_expanded = self._expand_t(t, x_0.dim())
        noise = torch.randn_like(x_0) if noise is None else noise
        v_t = v_t + self._tube_width_derivative(t_expanded) * noise

    return v_t
__call__(x_0, x_1, t)

Interpolated point and its target velocity, from a single epsilon.

Drawing epsilon twice would hand the model an input from one sample path and the derivative of a different one.

Source code in flowpde/flows/components/paths.py
def __call__(
    self,
    x_0: Tensor,
    x_1: Tensor,
    t: Tensor,
) -> Tuple[Tensor, Tensor]:
    """Interpolated point and its target velocity, from a single epsilon.

    Drawing epsilon twice would hand the model an input from one sample
    path and the derivative of a different one.
    """
    noise = torch.randn_like(x_0) if self.sigma > 0 else None
    return (
        self.interpolate(x_0, x_1, t, noise=noise),
        self.velocity(x_0, x_1, t, noise=noise),
    )

get_path(path, **kwargs)

Get a path interpolant by name or return if already an instance.

Parameters:

Name Type Description Default
path Union[str, PathInterpolant]

Path name ('linear', 'ot_conditional') or PathInterpolant instance

required
**kwargs Any

Additional arguments for path constructor

{}

Returns:

Type Description
PathInterpolant

PathInterpolant instance

Examples:

>>> path = get_path('linear')
>>> path = get_path('ot_conditional', sigma=0.01)
>>> path = get_path(LinearPath())  # Pass-through
Source code in flowpde/flows/components/paths.py
def get_path(path: Union[str, PathInterpolant], **kwargs: Any) -> PathInterpolant:
    """
    Get a path interpolant by name or return if already an instance.

    Args:
        path: Path name ('linear', 'ot_conditional') or PathInterpolant instance
        **kwargs: Additional arguments for path constructor

    Returns:
        PathInterpolant instance

    Examples:
        >>> path = get_path('linear')
        >>> path = get_path('ot_conditional', sigma=0.01)
        >>> path = get_path(LinearPath())  # Pass-through
    """
    if isinstance(path, PathInterpolant):
        return path

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

    return _PATH_REGISTRY[path](**kwargs)

Time Samplers

time_samplers

Time Sampling Strategies for Flow Matching.

Defines distributions for sampling time t ∈ [0, 1] during training. Different sampling strategies can improve training dynamics.

TimeSampler

Bases: ABC

Abstract base class for time sampling distributions.

Different time distributions can affect training: - Uniform: Standard, simple - Logit-normal: Concentrates samples near t=0 and t=1 - Beta: Flexible, can emphasize different regions

Source code in flowpde/flows/components/time_samplers.py
class TimeSampler(ABC):
    """
    Abstract base class for time sampling distributions.

    Different time distributions can affect training:
    - Uniform: Standard, simple
    - Logit-normal: Concentrates samples near t=0 and t=1
    - Beta: Flexible, can emphasize different regions
    """

    @abstractmethod
    def sample(self, batch_size: int, device: torch.device) -> Tensor:
        """
        Sample time values in [0, 1].

        Args:
            batch_size: Number of samples
            device: Device for tensor

        Returns:
            Time tensor of shape (batch_size, 1)
        """
        pass

    def __call__(self, batch_size: int, device: torch.device) -> Tensor:
        """Alias for sample()."""
        return self.sample(batch_size, device)
sample(batch_size, device) abstractmethod

Sample time values in [0, 1].

Parameters:

Name Type Description Default
batch_size int

Number of samples

required
device device

Device for tensor

required

Returns:

Type Description
Tensor

Time tensor of shape (batch_size, 1)

Source code in flowpde/flows/components/time_samplers.py
@abstractmethod
def sample(self, batch_size: int, device: torch.device) -> Tensor:
    """
    Sample time values in [0, 1].

    Args:
        batch_size: Number of samples
        device: Device for tensor

    Returns:
        Time tensor of shape (batch_size, 1)
    """
    pass
__call__(batch_size, device)

Alias for sample().

Source code in flowpde/flows/components/time_samplers.py
def __call__(self, batch_size: int, device: torch.device) -> Tensor:
    """Alias for sample()."""
    return self.sample(batch_size, device)

UniformSampler

Bases: TimeSampler

Uniform time sampling: t ~ U(low, high).

Standard sampling strategy for flow matching.

Parameters:

Name Type Description Default
low float

Lower bound (default: 0.0)

0.0
high float

Upper bound (default: 1.0)

1.0
Source code in flowpde/flows/components/time_samplers.py
class UniformSampler(TimeSampler):
    """
    Uniform time sampling: t ~ U(low, high).

    Standard sampling strategy for flow matching.

    Args:
        low: Lower bound (default: 0.0)
        high: Upper bound (default: 1.0)
    """

    def __init__(self, low: float = 0.0, high: float = 1.0):
        self.low = low
        self.high = high

    def sample(self, batch_size: int, device: torch.device) -> Tensor:
        """Sample t ~ U(low, high)."""
        return torch.rand(batch_size, 1, device=device) * (self.high - self.low) + self.low
sample(batch_size, device)

Sample t ~ U(low, high).

Source code in flowpde/flows/components/time_samplers.py
def sample(self, batch_size: int, device: torch.device) -> Tensor:
    """Sample t ~ U(low, high)."""
    return torch.rand(batch_size, 1, device=device) * (self.high - self.low) + self.low

LogitNormalSampler

Bases: TimeSampler

Logit-normal time sampling: t = sigmoid(z) where z ~ N(mean, std²).

This distribution concentrates samples near t=0 and t=1, which can help the model learn the behavior at the boundaries better.

Used in Rectified Flow and some diffusion models.

Parameters:

Name Type Description Default
mean float

Mean of the normal distribution (default: 0.0)

0.0
std float

Standard deviation (default: 1.0)

1.0
Source code in flowpde/flows/components/time_samplers.py
class LogitNormalSampler(TimeSampler):
    """
    Logit-normal time sampling: t = sigmoid(z) where z ~ N(mean, std²).

    This distribution concentrates samples near t=0 and t=1, which can
    help the model learn the behavior at the boundaries better.

    Used in Rectified Flow and some diffusion models.

    Args:
        mean: Mean of the normal distribution (default: 0.0)
        std: Standard deviation (default: 1.0)
    """

    def __init__(self, mean: float = 0.0, std: float = 1.0):
        self.mean = mean
        self.std = std

    def sample(self, batch_size: int, device: torch.device) -> Tensor:
        """Sample t = sigmoid(z) where z ~ N(mean, std²)."""
        z = torch.randn(batch_size, 1, device=device) * self.std + self.mean
        t = torch.sigmoid(z)
        return t
sample(batch_size, device)

Sample t = sigmoid(z) where z ~ N(mean, std²).

Source code in flowpde/flows/components/time_samplers.py
def sample(self, batch_size: int, device: torch.device) -> Tensor:
    """Sample t = sigmoid(z) where z ~ N(mean, std²)."""
    z = torch.randn(batch_size, 1, device=device) * self.std + self.mean
    t = torch.sigmoid(z)
    return t

BetaSampler

Bases: TimeSampler

Beta distribution time sampling: t ~ Beta(alpha, beta).

Flexible distribution that can emphasize different regions: - alpha=beta=1: Uniform - alpha=beta=0.5: U-shaped (emphasizes boundaries) - alpha=beta>1: Concentrated around 0.5

Parameters:

Name Type Description Default
alpha float

First shape parameter (default: 1.0)

1.0
beta float

Second shape parameter (default: 1.0)

1.0
Source code in flowpde/flows/components/time_samplers.py
class BetaSampler(TimeSampler):
    """
    Beta distribution time sampling: t ~ Beta(alpha, beta).

    Flexible distribution that can emphasize different regions:
    - alpha=beta=1: Uniform
    - alpha=beta=0.5: U-shaped (emphasizes boundaries)
    - alpha=beta>1: Concentrated around 0.5

    Args:
        alpha: First shape parameter (default: 1.0)
        beta: Second shape parameter (default: 1.0)
    """

    def __init__(self, alpha: float = 1.0, beta: float = 1.0):
        self.alpha = alpha
        self.beta = beta
        self._dist = torch.distributions.Beta(alpha, beta)

    def sample(self, batch_size: int, device: torch.device) -> Tensor:
        """Sample t ~ Beta(alpha, beta)."""
        # Sample and move to device
        t = self._dist.sample((batch_size, 1))
        return t.to(device)
sample(batch_size, device)

Sample t ~ Beta(alpha, beta).

Source code in flowpde/flows/components/time_samplers.py
def sample(self, batch_size: int, device: torch.device) -> Tensor:
    """Sample t ~ Beta(alpha, beta)."""
    # Sample and move to device
    t = self._dist.sample((batch_size, 1))
    return t.to(device)

get_time_sampler(sampler, **kwargs)

Get a time sampler by name or return if already an instance.

Parameters:

Name Type Description Default
sampler Union[str, TimeSampler]

Sampler name ('uniform', 'logit_normal', 'beta') or instance

required
**kwargs Any

Additional arguments for sampler constructor

{}

Returns:

Type Description
TimeSampler

TimeSampler instance

Examples:

>>> sampler = get_time_sampler('uniform')
>>> sampler = get_time_sampler('logit_normal', std=0.5)
>>> sampler = get_time_sampler('beta', alpha=0.5, beta=0.5)
Source code in flowpde/flows/components/time_samplers.py
def get_time_sampler(
    sampler: Union[str, TimeSampler],
    **kwargs: Any
) -> TimeSampler:
    """
    Get a time sampler by name or return if already an instance.

    Args:
        sampler: Sampler name ('uniform', 'logit_normal', 'beta') or instance
        **kwargs: Additional arguments for sampler constructor

    Returns:
        TimeSampler instance

    Examples:
        >>> sampler = get_time_sampler('uniform')
        >>> sampler = get_time_sampler('logit_normal', std=0.5)
        >>> sampler = get_time_sampler('beta', alpha=0.5, beta=0.5)
    """
    if isinstance(sampler, TimeSampler):
        return sampler

    if sampler not in _SAMPLER_REGISTRY:
        raise ValueError(
            f"Unknown time sampler: '{sampler}'. "
            f"Available: {list(_SAMPLER_REGISTRY.keys())}"
        )

    return _SAMPLER_REGISTRY[sampler](**kwargs)

Couplings

couplings

Coupling Strategies for Flow Matching.

Defines how noise samples (x_0) and data samples (x_1) are paired during training. Different couplings can affect training efficiency.

Coupling

Bases: ABC

Abstract base class for noise-data coupling strategies.

A coupling defines how to pair samples from the base distribution (noise, x_0) with samples from the data distribution (x_1).

Standard flow matching uses independent coupling, but optimal transport (OT) couplings can improve training.

Source code in flowpde/flows/components/couplings.py
class Coupling(ABC):
    """
    Abstract base class for noise-data coupling strategies.

    A coupling defines how to pair samples from the base distribution
    (noise, x_0) with samples from the data distribution (x_1).

    Standard flow matching uses independent coupling, but optimal
    transport (OT) couplings can improve training.
    """

    @abstractmethod
    def couple(
        self,
        x_0: Tensor,
        x_1: Tensor
    ) -> Tuple[Tensor, Tensor]:
        """
        Apply coupling strategy to pair noise and data.

        Args:
            x_0: Noise samples, shape (B, *)
            x_1: Data samples, shape (B, *)

        Returns:
            (x_0_coupled, x_1_coupled): Paired samples
        """
        pass

    def __call__(
        self,
        x_0: Tensor,
        x_1: Tensor
    ) -> Tuple[Tensor, Tensor]:
        """Alias for couple()."""
        return self.couple(x_0, x_1)
couple(x_0, x_1) abstractmethod

Apply coupling strategy to pair noise and data.

Parameters:

Name Type Description Default
x_0 Tensor

Noise samples, shape (B, *)

required
x_1 Tensor

Data samples, shape (B, *)

required

Returns:

Type Description
(x_0_coupled, x_1_coupled)

Paired samples

Source code in flowpde/flows/components/couplings.py
@abstractmethod
def couple(
    self,
    x_0: Tensor,
    x_1: Tensor
) -> Tuple[Tensor, Tensor]:
    """
    Apply coupling strategy to pair noise and data.

    Args:
        x_0: Noise samples, shape (B, *)
        x_1: Data samples, shape (B, *)

    Returns:
        (x_0_coupled, x_1_coupled): Paired samples
    """
    pass
__call__(x_0, x_1)

Alias for couple().

Source code in flowpde/flows/components/couplings.py
def __call__(
    self,
    x_0: Tensor,
    x_1: Tensor
) -> Tuple[Tensor, Tensor]:
    """Alias for couple()."""
    return self.couple(x_0, x_1)

IndependentCoupling

Bases: Coupling

Independent coupling: no reordering.

Each noise sample x_0[i] is paired with data sample x_1[i] as they appear in the batch. This is the standard approach.

Simple and fast, but may not be optimal for transport.

Source code in flowpde/flows/components/couplings.py
class IndependentCoupling(Coupling):
    """
    Independent coupling: no reordering.

    Each noise sample x_0[i] is paired with data sample x_1[i]
    as they appear in the batch. This is the standard approach.

    Simple and fast, but may not be optimal for transport.
    """

    def couple(
        self,
        x_0: Tensor,
        x_1: Tensor
    ) -> Tuple[Tensor, Tensor]:
        """Identity coupling: return inputs unchanged."""
        return x_0, x_1
couple(x_0, x_1)

Identity coupling: return inputs unchanged.

Source code in flowpde/flows/components/couplings.py
def couple(
    self,
    x_0: Tensor,
    x_1: Tensor
) -> Tuple[Tensor, Tensor]:
    """Identity coupling: return inputs unchanged."""
    return x_0, x_1

MiniBatchOTCoupling

Bases: Coupling

Mini-batch Optimal Transport coupling.

Reorders the noise samples within each mini-batch to minimize total transport cost (squared Euclidean distance).

This can lead to more efficient training by creating shorter transport paths on average.

Requires scipy for the linear assignment. scipy is a hard dependency of this package, so a missing install is an error rather than a reason to quietly fall back to independent coupling -- that fallback would turn an OT ablation into a duplicate of the baseline with nothing to show for it.

Parameters:

Name Type Description Default
cost_fn str

Cost function ('euclidean', 'cosine'). Default: 'euclidean'

'euclidean'
Source code in flowpde/flows/components/couplings.py
class MiniBatchOTCoupling(Coupling):
    """
    Mini-batch Optimal Transport coupling.

    Reorders the noise samples within each mini-batch to minimize
    total transport cost (squared Euclidean distance).

    This can lead to more efficient training by creating shorter
    transport paths on average.

    Requires scipy for the linear assignment.  scipy is a hard dependency of
    this package, so a missing install is an error rather than a reason to
    quietly fall back to independent coupling -- that fallback would turn an
    OT ablation into a duplicate of the baseline with nothing to show for it.

    Args:
        cost_fn: Cost function ('euclidean', 'cosine'). Default: 'euclidean'
    """

    def __init__(self, cost_fn: str = 'euclidean'):
        if cost_fn not in ('euclidean', 'cosine'):
            raise ValueError(
                f"Unknown cost function: {cost_fn}. "
                "Choose from ['euclidean', 'cosine']."
            )
        self.cost_fn = cost_fn

    @staticmethod
    def _linear_sum_assignment():
        """Return scipy's solver, or explain why it is missing."""
        try:
            from scipy.optimize import linear_sum_assignment
        except ImportError as error:
            raise ImportError(
                "MiniBatchOTCoupling needs scipy for the linear assignment. "
                "Install it with `pip install scipy` (it is already a declared "
                "dependency of flowpde)."
            ) from error
        return linear_sum_assignment


    def couple(
        self,
        x_0: Tensor,
        x_1: Tensor
    ) -> Tuple[Tensor, Tensor]:
        """
        Compute OT coupling via linear assignment.

        Finds permutation of x_0 that minimizes total transport cost.
        """
        linear_sum_assignment = self._linear_sum_assignment()

        batch_size = x_0.shape[0]

        # Flatten for distance computation
        x_0_flat = x_0.view(batch_size, -1)
        x_1_flat = x_1.view(batch_size, -1)

        # Compute cost matrix
        if self.cost_fn == 'euclidean':
            # Squared Euclidean distance
            cost = torch.cdist(x_0_flat, x_1_flat, p=2).pow(2)
        elif self.cost_fn == 'cosine':
            # Cosine distance
            x_0_norm = x_0_flat / (x_0_flat.norm(dim=1, keepdim=True) + 1e-8)
            x_1_norm = x_1_flat / (x_1_flat.norm(dim=1, keepdim=True) + 1e-8)
            cost = 1 - torch.mm(x_0_norm, x_1_norm.t())

        # Solve linear assignment (on CPU)
        cost_np = cost.detach().cpu().numpy()
        row_ind, col_ind = linear_sum_assignment(cost_np)

        # Reorder x_0 to match optimal assignment
        # col_ind[i] is the index in x_1 that x_0[i] should map to
        # We need to reorder x_0 so that x_0[col_ind] maps to x_1
        inverse_perm = torch.argsort(torch.tensor(col_ind, device=x_0.device))
        x_0_coupled = x_0[inverse_perm]

        return x_0_coupled, x_1
couple(x_0, x_1)

Compute OT coupling via linear assignment.

Finds permutation of x_0 that minimizes total transport cost.

Source code in flowpde/flows/components/couplings.py
def couple(
    self,
    x_0: Tensor,
    x_1: Tensor
) -> Tuple[Tensor, Tensor]:
    """
    Compute OT coupling via linear assignment.

    Finds permutation of x_0 that minimizes total transport cost.
    """
    linear_sum_assignment = self._linear_sum_assignment()

    batch_size = x_0.shape[0]

    # Flatten for distance computation
    x_0_flat = x_0.view(batch_size, -1)
    x_1_flat = x_1.view(batch_size, -1)

    # Compute cost matrix
    if self.cost_fn == 'euclidean':
        # Squared Euclidean distance
        cost = torch.cdist(x_0_flat, x_1_flat, p=2).pow(2)
    elif self.cost_fn == 'cosine':
        # Cosine distance
        x_0_norm = x_0_flat / (x_0_flat.norm(dim=1, keepdim=True) + 1e-8)
        x_1_norm = x_1_flat / (x_1_flat.norm(dim=1, keepdim=True) + 1e-8)
        cost = 1 - torch.mm(x_0_norm, x_1_norm.t())

    # Solve linear assignment (on CPU)
    cost_np = cost.detach().cpu().numpy()
    row_ind, col_ind = linear_sum_assignment(cost_np)

    # Reorder x_0 to match optimal assignment
    # col_ind[i] is the index in x_1 that x_0[i] should map to
    # We need to reorder x_0 so that x_0[col_ind] maps to x_1
    inverse_perm = torch.argsort(torch.tensor(col_ind, device=x_0.device))
    x_0_coupled = x_0[inverse_perm]

    return x_0_coupled, x_1

get_coupling(coupling, **kwargs)

Get a coupling strategy by name or return if already an instance.

Parameters:

Name Type Description Default
coupling Union[str, Coupling]

Coupling name ('independent', 'minibatch_ot') or instance

required
**kwargs Any

Additional arguments for coupling constructor

{}

Returns:

Type Description
Coupling

Coupling instance

Examples:

>>> coupling = get_coupling('independent')
>>> coupling = get_coupling('minibatch_ot', cost_fn='cosine')
Source code in flowpde/flows/components/couplings.py
def get_coupling(
    coupling: Union[str, Coupling],
    **kwargs: Any
) -> Coupling:
    """
    Get a coupling strategy by name or return if already an instance.

    Args:
        coupling: Coupling name ('independent', 'minibatch_ot') or instance
        **kwargs: Additional arguments for coupling constructor

    Returns:
        Coupling instance

    Examples:
        >>> coupling = get_coupling('independent')
        >>> coupling = get_coupling('minibatch_ot', cost_fn='cosine')
    """
    if isinstance(coupling, Coupling):
        return coupling

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

    return _COUPLING_REGISTRY[coupling](**kwargs)

Sources

Where trajectories start. BatchSource is what makes reflow correct — it consumes the exact noise z that produced each generated target instead of resampling.

sources

Source Distributions for Flow Matching.

The source distribution supplies \(x_0\), the point a trajectory starts from. It is the fourth pluggable component of the flow-matching objective, alongside the path, the time sampler and the coupling.

Making it configurable is what allows the objective to be trained on precomputed pairs rather than freshly drawn noise. That distinction is the whole mechanism behind reflow: reflow retrains on pairs \((z, \mathrm{ODE}(z))\) generated by the current model, and the straightening only happens if training uses that \(z\) rather than an independently resampled one. Drawing fresh noise silently turns reflow into a no-op.

Available sources:

  • GaussianSource — standard normal noise (the default).
  • BatchSource — read \(x_0\) from the training batch, falling back to another source when the key is absent (i.e. at inference time).

SourceDistribution

Bases: ABC

Abstract base class for flow source distributions.

A source answers one question: given a target shape and a training batch, where does the trajectory start?

Source code in flowpde/flows/components/sources.py
class SourceDistribution(ABC):
    """
    Abstract base class for flow source distributions.

    A source answers one question: given a target shape and a training batch,
    where does the trajectory start?
    """

    @abstractmethod
    def sample(
        self,
        shape: Tuple[int, ...],
        device: torch.device,
        batch: Optional[Dict[str, Tensor]] = None,
    ) -> Tensor:
        """
        Draw $x_0$.

        Args:
            shape: Shape of the target `x_1`, which `x_0` must match.
            device: Device to place the result on.
            batch: The training batch, when one is available.  Sources that
                read precomputed values use it; others ignore it.

        Returns:
            Tensor of shape `shape`.
        """
        raise NotImplementedError

    def __call__(
        self,
        shape: Tuple[int, ...],
        device: torch.device,
        batch: Optional[Dict[str, Tensor]] = None,
    ) -> Tensor:
        return self.sample(shape, device, batch)

    def get_config(self) -> Dict[str, Any]:
        return {"type": self.__class__.__name__}
sample(shape, device, batch=None) abstractmethod

Draw \(x_0\).

Parameters:

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

Shape of the target x_1, which x_0 must match.

required
device device

Device to place the result on.

required
batch Optional[Dict[str, Tensor]]

The training batch, when one is available. Sources that read precomputed values use it; others ignore it.

None

Returns:

Type Description
Tensor

Tensor of shape shape.

Source code in flowpde/flows/components/sources.py
@abstractmethod
def sample(
    self,
    shape: Tuple[int, ...],
    device: torch.device,
    batch: Optional[Dict[str, Tensor]] = None,
) -> Tensor:
    """
    Draw $x_0$.

    Args:
        shape: Shape of the target `x_1`, which `x_0` must match.
        device: Device to place the result on.
        batch: The training batch, when one is available.  Sources that
            read precomputed values use it; others ignore it.

    Returns:
        Tensor of shape `shape`.
    """
    raise NotImplementedError

GaussianSource

Bases: SourceDistribution

Standard Gaussian source, \(x_0 \sim \mathcal{N}(0, \sigma^2 I)\).

The default for flow matching, and what inference uses regardless of how training was configured.

Parameters:

Name Type Description Default
std float

Standard deviation (default: 1.0).

1.0
Source code in flowpde/flows/components/sources.py
class GaussianSource(SourceDistribution):
    """
    Standard Gaussian source, $x_0 \\sim \\mathcal{N}(0, \\sigma^2 I)$.

    The default for flow matching, and what inference uses regardless of how
    training was configured.

    Args:
        std: Standard deviation (default: 1.0).
    """

    def __init__(self, std: float = 1.0):
        self.std = std

    def sample(
        self,
        shape: Tuple[int, ...],
        device: torch.device,
        batch: Optional[Dict[str, Tensor]] = None,
    ) -> Tensor:
        noise = torch.randn(*shape, device=device)
        return noise * self.std if self.std != 1.0 else noise

    def get_config(self) -> Dict[str, Any]:
        return {"type": "GaussianSource", "std": self.std}

    def __repr__(self) -> str:
        return f"GaussianSource(std={self.std})"

BatchSource

Bases: SourceDistribution

Read \(x_0\) from the training batch.

Used for reflow, where the pairing between source and target is fixed in advance and must be preserved: the model is retrained on exactly the trajectories it generated, which is what straightens them.

When the key is absent — most importantly during sampling, where there is no batch — the fallback source is used instead. So a model trained with BatchSource still samples from \(\mathcal{N}(0, I)\) at inference, as it should.

Parameters:

Name Type Description Default
key str

Batch key holding the precomputed x_0 (default: 'x_0').

'x_0'
fallback Optional[SourceDistribution]

Source used when the key is missing. Defaults to GaussianSource.

None
strict bool

When True, raise instead of falling back if a batch is supplied but lacks the key. Catches silently-degraded reflow runs, and is the default for that reason.

True
Source code in flowpde/flows/components/sources.py
class BatchSource(SourceDistribution):
    """
    Read $x_0$ from the training batch.

    Used for reflow, where the pairing between source and target is fixed in
    advance and *must* be preserved: the model is retrained on exactly the
    trajectories it generated, which is what straightens them.

    When the key is absent — most importantly during sampling, where there is
    no batch — the `fallback` source is used instead.  So a model trained
    with `BatchSource` still samples from $\\mathcal{N}(0, I)$ at
    inference, as it should.

    Args:
        key: Batch key holding the precomputed `x_0` (default: `'x_0'`).
        fallback: Source used when the key is missing.  Defaults to
            `GaussianSource`.
        strict: When True, raise instead of falling back if a batch is
            supplied but lacks the key.  Catches silently-degraded reflow
            runs, and is the default for that reason.
    """

    def __init__(
        self,
        key: str = "x_0",
        fallback: Optional[SourceDistribution] = None,
        strict: bool = True,
    ):
        self.key = key
        self.fallback = fallback if fallback is not None else GaussianSource()
        self.strict = strict

    def sample(
        self,
        shape: Tuple[int, ...],
        device: torch.device,
        batch: Optional[Dict[str, Tensor]] = None,
    ) -> Tensor:
        if batch is None:
            # No batch at all: this is inference, so fall back silently.
            return self.fallback.sample(shape, device, batch)

        if self.key not in batch:
            if self.strict:
                available = ", ".join(sorted(batch.keys()))
                raise KeyError(
                    f"BatchSource expected key '{self.key}' in the batch but "
                    f"found only [{available}]. Reflow training must supply "
                    f"the precomputed x_0; resampling noise instead would "
                    f"silently disable path straightening. Pass strict=False "
                    f"to allow falling back to {type(self.fallback).__name__}."
                )
            return self.fallback.sample(shape, device, batch)

        x_0 = batch[self.key]
        x_0 = x_0.flatten(start_dim=1).to(device)

        if x_0.shape != tuple(shape):
            raise ValueError(
                f"Precomputed x_0 has shape {tuple(x_0.shape)} but the target "
                f"requires {tuple(shape)}."
            )
        return x_0

    def get_config(self) -> Dict[str, Any]:
        return {
            "type": "BatchSource",
            "key": self.key,
            "strict": self.strict,
            "fallback": self.fallback.get_config(),
        }

    def __repr__(self) -> str:
        return f"BatchSource(key='{self.key}', fallback={self.fallback!r})"

get_source(source, **kwargs)

Get a source distribution by name, or pass an instance through.

Parameters:

Name Type Description Default
source Union[str, SourceDistribution]

Name ('gaussian', 'batch') or a SourceDistribution instance.

required
**kwargs Any

Forwarded to the constructor when a name is given.

{}

Returns:

Type Description
SourceDistribution

A SourceDistribution.

Examples:

>>> get_source('gaussian')
GaussianSource(std=1.0)
>>> isinstance(get_source(GaussianSource()), GaussianSource)
True
Source code in flowpde/flows/components/sources.py
def get_source(
    source: Union[str, SourceDistribution], **kwargs: Any
) -> SourceDistribution:
    """
    Get a source distribution by name, or pass an instance through.

    Args:
        source: Name (`'gaussian'`, `'batch'`) or a
            `SourceDistribution` instance.
        **kwargs: Forwarded to the constructor when a name is given.

    Returns:
        A `SourceDistribution`.

    Examples:
        >>> get_source('gaussian')
        GaussianSource(std=1.0)
        >>> isinstance(get_source(GaussianSource()), GaussianSource)
        True
    """
    if isinstance(source, SourceDistribution):
        return source

    if source not in _SOURCE_REGISTRY:
        raise ValueError(
            f"Unknown source: '{source}'. Available: {list(_SOURCE_REGISTRY)}"
        )

    return _SOURCE_REGISTRY[source](**kwargs)