Skip to content

ODE Solvers

ODE integration methods for flow-based sampling.

ode_solvers

ODE Solvers for Flow Matching Inference using torchdiffeq.

This module provides a unified interface for using various ODE solvers from the torchdiffeq library for flow matching inference. The solvers integrate the velocity field from t=0 (noise) to t=1 (solution).

Available solvers: - Explicit methods: euler, midpoint, rk4, dopri5 (adaptive) - Implicit methods: dopri8, adaptive_heun - Specialized: bosh3, tsit5

Refactored to inherit from flowpde.core.base_solver.ODESolver

VelocityField

Bases: Module

Wraps a flow matching model as an ODE velocity field.

The flow matching model predicts \(v(x_t, ext{condition}, t)\), which is the time derivative \(dx/dt\) at time \(t\). This wrapper makes it compatible with torchdiffeq's interface.

Source code in flowpde/solvers/ode_solvers.py
class VelocityField(nn.Module):
    """
    Wraps a flow matching model as an ODE velocity field.

    The flow matching model predicts $v(x_t, \text{condition}, t)$, which is the
    time derivative $dx/dt$ at time $t$. This wrapper makes it compatible
    with torchdiffeq's interface.
    """

    def __init__(self, model: nn.Module, condition: Tensor):
        """
        Args:
            model: Flow matching model with signature model(x, condition, t)
            condition: Conditioning tensor (batch_size, dim) - fixed during integration
        """
        super().__init__()
        self.model = model
        self.condition = condition

    def forward(self, t: Tensor, x: Tensor) -> Tensor:
        """
        Compute velocity field at time $t$.

        Args:
            t: Current time (scalar tensor)
            x: Current state (batch_size, dim)

        Returns:
            Velocity $dx/dt$ at time $t$
        """
        # torchdiffeq passes scalar t, but model expects (batch_size, 1)
        batch_size = x.shape[0]
        t_batch = t.expand(batch_size, 1)

        # Compute velocity. The caller owns the grad context: sample()
        # sets it, and training needs the graph kept.
        return self.model(x, self.condition, t_batch)
__init__(model, condition)

Parameters:

Name Type Description Default
model Module

Flow matching model with signature model(x, condition, t)

required
condition Tensor

Conditioning tensor (batch_size, dim) - fixed during integration

required
Source code in flowpde/solvers/ode_solvers.py
def __init__(self, model: nn.Module, condition: Tensor):
    """
    Args:
        model: Flow matching model with signature model(x, condition, t)
        condition: Conditioning tensor (batch_size, dim) - fixed during integration
    """
    super().__init__()
    self.model = model
    self.condition = condition
forward(t, x)

Compute velocity field at time \(t\).

Parameters:

Name Type Description Default
t Tensor

Current time (scalar tensor)

required
x Tensor

Current state (batch_size, dim)

required

Returns:

Type Description
Tensor

Velocity \(dx/dt\) at time \(t\)

Source code in flowpde/solvers/ode_solvers.py
def forward(self, t: Tensor, x: Tensor) -> Tensor:
    """
    Compute velocity field at time $t$.

    Args:
        t: Current time (scalar tensor)
        x: Current state (batch_size, dim)

    Returns:
        Velocity $dx/dt$ at time $t$
    """
    # torchdiffeq passes scalar t, but model expects (batch_size, 1)
    batch_size = x.shape[0]
    t_batch = t.expand(batch_size, 1)

    # Compute velocity. The caller owns the grad context: sample()
    # sets it, and training needs the graph kept.
    return self.model(x, self.condition, t_batch)

ODEFlowSolver

Bases: ODESolver

ODE solver for flow matching inference using torchdiffeq.

This class provides a high-level interface for sampling from flow matching models using various ODE solvers from torchdiffeq.

Inherits from flowpde.core.base_solver.ODESolver

