Skip to content

Base Solver

Abstract base classes for ODE and SDE solvers.

base_solver

Base Solver Abstract Class

Defines the interface for all ODE/SDE solvers in FlowPDE.

BaseSolver

Bases: ABC

Abstract base class for ODE solvers.

This defines the common interface for numerical integration methods used to solve differential equations in normalizing flows.

Solvers can be: - Fixed-step (e.g., Euler, RK4) - Adaptive step-size (e.g., Dopri5, Dopri8) - Stochastic (SDE solvers) (Future work)

Source code in flowpde/core/base_solver.py
class BaseSolver(ABC):
    """
    Abstract base class for ODE solvers.

    This defines the common interface for numerical integration methods
    used to solve differential equations in normalizing flows.

    Solvers can be:
    - Fixed-step (e.g., Euler, RK4)
    - Adaptive step-size (e.g., Dopri5, Dopri8)
    - Stochastic (SDE solvers) (Future work)
    """

    def __init__(
        self,
        rtol: float = 1e-5,
        atol: float = 1e-7,
        method_options: Optional[Dict[str, Any]] = None,
        **kwargs: Any
    ):
        """
        Initialize solver.

        Args:
            rtol: Relative tolerance for adaptive solvers
            atol: Absolute tolerance for adaptive solvers
            method_options: Additional solver-specific options
            **kwargs: Additional parameters
        """
        self.rtol = rtol
        self.atol = atol
        self.method_options = method_options or {}
        self._extra_kwargs = kwargs

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

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

        Returns:
            $y_{\text{final}}$: Final state at $t_{\text{end}}$ (batch_size, dim)
            info (optional): Dictionary with solving statistics
        """
        raise NotImplementedError

    def solve_trajectory(
        self,
        func: Callable[[Tensor, Tensor], Tensor],
        y0: Tensor,
        t_eval: Tensor,
        **kwargs: Any
    ) -> Union[Tensor, Tuple[Tensor, Dict[str, Any]]]:
        """
        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)
            info (optional): Dictionary with solving statistics
        """
        # Default implementation: solve multiple intervals
        # Subclasses can override for more efficient implementations
        trajectory = [y0]
        y_current = y0

        for i in range(len(t_eval) - 1):
            t_span = (t_eval[i].item(), t_eval[i+1].item())
            result = self.solve(func, y_current, t_span, **kwargs)
            if isinstance(result, tuple):
                y_current, _ = result
            else:
                y_current = result
            trajectory.append(y_current)

        return torch.stack(trajectory, dim=0)

    @abstractmethod
    def get_solver_info(self) -> Dict[str, Any]:
        """
        Get information about the solver configuration.

        Returns:
            Dictionary with solver properties
        """
        raise NotImplementedError

    def set_tolerance(self, rtol: float, atol: float):
        """Update solver tolerances."""
        self.rtol = rtol
        self.atol = atol

    @property
    @abstractmethod
    def is_adaptive(self) -> bool:
        """Whether this is an adaptive step-size solver."""
        raise NotImplementedError

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

    def __repr__(self) -> str:
        return (
            f"{self.__class__.__name__}("
            f"rtol={self.rtol}, atol={self.atol}, "
            f"adaptive={self.is_adaptive})"
        )
is_adaptive abstractmethod property

Whether this is an adaptive step-size solver.

supports_adjoint property

Whether this solver supports adjoint method for backprop.

__init__(rtol=1e-05, atol=1e-07, method_options=None, **kwargs)

Initialize solver.

Parameters:

Name Type Description Default
rtol float

Relative tolerance for adaptive solvers

1e-05
atol float

Absolute tolerance for adaptive solvers

1e-07
method_options Optional[Dict[str, Any]]

Additional solver-specific options

None
**kwargs Any

Additional parameters

{}
Source code in flowpde/core/base_solver.py
def __init__(
    self,
    rtol: float = 1e-5,
    atol: float = 1e-7,
    method_options: Optional[Dict[str, Any]] = None,
    **kwargs: Any
):
    """
    Initialize solver.

    Args:
        rtol: Relative tolerance for adaptive solvers
        atol: Absolute tolerance for adaptive solvers
        method_options: Additional solver-specific options
        **kwargs: Additional parameters
    """
    self.rtol = rtol
    self.atol = atol
    self.method_options = method_options or {}
    self._extra_kwargs = kwargs
solve(func, y0, t_span, **kwargs) abstractmethod

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

Parameters:

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

Function computing \(dy/dt\) given \((t, y)\) Signature: func(t: Tensor, y: Tensor) -> Tensor

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
Union[Tensor, Tuple[Tensor, Dict[str, Any]]]

\(y_{ ext{final}}\): Final state at \(t_{ ext{end}}\) (batch_size, dim)

info optional

Dictionary with solving statistics

Source code in flowpde/core/base_solver.py
@abstractmethod
def solve(
    self,
    func: Callable[[Tensor, Tensor], Tensor],
    y0: Tensor,
    t_span: Tuple[float, float],
    **kwargs: Any
) -> Union[Tensor, Tuple[Tensor, Dict[str, Any]]]:
    """
    Solve the differential equation $dy/dt = f(t, y)$.

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

    Returns:
        $y_{\text{final}}$: Final state at $t_{\text{end}}$ (batch_size, dim)
        info (optional): Dictionary with solving statistics
    """
    raise NotImplementedError
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 Union[Tensor, Tuple[Tensor, Dict[str, Any]]]

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

info optional

Dictionary with solving statistics

Source code in flowpde/core/base_solver.py
def solve_trajectory(
    self,
    func: Callable[[Tensor, Tensor], Tensor],
    y0: Tensor,
    t_eval: Tensor,
    **kwargs: Any
) -> Union[Tensor, Tuple[Tensor, Dict[str, Any]]]:
    """
    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)
        info (optional): Dictionary with solving statistics
    """
    # Default implementation: solve multiple intervals
    # Subclasses can override for more efficient implementations
    trajectory = [y0]
    y_current = y0

    for i in range(len(t_eval) - 1):
        t_span = (t_eval[i].item(), t_eval[i+1].item())
        result = self.solve(func, y_current, t_span, **kwargs)
        if isinstance(result, tuple):
            y_current, _ = result
        else:
            y_current = result
        trajectory.append(y_current)

    return torch.stack(trajectory, dim=0)
get_solver_info() abstractmethod

Get information about the solver configuration.

Returns:

Type Description
Dict[str, Any]

Dictionary with solver properties

Source code in flowpde/core/base_solver.py
@abstractmethod
def get_solver_info(self) -> Dict[str, Any]:
    """
    Get information about the solver configuration.

    Returns:
        Dictionary with solver properties
    """
    raise NotImplementedError
set_tolerance(rtol, atol)

Update solver tolerances.

Source code in flowpde/core/base_solver.py
def set_tolerance(self, rtol: float, atol: float):
    """Update solver tolerances."""
    self.rtol = rtol
    self.atol = atol

ODESolver

Bases: BaseSolver

Base class specifically for ODE solvers.

Ordinary Differential Equation solvers for deterministic flows.

Source code in flowpde/core/base_solver.py
class ODESolver(BaseSolver):
    """
    Base class specifically for ODE solvers.

    Ordinary Differential Equation solvers for deterministic flows.
    """

    def __init__(
        self,
        method: str,
        rtol: float = 1e-5,
        atol: float = 1e-7,
        **kwargs: Any
    ):
        """
        Initialize ODE solver.

        Args:
            method: Integration method name
            rtol: Relative tolerance
            atol: Absolute tolerance
            **kwargs: Additional parameters
        """
        super().__init__(rtol=rtol, atol=atol, **kwargs)
        self.method = method

    def get_solver_info(self) -> Dict[str, Any]:
        """Get ODE solver information."""
        return {
            'type': 'ODE',
            'method': self.method,
            'rtol': self.rtol,
            'atol': self.atol,
            'adaptive': self.is_adaptive,
            'adjoint_support': self.supports_adjoint,
        }
__init__(method, rtol=1e-05, atol=1e-07, **kwargs)

Initialize ODE solver.

Parameters:

Name Type Description Default
method str

Integration method name

required
rtol float

Relative tolerance

1e-05
atol float

Absolute tolerance

1e-07
**kwargs Any

Additional parameters

{}
Source code in flowpde/core/base_solver.py
def __init__(
    self,
    method: str,
    rtol: float = 1e-5,
    atol: float = 1e-7,
    **kwargs: Any
):
    """
    Initialize ODE solver.

    Args:
        method: Integration method name
        rtol: Relative tolerance
        atol: Absolute tolerance
        **kwargs: Additional parameters
    """
    super().__init__(rtol=rtol, atol=atol, **kwargs)
    self.method = method
get_solver_info()

Get ODE solver information.

Source code in flowpde/core/base_solver.py
def get_solver_info(self) -> Dict[str, Any]:
    """Get ODE solver information."""
    return {
        'type': 'ODE',
        'method': self.method,
        'rtol': self.rtol,
        'atol': self.atol,
        'adaptive': self.is_adaptive,
        'adjoint_support': self.supports_adjoint,
    }