1use std::fmt::Display;
4use std::hash::Hash;
5use std::ops::{Add, Neg};
6
7use derivative::Derivative;
8use indexmap::IndexMap;
9use nalgebra::DVector;
10use num_traits::{One, Pow, Zero};
11
12#[cfg(feature = "serde")]
13use serde::{Deserialize, Serialize};
14#[cfg(feature = "serde-wasm")]
15use tsify::Tsify;
16
17#[cfg(test)]
18use super::ODEProblem;
19use super::ODESystem;
20use crate::zero::{alg::Polynomial, rig::DisplayCoef};
21
22#[derive(Clone, Derivative)]
24#[derivative(Default(bound = ""))]
25pub struct PolynomialSystem<Var, Coef, Exp> {
26 pub components: IndexMap<Var, Polynomial<Var, Coef, Exp>>,
28}
29
30impl<Var, Coef, Exp> PolynomialSystem<Var, Coef, Exp>
31where
32 Var: Hash + Ord,
33 Exp: Ord,
34{
35 pub fn new() -> Self {
37 Default::default()
38 }
39
40 pub fn add_term(&mut self, var: Var, term: Polynomial<Var, Coef, Exp>)
42 where
43 Coef: Add<Output = Coef>,
44 {
45 if let Some(component) = self.components.get_mut(&var) {
46 *component = std::mem::take(component) + term;
47 } else {
48 self.components.insert(var, term);
49 }
50 }
51
52 pub fn extend_scalars<NewCoef, F>(self, f: F) -> PolynomialSystem<Var, NewCoef, Exp>
54 where
55 F: Clone + FnMut(Coef) -> NewCoef,
56 {
57 self.map(|poly| poly.extend_scalars(f.clone()))
58 }
59
60 pub fn map_variables<NewVar, F>(&self, mut f: F) -> PolynomialSystem<NewVar, Coef, Exp>
62 where
63 NewVar: Clone + Hash + Ord,
64 Coef: Clone + Add<Output = Coef>,
65 Exp: Clone + Add<Output = Exp>,
66 F: FnMut(&Var) -> NewVar,
67 {
68 let components = self
69 .components
70 .iter()
71 .map(|(var, poly)| {
72 let new_var = f(var);
73 (new_var, poly.map_variables(|v| f(v)))
74 })
75 .collect();
76 PolynomialSystem { components }
77 }
78
79 pub fn normalize(self) -> Self
81 where
82 Coef: Zero,
83 Exp: Zero,
84 {
85 self.map(|poly| poly.normalize())
86 }
87
88 pub fn map<NewCoef, NewExp, F>(self, mut f: F) -> PolynomialSystem<Var, NewCoef, NewExp>
90 where
91 F: FnMut(Polynomial<Var, Coef, Exp>) -> Polynomial<Var, NewCoef, NewExp>,
92 {
93 let components = self.components.into_iter().map(|(var, poly)| (var, f(poly))).collect();
94 PolynomialSystem { components }
95 }
96
97 pub fn to_latex_equations(&self) -> Vec<LatexEquation>
99 where
100 Var: Display,
101 Coef: Display + DisplayCoef + Clone + PartialEq + One + Neg<Output = Coef>,
102 Exp: Display + PartialEq + One,
103 {
104 self.components
105 .iter()
106 .map(|(var, poly)| LatexEquation {
107 lhs: format!("\\frac{{\\mathrm{{d}}}}{{\\mathrm{{d}}t}} {var}"),
108 rhs: poly.to_latex(),
109 })
110 .collect()
111 }
112}
113
114#[derive(Debug, PartialEq, Eq)]
115#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
116#[cfg_attr(feature = "serde-wasm", derive(Tsify))]
117#[cfg_attr(feature = "serde-wasm", tsify(into_wasm_abi, from_wasm_abi))]
118pub struct LatexEquation {
120 pub lhs: String,
122 pub rhs: String,
124}
125
126impl<Var, Exp> PolynomialSystem<Var, f32, Exp>
127where
128 Var: Clone + Hash + Ord,
129 Exp: Clone + Ord + Add<Output = Exp>,
130{
131 pub fn to_numerical(&self) -> NumericalPolynomialSystem<Exp> {
136 let indices: IndexMap<Var, usize> =
137 self.components.keys().enumerate().map(|(i, var)| (var.clone(), i)).collect();
138 let components = self
139 .components
140 .values()
141 .map(|poly| poly.map_variables(|var| *indices.get(var).unwrap()))
142 .collect();
143 NumericalPolynomialSystem { components }
144 }
145}
146
147impl<Var, Coef, Exp> Display for PolynomialSystem<Var, Coef, Exp>
148where
149 Var: Display,
150 Polynomial<Var, Coef, Exp>: Display,
151{
152 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
153 for (var, component) in self.components.iter() {
154 writeln!(f, "d{var} = {component}")?;
155 }
156 Ok(())
157 }
158}
159
160impl<Var, Coef, Exp> FromIterator<(Var, Polynomial<Var, Coef, Exp>)>
161 for PolynomialSystem<Var, Coef, Exp>
162where
163 Var: Hash + Ord,
164 Coef: Add<Output = Coef>,
165 Exp: Ord,
166{
167 fn from_iter<T: IntoIterator<Item = (Var, Polynomial<Var, Coef, Exp>)>>(iter: T) -> Self {
168 let mut system: Self = Default::default();
169 for (var, term) in iter {
170 system.add_term(var, term);
171 }
172 system
173 }
174}
175
176pub struct NumericalPolynomialSystem<Exp> {
181 pub components: Vec<Polynomial<usize, f32, Exp>>,
183}
184
185impl<Exp> ODESystem for NumericalPolynomialSystem<Exp>
186where
187 Exp: Clone + Ord,
188 f32: Pow<Exp, Output = f32>,
189{
190 fn vector_field(&self, dx: &mut DVector<f32>, x: &DVector<f32>, _t: f32) {
191 for i in 0..dx.len() {
192 dx[i] = self.components[i].eval(|var| x[*var])
193 }
194 }
195}
196
197impl<Exp> Display for NumericalPolynomialSystem<Exp>
198where
199 Exp: Clone + Ord + Add<Output = Exp> + One + Display,
200{
201 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
202 let var_name = |i: usize| format!("x{i}");
203 for (var, component) in self.components.iter().enumerate() {
204 let var = var_name(var);
205 let component = component.map_variables(|i| var_name(*i));
206 writeln!(f, "d{var} = {component}")?;
207 }
208 Ok(())
209 }
210}
211
212#[cfg(test)]
213mod tests {
214 use expect_test::expect;
215
216 use super::super::textplot_ode_result;
217 use super::*;
218
219 type Parameter<Id> = Polynomial<Id, f32, u8>;
220
221 #[test]
222 fn sir() {
223 let param = |c: char| Parameter::<_>::generator(c);
224 let var = |c: char| Polynomial::<_, Parameter<_>, u8>::generator(c);
225 let terms = [
226 ('S', -var('S') * var('I') * param('β')),
227 ('I', var('S') * var('I') * param('β')),
228 ('I', -var('I') * param('γ')),
229 ('R', var('I') * param('γ')),
230 ];
231 let sys: PolynomialSystem<_, _, _> = terms.into_iter().collect();
232 let expected = expect![[r#"
233 dS = -β I S
234 dI = -γ I + β I S
235 dR = γ I
236 "#]];
237 expected.assert_eq(&sys.to_string());
238
239 let sys = sys.extend_scalars(|p| p.eval(|_| 1.0));
240 let expected = expect![[r#"
241 dS = -I S
242 dI = -I + I S
243 dR = I
244 "#]];
245 expected.assert_eq(&sys.to_string());
246
247 let initial = DVector::from_column_slice(&[4.0, 1.0, 0.0]);
248 let problem = ODEProblem::new(sys.to_numerical(), initial).end_time(5.0);
249 let result = problem.solve_rk4(0.1).unwrap();
250 let expected = expect![[r#"
251 ⡁⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⣀⣀⣀⠤⠤⠤⠒⠒⠒⠒⠒⠉⠉⠉⠉⠁ 4.9
252 ⠄⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⣀⣀⠤⠒⠒⠉⠉⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
253 ⠂⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⣀⠤⠒⠉⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
254 ⡁⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⢀⠤⠊⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
255 ⢇⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⡠⠒⠁⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
256 ⠚⡄⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⡠⠊⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
257 ⡁⢣⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⢀⠎⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
258 ⠄⠘⡄⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⡔⠁⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
259 ⠂⠀⢣⠀⠀⠀⠀⠀⠀⠀⠀⠀⢀⠎⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
260 ⡁⠀⠘⡄⠀⢀⠤⠒⠤⡀⠀⢠⠃⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
261 ⠄⠀⠀⢣⡔⠁⠀⠀⠀⠈⢦⠃⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
262 ⠂⠀⠀⡜⡄⠀⠀⠀⠀⢠⠃⠑⢄⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
263 ⡁⠀⡸⠀⢣⠀⠀⠀⢠⠃⠀⠀⠀⠣⡀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
264 ⠄⢠⠃⠀⠘⡄⠀⢠⠃⠀⠀⠀⠀⠀⠈⠢⡀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
265 ⢂⠇⠀⠀⠀⠱⣠⠃⠀⠀⠀⠀⠀⠀⠀⠀⠈⠢⡀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
266 ⡝⠀⠀⠀⠀⢠⢣⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠈⠢⣀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
267 ⠅⠀⠀⠀⢠⠃⠀⢣⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠑⠤⣀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
268 ⠂⠀⠀⢠⠃⠀⠀⠀⠣⡀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠉⠒⠤⣀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
269 ⡁⠀⡠⠃⠀⠀⠀⠀⠀⠈⠒⠤⣀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠉⠒⠒⠤⠤⣀⣀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀
270 ⢄⠔⠁⠀⠀⠀⠀⠀⠀⠀⠀⠀⠀⠉⠒⠒⠤⠤⠤⠤⣀⣀⣀⣀⣀⣀⣀⣀⣀⣀⣀⣀⣀⣀⣀⣉⣉⣒⣒⣒⣒⣤⣤⣤⣤⠤⣀⣀⣀⣀⡀
271 ⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠁⠈⠀⠉⠉⠉⠉⠉⠁ 0.0
272 0.0 5.0
273 "#]];
274 expected.assert_eq(&textplot_ode_result(&problem, &result));
275 }
276}