Source code in flowpde/solvers/ode_solvers.py
class ODEFlowSolver(ODESolver):
    """
    ODE solver for flow matching inference using torchdiffeq.

    This class provides a high-level interface for sampling from flow
    matching models using various ODE solvers from torchdiffeq.

    Inherits from flowpde.core.base_solver.ODESolver
    """

    # Available solvers and their properties
    ADAPTIVE_SOLVERS = ['dopri5', 'dopri8', 'bosh3', 'adaptive_heun', 'tsit5']
    FIXED_STEP_SOLVERS = ['euler', 'midpoint', 'rk4', 'explicit_adams', 'implicit_adams']

    def __init__(
        self,
        model: nn.Module,
        method: str = 'dopri5',
        rtol: float = 1e-5,
        atol: float = 1e-7,
        adjoint: bool = False,
        method_options: Optional[dict] = None
    ):
        """
        Initialize ODE solver.

        Args:
            model: Flow matching model with signature model(x, condition, t)
            method: ODE solver method. Options:
                - 'dopri5': Runge-Kutta 4(5) (adaptive, recommended)
                - 'dopri8': Runge-Kutta 7(8) (high accuracy)
                - 'bosh3': Bogacki-Shampine 2(3) (faster, less accurate)
                - 'tsit5': Tsitouras 5(4) (good balance)
                - 'euler': Explicit Euler (simple, fast, less accurate)
                - 'midpoint': Explicit midpoint
                - 'rk4': Classic 4th order Runge-Kutta
                - 'adaptive_heun': Adaptive Heun's method
            rtol: Relative tolerance (for adaptive solvers)
            atol: Absolute tolerance (for adaptive solvers)
            adjoint: If True, use adjoint method for memory-efficient backprop
            method_options: Additional options for the solver
        """
        super().__init__(
            method=method,
            rtol=rtol,
            atol=atol,
            method_options=method_options
        )

        self.model = model
        self.adjoint = adjoint

        # Select integration function
        self.odeint_fn = odeint_adjoint if adjoint else odeint

        # Validate method
        all_methods = self.ADAPTIVE_SOLVERS + self.FIXED_STEP_SOLVERS
        if method not in all_methods:
            raise ValueError(
                f"Unknown method '{method}'. Available: {all_methods}"
            )

    @property
    def is_adaptive(self) -> bool:
        """Whether this is an adaptive step-size solver."""
        return self.method in self.ADAPTIVE_SOLVERS

    @property
    def supports_adjoint(self) -> bool:
        """Whether this solver supports adjoint method for backprop."""
        return True

    def solve(
        self,
        func: Callable[[Tensor, Tensor], Tensor],
        y0: Tensor,
        t_span: Tuple[float, float],
        **kwargs: Any
    ) -> Tensor:
        """
        Solve the ODE $dy/dt = f(t, y)$.

        Args:
            func: Function computing $dy/dt$ given $(t, y)$
            y0: Initial state (batch_size, dim)
            t_span: Time interval $(t_{\text{start}}, t_{\text{end}})$
            **kwargs: Additional solving parameters

        Returns:
            y_final: Final state at t_end (batch_size, dim)
        """
        device = y0.device
        t_eval = torch.tensor(list(t_span), device=device)

        # Build options. 'dtype' is an adaptive-solver option; fixed-step
        # solvers reject it with a warning on every call.
        options = {} if self.method in self.FIXED_STEP_SOLVERS else {'dtype': torch.float32}
        options.update(self.method_options)
        options.update(kwargs)

        # Solve
        trajectory = self.odeint_fn(
            func,
            y0,
            t_eval,
            method=self.method,
            rtol=self.rtol,
            atol=self.atol,
            options=options
        )

        return trajectory[-1]  # Return final state

    def solve_trajectory(
        self,
        func: Callable[[Tensor, Tensor], Tensor],
        y0: Tensor,
        t_eval: Tensor,
        **kwargs: Any
    ) -> Tensor:
        """
        Solve and return trajectory at specified time points.

        Args:
            func: Function computing dy/dt
            y0: Initial state (batch_size, dim)
            t_eval: Time points for evaluation (n_steps,)
            **kwargs: Additional parameters

        Returns:
            trajectory: States at each time point (n_steps, batch_size, dim)
        """
        # Build options. 'dtype' is an adaptive-solver option; fixed-step
        # solvers reject it with a warning on every call.
        if self.method in self.FIXED_STEP_SOLVERS:
            options = {}
            n_steps = len(t_eval) - 1
            dt = (t_eval[-1] - t_eval[0]) / n_steps
            # torchdiffeq infers the direction from t_eval and wants a
            # positive magnitude; a signed dt makes backward integration
            # (data -> latent, t: 1 -> 0) fail outright.
            options['step_size'] = abs(dt.item())
        else:
            options = {'dtype': torch.float32}

        options.update(self.method_options)
        options.update(kwargs)

        # Solve
        trajectory = self.odeint_fn(
            func,
            y0,
            t_eval,
            method=self.method,
            rtol=self.rtol,
            atol=self.atol,
            options=options
        )

        return trajectory

    def sample(
        self,
        condition: Tensor,
        x_init: Optional[Tensor] = None,
        t_span: Tuple[float, float] = (0.0, 1.0),
        return_trajectory: bool = False,
        n_steps: Optional[int] = None,
        no_grad: bool = True,
    ) -> Union[Tensor, Tuple[Tensor, Tensor]]:
        """
        Sample from flow matching model by solving ODE.

        Args:
            condition: Conditioning tensor (batch_size, dim) or (batch_size, H, W)
            x_init: Initial noise (batch_size, dim). If None, sample from N(0, I)
            t_span: Time interval (t_start, t_end), typically (0.0, 1.0)
            return_trajectory: If True, return full trajectory
            n_steps: Number of evaluation points (for fixed-step solvers or trajectory)
                     If None, adaptive solvers choose steps automatically
            no_grad: Integrate under `torch.no_grad()` (default).  Pass False
                     to keep the autograd graph, which is what differentiating
                     through sampling and the adjoint solver both require.

        Note:
            The model is switched to eval mode for the duration of the call and
            restored afterwards, so sampling mid-training does not silently
            disable dropout or freeze BatchNorm statistics.

        Returns:
            samples: Final samples at t_end (batch_size, dim) or (batch_size, H, W)
            trajectory: (optional) Full trajectory if return_trajectory=True
        """
        was_training = self.model.training
        self.model.eval()
        try:
            with torch.no_grad() if no_grad else nullcontext():
                return self._sample(
                    condition, x_init, t_span, return_trajectory, n_steps
                )
        finally:
            # Eval mode is scoped to this call, never a permanent change to the
            # caller's model.
            self.model.train(was_training)

    def _sample(
        self,
        condition: Tensor,
        x_init: Optional[Tensor],
        t_span: Tuple[float, float],
        return_trajectory: bool,
        n_steps: Optional[int],
    ) -> Union[Tensor, Tuple[Tensor, Tensor]]:
        """Body of `sample()`; assumes mode and grad context are already set."""
        # Store original shapes for reshaping
        condition_original_shape = condition.shape
        if condition.dim() > 2:
            condition = condition.flatten(start_dim=1)

        condition = condition.float()
        batch_size = condition.shape[0]
        device = condition.device

        # Initialize from noise and store its shape
        if x_init is None:
            # Default: same spatial shape as condition but need to infer solution channels
            # For now, assume solution has 1 channel (common for PDEs)
            if len(condition_original_shape) > 2:
                # Spatial input: (B, C, H, W) -> solution is (B, 1, H, W)
                solution_shape = (batch_size, 1, *condition_original_shape[2:])
                x_init = torch.randn(solution_shape, device=device)
                x_init_original_shape = solution_shape
            else:
                # Flattened input: same dim as condition
                x_init = torch.randn_like(condition)
                x_init_original_shape = condition_original_shape
        else:
            x_init_original_shape = x_init.shape

        # Flatten x_init for ODE integration
        if x_init.dim() > 2:
            x_init = x_init.flatten(start_dim=1).float()
        else:
            x_init = x_init.float()

        # Create time span
        t_start, t_end = t_span
        if self.method in self.FIXED_STEP_SOLVERS and n_steps is None and not return_trajectory:
            raise ValueError(
                f"n_steps is required for the fixed-step solver "
                f"'{self.method}'. Without it the integration silently "
                f"collapses to a single step across the whole interval."
            )
        if return_trajectory or (self.method in self.FIXED_STEP_SOLVERS and n_steps is not None):
            # Fixed evaluation points
            if n_steps is None:
                n_steps = 50  # Default
            t_eval = torch.linspace(t_start, t_end, n_steps + 1, device=device)
        else:
            # For adaptive solvers, just specify endpoints
            t_eval = torch.tensor([t_start, t_end], device=device)

        # Create velocity field
        velocity_field = VelocityField(self.model, condition)

        # Solve ODE
        trajectory = self.solve_trajectory(
            velocity_field,
            x_init,
            t_eval
        )

        # trajectory shape: (n_steps+1, batch_size, dim) or (2, batch_size, dim)
        samples = trajectory[-1]  # Final state

        # Reshape to original solution shape (not condition shape!)
        if len(x_init_original_shape) > 2:
            samples = samples.view(x_init_original_shape)
            if return_trajectory:
                trajectory = trajectory.view(-1, *x_init_original_shape)

        if return_trajectory:
            return samples, trajectory

        return samples

    def get_solver_info(self) -> Dict[str, Any]:
        """Get information about current solver configuration."""
        info = super().get_solver_info()
        info['adjoint'] = self.adjoint
        return info
