catlog/simulate/ode/
mod.rs1use nalgebra::DVector;
4use ode_solvers::{
5 self,
6 dop_shared::{IntegrationError, SolverResult},
7};
8
9#[cfg(test)]
10use textplots::{Chart, Plot, Shape};
11
12pub trait ODESystem {
16 fn vector_field(&self, dx: &mut DVector<f32>, x: &DVector<f32>, t: f32);
18
19 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#[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 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 rtol: 0.001,
51 atol: 1e-6,
52 }
53 }
54
55 pub fn start_time(mut self, t: f32) -> Self {
57 self.start_time = t;
58 self
59 }
60
61 pub fn end_time(mut self, t: f32) -> Self {
63 self.end_time = t;
64 self
65 }
66
67 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 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 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::*;