catlog/simulate/ode/
mod.rs

1//! Simulation of dynamical systems defined by ODEs.
2
3use nalgebra::DVector;
4use ode_solvers::{
5    self,
6    dop_shared::{IntegrationError, SolverResult},
7};
8
9#[cfg(test)]
10use textplots::{Chart, Plot, Shape};
11
12/// A system of ordinary differential equations (ODEs).
13///
14/// An ODE system is anything that can compute a vector field.
15pub trait ODESystem {
16    /// Compute the vector field at the given time and state in place.
17    fn vector_field(&self, dx: &mut DVector<f32>, x: &DVector<f32>, t: f32);
18
19    /// Compute and return the vector field at the given time and state.
20    fn eval_vector_field(&self, x: &DVector<f32>, t: f32) -> DVector<f32> {
21        let mut dx = DVector::from_element(x.len(), 0.0f32);
22        self.vector_field(&mut dx, x, t);
23        dx
24    }
25}
26
27/// An ODE problem ready to be solved.
28///
29/// An ODE problem comprises an [ODE system](ODESystem) plus the extra information
30/// needed to solve the system, namely the initial values and the time span.
31#[derive(Clone, Debug, PartialEq)]
32pub struct ODEProblem<Sys> {
33    pub(crate) system: Sys,
34    pub(crate) initial_values: DVector<f32>,
35    pub(crate) start_time: f32,
36    pub(crate) end_time: f32,
37    rtol: f32,
38    atol: f32,
39}
40
41impl<Sys> ODEProblem<Sys> {
42    /// Creates a new ODE problem.
43    pub fn new(system: Sys, initial_values: DVector<f32>) -> Self {
44        ODEProblem {
45            system,
46            initial_values,
47            start_time: 0.0,
48            end_time: 0.0,
49            // Same defaults as `scipy.integrate.RK45`.
50            rtol: 0.001,
51            atol: 1e-6,
52        }
53    }
54
55    /// Sets the start time for the problem.
56    pub fn start_time(mut self, t: f32) -> Self {
57        self.start_time = t;
58        self
59    }
60
61    /// Sets the end time for the problem.
62    pub fn end_time(mut self, t: f32) -> Self {
63        self.end_time = t;
64        self
65    }
66
67    /// Sets the time span (start and end time) for the problem.
68    pub fn time_span(mut self, tspan: (f32, f32)) -> Self {
69        (self.start_time, self.end_time) = tspan;
70        self
71    }
72}
73
74impl<Sys> ODEProblem<Sys>
75where
76    Sys: ODESystem,
77{
78    /// Solves the ODE system using the Runge-Kutta method.
79    ///
80    /// Returns the solver results if successful and an integration error otherwise.
81    pub fn solve_rk4(
82        &self,
83        step_size: f32,
84    ) -> Result<SolverResult<f32, DVector<f32>>, IntegrationError> {
85        let mut stepper = ode_solvers::Rk4::new(
86            self,
87            self.start_time,
88            self.initial_values.clone(),
89            self.end_time,
90            step_size,
91        );
92        stepper.integrate()?;
93        Ok(stepper.into())
94    }
95
96    /// Solves the ODE system using the Dormand-Prince method.
97    ///
98    /// A variant of Runge-Kutta with adaptive step size control and automatic
99    /// selection of initial step size.
100    pub fn solve_dopri5(
101        &self,
102        output_step_size: f32,
103    ) -> Result<SolverResult<f32, DVector<f32>>, IntegrationError> {
104        let mut stepper = ode_solvers::Dopri5::new(
105            self,
106            self.start_time,
107            self.end_time,
108            output_step_size,
109            self.initial_values.clone(),
110            self.rtol,
111            self.atol,
112        );
113        stepper.integrate()?;
114        Ok(stepper.into())
115    }
116}
117
118impl<Sys> ode_solvers::dop_shared::System<f32, DVector<f32>> for &ODEProblem<Sys>
119where
120    Sys: ODESystem,
121{
122    fn system(&self, x: f32, y: &DVector<f32>, dy: &mut DVector<f32>) {
123        self.system.vector_field(dy, y, x);
124    }
125}
126
127#[cfg(test)]
128pub(crate) fn textplot_ode_result<Sys>(
129    problem: &ODEProblem<Sys>,
130    result: &SolverResult<f32, DVector<f32>>,
131) -> String {
132    textplot_mapped_ode_result(problem, result, |_| true, |x, i| x[i])
133}
134
135#[cfg(test)]
136pub(crate) fn textplot_mapped_ode_result<Sys>(
137    problem: &ODEProblem<Sys>,
138    result: &SolverResult<f32, DVector<f32>>,
139    include_var: impl Fn(usize) -> bool,
140    f: impl Fn(&DVector<f32>, usize) -> f32,
141) -> String {
142    let mut chart = Chart::new(100, 80, 0.0, problem.end_time);
143    let (t_out, x_out) = result.get();
144
145    let dim = problem.initial_values.len();
146    let line_data: Vec<_> = (0..dim)
147        .filter(|i| include_var(*i))
148        .map(|i| {
149            std::iter::zip(t_out.iter().copied(), x_out.iter().map(|x| f(x, i))).collect::<Vec<_>>()
150        })
151        .collect();
152
153    let lines: Vec<_> = line_data.iter().map(|data| Shape::Lines(data)).collect();
154    let chart = lines.iter().fold(&mut chart, |chart, line| chart.lineplot(line));
155    chart.axis();
156    chart.figures();
157    chart.to_string()
158}
159
160pub mod kuramoto;
161pub mod polynomial;
162
163pub use kuramoto::*;
164pub use polynomial::*;