is_adaptive property

Whether this is an adaptive step-size solver.

supports_adjoint property

Whether this solver supports adjoint method for backprop.

__init__(model, method='dopri5', rtol=1e-05, atol=1e-07, adjoint=False, method_options=None)

Initialize ODE solver.

Parameters:

Name Type Description Default
model Module

Flow matching model with signature model(x, condition, t)

required
method str

ODE solver method. Options: - 'dopri5': Runge-Kutta 4(5) (adaptive, recommended) - 'dopri8': Runge-Kutta 7(8) (high accuracy) - 'bosh3': Bogacki-Shampine 2(3) (faster, less accurate) - 'tsit5': Tsitouras 5(4) (good balance) - 'euler': Explicit Euler (simple, fast, less accurate) - 'midpoint': Explicit midpoint - 'rk4': Classic 4th order Runge-Kutta - 'adaptive_heun': Adaptive Heun's method

'dopri5'
rtol float

Relative tolerance (for adaptive solvers)

1e-05
atol float

Absolute tolerance (for adaptive solvers)

1e-07
adjoint bool

If True, use adjoint method for memory-efficient backprop

False
method_options Optional[dict]

Additional options for the solver

None
Source code in flowpde/solvers/ode_solvers.py
def __init__(
    self,
    model: nn.Module,
    method: str = 'dopri5',
    rtol: float = 1e-5,
    atol: float = 1e-7,
    adjoint: bool = False,
    method_options: Optional[dict] = None
):
    """
    Initialize ODE solver.

    Args:
        model: Flow matching model with signature model(x, condition, t)
        method: ODE solver method. Options:
            - 'dopri5': Runge-Kutta 4(5) (adaptive, recommended)
            - 'dopri8': Runge-Kutta 7(8) (high accuracy)
            - 'bosh3': Bogacki-Shampine 2(3) (faster, less accurate)
            - 'tsit5': Tsitouras 5(4) (good balance)
            - 'euler': Explicit Euler (simple, fast, less accurate)
            - 'midpoint': Explicit midpoint
            - 'rk4': Classic 4th order Runge-Kutta
            - 'adaptive_heun': Adaptive Heun's method
        rtol: Relative tolerance (for adaptive solvers)
        atol: Absolute tolerance (for adaptive solvers)
        adjoint: If True, use adjoint method for memory-efficient backprop
        method_options: Additional options for the solver
    """
    super().__init__(
        method=method,
        rtol=rtol,
        atol=atol,
        method_options=method_options
    )

    self.model = model
    self.adjoint = adjoint

    # Select integration function
    self.odeint_fn = odeint_adjoint if adjoint else odeint

    # Validate method
    all_methods = self.ADAPTIVE_SOLVERS + self.FIXED_STEP_SOLVERS
    if method not in all_methods:
        raise ValueError(
            f"Unknown method '{method}'. Available: {all_methods}"
        )
