use core::fmt::Display;
use crate::Utils::plots::plots;
use crate::numerical::NR_for_Euler::NRE;
use crate::symbolic::symbolic_engine::Expr;
use nalgebra::{DMatrix, DVector};
use log::info;
use std::collections::HashMap;
use std::time::Instant;
pub enum Equation {
LHS(Vec<Expr>),
RHS(Vec<Expr>),
}
pub struct BE {
pub newton: NRE,
y0: DVector<f64>,
t0: f64,
t_bound: f64,
t: f64,
y: DVector<f64>,
t_old: Option<f64>,
t_result: DVector<f64>,
y_result: DMatrix<f64>,
status: String,
message: Option<String>,
h: Option<f64>,
global_timestepping: bool,
stop_condition: Option<HashMap<String, f64>>,
neighborhood_check: f64,
}
impl Display for BE {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"BE {{ t0: {}, t_bound: {}, t: {}, y: {:?} }}",
self.t0, self.t_bound, self.t, self.y
)
}
}
impl BE {
pub fn new() -> BE {
let nr_new = NRE::new(
Vec::new(),
DVector::zeros(0),
Vec::new(),
String::new(),
0.0,
0,
0.0,
true,
None,
);
BE {
newton: nr_new,
y0: DVector::zeros(0),
t0: 0.0,
t_bound: 0.0,
t: 0.0,
y: DVector::zeros(0),
t_old: None,
t_result: DVector::zeros(0),
y_result: DMatrix::zeros(0, 0),
status: "running".to_string(),
message: None,
h: None,
global_timestepping: true,
stop_condition: None,
neighborhood_check: 1e-6,
}
}
pub fn set_initial(
&mut self,
eq_system: Vec<Expr>, values: Vec<String>,
arg: String,
tolerance: f64, max_iterations: usize, h: Option<f64>,
t0: f64,
t_bound: f64,
y0: DVector<f64>,
) -> () {
let initial_guess = y0.clone();
let nr = if let Some(dt) = h {
self.global_timestepping = true;
NRE::new(
eq_system,
initial_guess,
values,
arg,
tolerance,
max_iterations,
dt,
true,
None,
)
} else {
info!("global_timestepping = false");
self.global_timestepping = false;
NRE::new(
eq_system,
initial_guess,
values,
arg,
tolerance,
max_iterations,
1e-4,
false,
Some(t_bound),
)
};
self.newton = nr;
self.t0 = t0;
self.t_bound = t_bound;
self.y0 = y0.clone();
self.t = t0;
self.y = y0.clone();
self.check();
}
pub fn set_stop_condition(&mut self, stop_condition: HashMap<String, f64>) {
self.stop_condition = Some(stop_condition);
}
pub fn set_neighborhood_check(&mut self, tolerance: f64) {
self.neighborhood_check = tolerance;
}
fn check_stop_condition(&self, y: &DVector<f64>) -> bool {
if let Some(ref conditions) = self.stop_condition {
for (var_name, target_value) in conditions {
if let Some(var_index) = self.newton.values.iter().position(|v| v == var_name) {
let current_value = y[var_index];
if (current_value - target_value).abs() <= self.neighborhood_check {
return true;
}
}
}
}
false
}
pub fn check(&self) -> () {
assert_eq!(!self.y.is_empty(), true, "initial y is empty");
assert_eq!(!self.newton.eq_system.is_empty(), true, "system is empty");
assert_eq!(
!self.newton.initial_guess.is_empty(),
true,
"guess is empty"
);
assert_eq!(!self.newton.values.is_empty(), true, "values are empty");
assert_eq!(!self.newton.arg.is_empty(), true, "arg is empty");
assert_eq!(self.newton.tolerance >= 0.0, true, "tolerance is empty");
assert_eq!(
self.newton.max_iterations >= 1,
true,
"max_iterations is empty"
);
assert_eq!(self.newton.dt >= 0.0, true, "h is empty");
assert_eq!(
self.global_timestepping == true || self.h.is_none(),
true,
"for global timestepping h must be set"
);
}
pub fn _step_impl(&mut self) -> (bool, Option<String>) {
let nr = &mut self.newton;
let guess_i = self.y.clone();
let t_i = nr.dt + self.t;
nr.set_t(t_i);
nr.set_initial_guess(guess_i);
nr.solve();
let result = nr.get_result();
if result.is_none() {
info!("result is None");
return (
false,
Some("maximum number of iterations reached".to_string()),
);
} else {
self.y = nr.get_result().expect("REASON");
self.t = t_i;
return (true, None);
}
}
pub fn step(&mut self) {
let t = self.t;
if t == self.t_bound {
self.t_old = Some(t);
self.status = "finished".to_string();
} else {
let (success, message_) = self._step_impl();
if let Some(message_str) = message_ {
self.message = Some(message_str.to_string());
} else {
self.message = None;
}
if success == false {
self.status = "failed".to_string();
} else {
self.t_old = Some(t);
let _status: String = "running".to_string();
if (self.t - self.t_bound) >= 0.0 {
self.status = "finished".to_string();
}
}
}
}
pub fn main_loop(&mut self) -> () {
let start = Instant::now();
let mut integr_status: Option<i8> = None;
let mut y: Vec<DVector<f64>> = Vec::new();
let mut t: Vec<f64> = Vec::new();
let mut _i: i64 = 0;
while integr_status.is_none() {
self.step();
let _status: i8 = 0;
_i += 1;
if self.status == "finished".to_string() {
integr_status = Some(0)
} else if self.status == "failed".to_string() {
integr_status = Some(-1);
break;
}
if self.check_stop_condition(&self.y) {
self.status = "stopped_by_condition".to_string();
integr_status = Some(0);
}
t.push(self.t);
y.push(self.y.clone());
}
let rows = &y.len();
let cols = &y[0].len();
let mut flat_vec: Vec<f64> = Vec::new();
for vector in y.iter() {
flat_vec.extend(vector)
}
let y_res: DMatrix<f64> = DMatrix::from_vec(*cols, *rows, flat_vec).transpose();
let t_res = DVector::from_vec(t);
let duration = start.elapsed();
info!("Program took {} milliseconds to run", duration.as_millis());
self.t_result = t_res.clone();
self.y_result = y_res.clone();
}
pub fn solve(&mut self) -> () {
self.newton.eq_generate();
self.main_loop();
}
pub fn plot_result(&self) -> () {
plots(
self.newton.arg.clone(),
self.newton.values.clone(),
self.t_result.clone(),
self.y_result.clone(),
);
info!("result plotted");
}
pub fn get_result(&self) -> (Option<DVector<f64>>, Option<DMatrix<f64>>) {
(Some(self.t_result.clone()), Some(self.y_result.clone()))
}
pub fn get_status(&self) -> &String {
&self.status
}
}
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn test_newton_raphson_solver_for_Euler_1() {
let eq1 = Expr::parse_expression("z+y-10.0*x");
let eq2 = Expr::parse_expression("z*y-4.0*x");
let eq_system = vec![eq1, eq2];
info!("eq_system = {:?}", eq_system);
let y0 = DVector::from_vec(vec![1.0, 1.0]);
let values = vec!["z".to_string(), "y".to_string()];
let arg = "x".to_string();
let tolerance = 1e-2;
let max_iterations = 50;
let h = Some(1e-2);
let t0 = 0.0;
let t_bound = 1.0;
let mut solver = BE::new();
solver.set_initial(
eq_system,
values,
arg,
tolerance,
max_iterations,
h,
t0,
t_bound,
y0,
);
info!(
"y = {:?}, initial_guess = {:?}",
solver.newton.y, solver.newton.initial_guess
);
solver.newton.eq_generate();
let (success, message) = solver._step_impl();
assert_eq!(solver.y.len(), 2);
assert_eq!(success, true, "success = {} must be true", success);
assert_eq!(message, None, "message = {:?} must be None", message);
}
#[test]
fn test_newton_raphson_solver_for_Euler_2() {
let eq1 = Expr::parse_expression("z+y-10.0*x");
let eq2 = Expr::parse_expression("z*y-4.0*x");
let eq_system = vec![eq1, eq2];
info!("eq_system = {:?}", eq_system);
let y0 = DVector::from_vec(vec![1.0, 1.0]);
let values = vec!["z".to_string(), "y".to_string()];
let arg = "x".to_string();
let tolerance = 1e-2;
let max_iterations = 50;
let h = Some(1e-2);
let t0 = 0.0;
let t_bound = 1.0;
let mut solver = BE::new();
solver.set_initial(
eq_system,
values,
arg,
tolerance,
max_iterations,
h,
t0,
t_bound,
y0,
);
info!(
"y = {:?}, initial_guess = {:?}",
solver.newton.y, solver.newton.initial_guess
);
solver.newton.eq_generate();
solver.step();
assert_eq!(solver.status, "running".to_string());
}
#[test]
fn test_newton_raphson_solver_for_Euler_3() {
let eq1 = Expr::parse_expression("z+y-10.0*x");
let eq2 = Expr::parse_expression("z*y-4.0*x");
let eq_system = vec![eq1, eq2];
info!("eq_system = {:?}", eq_system);
let y0 = DVector::from_vec(vec![1.0, 1.0]);
let values = vec!["z".to_string(), "y".to_string()];
let arg = "x".to_string();
let tolerance = 1e-2;
let max_iterations = 50;
let h = Some(1e-2);
let t0 = 0.0;
let t_bound = 1.0;
let mut solver = BE::new();
solver.set_initial(
eq_system,
values,
arg,
tolerance,
max_iterations,
h,
t0,
t_bound,
y0,
);
info!(
"y = {:?}, initial_guess = {:?}",
solver.newton.y, solver.newton.initial_guess
);
solver.newton.eq_generate();
solver.solve();
assert_eq!(solver.status, "finished".to_string());
}
#[test]
fn test_newton_raphson_solver_for_Euler_4() {
let eq1 = Expr::parse_expression("z+y-10.0*x");
let eq2 = Expr::parse_expression("z*y-4.0*x");
let eq_system = vec![eq1, eq2];
info!("eq_system = {:?}", eq_system);
let y0 = DVector::from_vec(vec![1.0, 1.0]);
let values = vec!["z".to_string(), "y".to_string()];
let arg = "x".to_string();
let tolerance = 1e-2;
let max_iterations = 50;
let h = None;
let t0 = 0.0;
let t_bound = 1.0;
let mut solver = BE::new();
solver.set_initial(
eq_system,
values,
arg,
tolerance,
max_iterations,
h,
t0,
t_bound,
y0,
);
info!(
"y = {:?}, initial_guess = {:?}",
solver.newton.y, solver.newton.initial_guess
);
solver.newton.eq_generate();
solver.solve();
let res = solver.get_result();
let _result = res.1.unwrap();
assert_eq!(solver.status, "finished".to_string());
}
#[test]
fn test_be_stop_condition_single_variable() {
let eq1 = Expr::parse_expression("-z+2.0*x"); let eq_system = vec![eq1];
let y0 = DVector::from_vec(vec![1.0]);
let values = vec!["z".to_string()];
let arg = "x".to_string();
let tolerance = 1e-6;
let max_iterations = 50;
let h = Some(0.01);
let t0 = 0.0;
let t_bound = 10.0;
let mut solver = BE::new();
solver.set_initial(
eq_system,
values,
arg,
tolerance,
max_iterations,
h,
t0,
t_bound,
y0,
);
let mut stop_condition = HashMap::new();
stop_condition.insert("z".to_string(), 1.5);
solver.set_stop_condition(stop_condition);
solver.set_neighborhood_check(1e-2);
solver.solve();
assert_eq!(solver.get_status(), "stopped_by_condition");
let (_, y_result) = solver.get_result();
let y_res = y_result.unwrap();
let final_y = y_res[(y_res.nrows() - 1, 0)];
assert!((final_y - 1.5).abs() <= 1e-2);
}
#[test]
fn test_be_stop_condition_multiple_variables() {
let eq1 = Expr::parse_expression("z+y-2.0*x");
let eq2 = Expr::parse_expression("-z*y+3.0*x");
let eq_system = vec![eq1, eq2];
let y0 = DVector::from_vec(vec![1.0, 1.0]);
let values = vec!["z".to_string(), "y".to_string()];
let arg = "x".to_string();
let tolerance = 1e-6;
let max_iterations = 50;
let h = Some(0.01);
let t0 = 0.0;
let t_bound = 10.0;
let mut solver = BE::new();
solver.set_initial(
eq_system,
values,
arg,
tolerance,
max_iterations,
h,
t0,
t_bound,
y0,
);
let mut stop_condition = HashMap::new();
stop_condition.insert("z".to_string(), 1.2);
solver.set_stop_condition(stop_condition);
solver.set_neighborhood_check(1e-2);
solver.solve();
assert_eq!(solver.get_status(), "stopped_by_condition");
let (_, y_result) = solver.get_result();
let y_res = y_result.unwrap();
let final_z = y_res[(y_res.nrows() - 1, 0)];
assert!((final_z - 1.2).abs() <= 1e-2);
}
#[test]
fn test_be_no_stop_condition() {
let eq1 = Expr::parse_expression("z+y-10.0*x");
let eq2 = Expr::parse_expression("z*y-4.0*x");
let eq_system = vec![eq1, eq2];
let y0 = DVector::from_vec(vec![1.0, 1.0]);
let values = vec!["z".to_string(), "y".to_string()];
let arg = "x".to_string();
let tolerance = 1e-2;
let max_iterations = 50;
let h = Some(1e-2);
let t0 = 0.0;
let t_bound = 0.1;
let mut solver = BE::new();
solver.set_initial(
eq_system,
values,
arg,
tolerance,
max_iterations,
h,
t0,
t_bound,
y0,
);
solver.solve();
assert_eq!(solver.get_status(), "finished");
let (t_result, _) = solver.get_result();
let t_res = t_result.unwrap();
let final_t = t_res[t_res.len() - 1];
assert!((final_t - t_bound).abs() <= h.unwrap());
}
}