solve(func, y0, t_span, **kwargs)

Solve the ODE \(dy/dt = f(t, y)\).

Parameters:

Name Type Description Default
func Callable[[Tensor, Tensor], Tensor]

Function computing \(dy/dt\) given \((t, y)\)

required
y0 Tensor

Initial state (batch_size, dim)

required
t_span Tuple[float, float]

Time interval \((t_{ ext{start}}, t_{ ext{end}})\)

required
**kwargs Any

Additional solving parameters

{}

Returns:

Name Type Description
y_final Tensor

Final state at t_end (batch_size, dim)

Source code in flowpde/solvers/ode_solvers.py
def solve(
    self,
    func: Callable[[Tensor, Tensor], Tensor],
    y0: Tensor,
    t_span: Tuple[float, float],
    **kwargs: Any
) -> Tensor:
    """
    Solve the ODE $dy/dt = f(t, y)$.

    Args:
        func: Function computing $dy/dt$ given $(t, y)$
        y0: Initial state (batch_size, dim)
        t_span: Time interval $(t_{\text{start}}, t_{\text{end}})$
        **kwargs: Additional solving parameters

    Returns:
        y_final: Final state at t_end (batch_size, dim)
    """
    device = y0.device
    t_eval = torch.tensor(list(t_span), device=device)

    # Build options. 'dtype' is an adaptive-solver option; fixed-step
    # solvers reject it with a warning on every call.
    options = {} if self.method in self.FIXED_STEP_SOLVERS else {'dtype': torch.float32}
    options.update(self.method_options)
    options.update(kwargs)

    # Solve
    trajectory = self.odeint_fn(
        func,
        y0,
        t_eval,
        method=self.method,
        rtol=self.rtol,
        atol=self.atol,
        options=options
    )

    return trajectory[-1]  # Return final state
solve_trajectory(func, y0, t_eval, **kwargs)

Solve and return trajectory at specified time points.

Parameters:

Name Type Description Default
func Callable[[Tensor, Tensor], Tensor]

Function computing dy/dt

required
y0 Tensor

Initial state (batch_size, dim)

required
t_eval Tensor

Time points for evaluation (n_steps,)

required
**kwargs Any

Additional parameters

{}

Returns:

Name Type Description
trajectory Tensor

States at each time point (n_steps, batch_size, dim)

Source code in flowpde/solvers/ode_solvers.py
def solve_trajectory(
    self,
    func: Callable[[Tensor, Tensor], Tensor],
    y0: Tensor,
    t_eval: Tensor,
    **kwargs: Any
) -> Tensor:
    """
    Solve and return trajectory at specified time points.

    Args:
        func: Function computing dy/dt
        y0: Initial state (batch_size, dim)
        t_eval: Time points for evaluation (n_steps,)
        **kwargs: Additional parameters

    Returns:
        trajectory: States at each time point (n_steps, batch_size, dim)
    """
    # Build options. 'dtype' is an adaptive-solver option; fixed-step
    # solvers reject it with a warning on every call.
    if self.method in self.FIXED_STEP_SOLVERS:
        options = {}
        n_steps = len(t_eval) - 1
        dt = (t_eval[-1] - t_eval[0]) / n_steps
        # torchdiffeq infers the direction from t_eval and wants a
        # positive magnitude; a signed dt makes backward integration
        # (data -> latent, t: 1 -> 0) fail outright.
        options['step_size'] = abs(dt.item())
    else:
        options = {'dtype': torch.float32}

    options.update(self.method_options)
    options.update(kwargs)

    # Solve
    trajectory = self.odeint_fn(
        func,
        y0,
        t_eval,
        method=self.method,
        rtol=self.rtol,
        atol=self.atol,
        options=options
    )

    return trajectory
sample(condition, x_init=None, t_span=(0.0, 1.0), return_trajectory=False, n_steps=None, no_grad=True)

Sample from flow matching model by solving ODE.

Parameters:

Name Type Description Default
condition Tensor

Conditioning tensor (batch_size, dim) or (batch_size, H, W)

required
x_init Optional[Tensor]

Initial noise (batch_size, dim). If None, sample from N(0, I)

None
t_span Tuple[float, float]

Time interval (t_start, t_end), typically (0.0, 1.0)

(0.0, 1.0)
return_trajectory bool

If True, return full trajectory

False
n_steps Optional[int]

Number of evaluation points (for fixed-step solvers or trajectory) If None, adaptive solvers choose steps automatically

None
no_grad bool

Integrate under torch.no_grad() (default). Pass False to keep the autograd graph, which is what differentiating through sampling and the adjoint solver both require.

True
Note

The model is switched to eval mode for the duration of the call and restored afterwards, so sampling mid-training does not silently disable dropout or freeze BatchNorm statistics.

Returns:

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

Final samples at t_end (batch_size, dim) or (batch_size, H, W)

trajectory Union[Tensor, Tuple[Tensor, Tensor]]

(optional) Full trajectory if return_trajectory=True

Source code in flowpde/solvers/ode_solvers.py
def sample(
    self,
    condition: Tensor,
    x_init: Optional[Tensor] = None,
    t_span: Tuple[float, float] = (0.0, 1.0),
    return_trajectory: bool = False,
    n_steps: Optional[int] = None,
    no_grad: bool = True,
) -> Union[Tensor, Tuple[Tensor, Tensor]]:
    """
    Sample from flow matching model by solving ODE.

    Args:
        condition: Conditioning tensor (batch_size, dim) or (batch_size, H, W)
        x_init: Initial noise (batch_size, dim). If None, sample from N(0, I)
        t_span: Time interval (t_start, t_end), typically (0.0, 1.0)
        return_trajectory: If True, return full trajectory
        n_steps: Number of evaluation points (for fixed-step solvers or trajectory)
                 If None, adaptive solvers choose steps automatically
        no_grad: Integrate under `torch.no_grad()` (default).  Pass False
                 to keep the autograd graph, which is what differentiating
                 through sampling and the adjoint solver both require.

    Note:
        The model is switched to eval mode for the duration of the call and
        restored afterwards, so sampling mid-training does not silently
        disable dropout or freeze BatchNorm statistics.

    Returns:
        samples: Final samples at t_end (batch_size, dim) or (batch_size, H, W)
        trajectory: (optional) Full trajectory if return_trajectory=True
    """
    was_training = self.model.training
    self.model.eval()
    try:
        with torch.no_grad() if no_grad else nullcontext():
            return self._sample(
                condition, x_init, t_span, return_trajectory, n_steps
            )
    finally:
        # Eval mode is scoped to this call, never a permanent change to the
        # caller's model.
        self.model.train(was_training)
get_solver_info()

Get information about current solver configuration.

Source code in flowpde/solvers/ode_solvers.py
def get_solver_info(self) -> Dict[str, Any]:
    """Get information about current solver configuration."""
    info = super().get_solver_info()
    info['adjoint'] = self.adjoint
    return info

sample_with_ode_solver(model, condition, solver='dopri5', rtol=1e-05, atol=1e-07, n_steps=None, device=None, return_trajectory=False)

Convenience function for sampling with ODE solver.

This is a simpler interface to ODEFlowSolver for one-off sampling.

Parameters:

Name Type Description Default
model Module

Flow matching model

required
condition Tensor

Conditioning tensor

required
solver str

ODE solver method (see ODEFlowSolver for options)

'dopri5'
rtol float

Relative tolerance

1e-05
atol float

Absolute tolerance

1e-07
n_steps Optional[int]

Number of steps (for fixed-step solvers)

None
device Optional[str]

Device for computation. Defaults to CUDA when available, CPU otherwise.

None
return_trajectory bool

If True, return full trajectory

False

Returns:

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

Final samples

trajectory Union[Tensor, Tuple[Tensor, Tensor]]

(optional) Full trajectory if return_trajectory=True

Example

samples = sample_with_ode_solver( ... model=trained_model, ... condition=f, ... solver='dopri5', ... device='cuda' ... )

Source code in flowpde/solvers/ode_solvers.py
@torch.no_grad()
def sample_with_ode_solver(
    model: nn.Module,
    condition: Tensor,
    solver: str = 'dopri5',
    rtol: float = 1e-5,
    atol: float = 1e-7,
    n_steps: Optional[int] = None,
    device: Optional[str] = None,
    return_trajectory: bool = False
) -> Union[Tensor, Tuple[Tensor, Tensor]]:
    """
    Convenience function for sampling with ODE solver.

    This is a simpler interface to ODEFlowSolver for one-off sampling.

    Args:
        model: Flow matching model
        condition: Conditioning tensor
        solver: ODE solver method (see ODEFlowSolver for options)
        rtol: Relative tolerance
        atol: Absolute tolerance
        n_steps: Number of steps (for fixed-step solvers)
        device: Device for computation.  Defaults to CUDA when available,
            CPU otherwise.
        return_trajectory: If True, return full trajectory

    Returns:
        samples: Final samples
        trajectory: (optional) Full trajectory if return_trajectory=True

    Example:
        >>> samples = sample_with_ode_solver(
        ...     model=trained_model,
        ...     condition=f,
        ...     solver='dopri5',
        ...     device='cuda'
        ... )
    """
    device = resolve_device(device)
    condition = condition.to(device)

    solver_instance = ODEFlowSolver(
        model=model,
        method=solver,
        rtol=rtol,
        atol=atol
    )

    return solver_instance.sample(
        condition=condition,
        return_trajectory=return_trajectory,
        n_steps=n_steps
    )

compare_solvers(model, condition, ground_truth=None, solvers=None, device=None, n_steps=50)

Compare different ODE solvers on the same input.

Parameters:

Name Type Description Default
model Module

Flow matching model

required
condition Tensor

Conditioning tensor

required
ground_truth Optional[Tensor]

Optional ground truth for error computation

None
solvers Optional[List[str]]

List of solver names to compare. If None, uses default set

None
device Optional[str]

Device for computation. Defaults to CUDA when available, CPU otherwise.

None
n_steps int

Number of steps for fixed-step solvers

50

Returns:

Type Description
dict

Dictionary with results for each solver including:

dict
  • samples: Generated samples
dict
  • time: Computation time
dict
  • error: L2 error vs ground truth (if provided)
Example

results = compare_solvers( ... model=trained_model, ... condition=f, ... ground_truth=u_true, ... solvers=['euler', 'rk4', 'dopri5'] ... ) for solver, info in results.items(): ... print(f"{solver}: error={info['error']:.6f}, time={info['time']:.3f}s")

Source code in flowpde/solvers/ode_solvers.py
def compare_solvers(
    model: nn.Module,
    condition: Tensor,
    ground_truth: Optional[Tensor] = None,
    solvers: Optional[List[str]] = None,
    device: Optional[str] = None,
    n_steps: int = 50
) -> dict:
    """
    Compare different ODE solvers on the same input.

    Args:
        model: Flow matching model
        condition: Conditioning tensor
        ground_truth: Optional ground truth for error computation
        solvers: List of solver names to compare. If None, uses default set
        device: Device for computation.  Defaults to CUDA when available,
            CPU otherwise.
        n_steps: Number of steps for fixed-step solvers

    Returns:
        Dictionary with results for each solver including:
        - samples: Generated samples
        - time: Computation time
        - error: L2 error vs ground truth (if provided)

    Example:
        >>> results = compare_solvers(
        ...     model=trained_model,
        ...     condition=f,
        ...     ground_truth=u_true,
        ...     solvers=['euler', 'rk4', 'dopri5']
        ... )
        >>> for solver, info in results.items():
        ...     print(f"{solver}: error={info['error']:.6f}, time={info['time']:.3f}s")
    """
    import time

    if solvers is None:
        solvers = ['euler', 'midpoint', 'rk4', 'dopri5']

    device = resolve_device(device)
    model.eval()
    condition = condition.to(device)
    if ground_truth is not None:
        ground_truth = ground_truth.to(device)

    results = {}

    for solver_name in solvers:
        # Create solver
        ode_solver = ODEFlowSolver(
            model=model,
            method=solver_name,
            rtol=1e-5,
            atol=1e-7
        )

        # Time the sampling
        if device == 'cuda':
            torch.cuda.synchronize()
        start_time = time.time()

        samples = ode_solver.sample(
            condition=condition,
            n_steps=n_steps if solver_name in ODEFlowSolver.FIXED_STEP_SOLVERS else None
        )

        if device == 'cuda':
            torch.cuda.synchronize()
        elapsed = time.time() - start_time

        # Compute error if ground truth provided
        error = None
        if ground_truth is not None:
            samples_flat = samples.flatten(start_dim=1)
            gt_flat = ground_truth.flatten(start_dim=1)
            error = (samples_flat - gt_flat).norm(dim=1).mean().item()

        results[solver_name] = {
            'samples': samples.cpu(),
            'time': elapsed,
            'error': error,
            'solver_info': ode_solver.get_solver_info()
        }

    return results