use std::collections::HashMap;
use nalgebra::DVector;
use crate::context::error::OxiflowError;
use crate::context::value::ContextValue;
use crate::context::ContextCalculator;
use crate::solver::chain::build_calculator_chain;
use crate::solver::config::StepControl;
use crate::solver::methods::step_control::StepSizeController;
use crate::solver::methods::{check_finite, evaluate_derivative};
use crate::solver::scenario::{Domain, Scenario};
use crate::solver::{SimulationResult, Solver, SolverConfiguration};
const C2: f64 = 1.0 / 5.0;
const C3: f64 = 3.0 / 10.0;
const C4: f64 = 4.0 / 5.0;
const C5: f64 = 8.0 / 9.0;
const C6: f64 = 1.0;
const C7: f64 = 1.0;
const A21: f64 = 1.0 / 5.0;
const A31: f64 = 3.0 / 40.0;
const A32: f64 = 9.0 / 40.0;
const A41: f64 = 44.0 / 45.0;
const A42: f64 = -56.0 / 15.0;
const A43: f64 = 32.0 / 9.0;
const A51: f64 = 19372.0 / 6561.0;
const A52: f64 = -25360.0 / 2187.0;
const A53: f64 = 64448.0 / 6561.0;
const A54: f64 = -212.0 / 729.0;
const A61: f64 = 9017.0 / 3168.0;
const A62: f64 = -355.0 / 33.0;
const A63: f64 = 46732.0 / 5247.0;
const A64: f64 = 49.0 / 176.0;
const A65: f64 = -5103.0 / 18656.0;
const A71: f64 = 35.0 / 384.0;
const A73: f64 = 500.0 / 1113.0;
const A74: f64 = 125.0 / 192.0;
const A75: f64 = -2187.0 / 6784.0;
const A76: f64 = 11.0 / 84.0;
const E1: f64 = 71.0 / 57600.0;
const E2: f64 = 0.0;
const E3: f64 = -71.0 / 16695.0;
const E4: f64 = 71.0 / 1920.0;
const E5: f64 = -17253.0 / 339200.0;
const E6: f64 = 22.0 / 525.0;
const E7: f64 = -1.0 / 40.0;
const ERROR_ESTIMATOR_ORDER: f64 = 4.0;
const MAX_REJECTIONS_PER_STEP: usize = 50;
pub struct DoPri45Solver;
impl Solver for DoPri45Solver {
fn solve(
&self,
scenario: &Scenario,
config: &SolverConfiguration,
) -> Result<SimulationResult, OxiflowError> {
scenario.validate()?;
let domain = scenario.single_domain()?;
let (dt_init, dt_min, dt_max, rtol, atol) = match &config.time.step_control {
StepControl::Adaptive {
dt_init,
dt_min,
dt_max,
rtol,
atol,
} => (*dt_init, *dt_min, *dt_max, *rtol, *atol),
_ => {
return Err(OxiflowError::InvalidDomain(
"DoPri45Solver only supports StepControl::Adaptive".into(),
))
}
};
let t_end = config.time.t_end;
let t_start = scenario.t_start;
if dt_init <= 0.0 || dt_min <= 0.0 || dt_max < dt_min {
return Err(OxiflowError::InvalidDomain(
"dt_init and dt_min must be strictly positive, and dt_max >= dt_min".into(),
));
}
if t_end <= t_start {
return Err(OxiflowError::InvalidDomain(
"t_end must be greater than t_start".into(),
));
}
let requirements = scenario.context_requirements();
let chain = build_calculator_chain(&requirements, &config.calculators)?;
let mut u = domain.model.initial_state(domain.mesh.as_ref());
let mut t = t_start;
let mut dt = dt_init;
let mut controller =
StepSizeController::new(rtol, atol, dt_min, dt_max, ERROR_ESTIMATOR_ORDER);
let save_every = config.time.save_every.unwrap_or(1);
let mut accepted_since_save = 0usize;
let mut states: Vec<ContextValue> = vec![u.clone()];
let mut times: Vec<f64> = vec![t_start];
let mut accepted_steps = 0usize;
let mut rejected_steps = 0usize;
while t < t_end - 1e-12 {
let mut local_dt = dt.min(t_end - t);
let mut rejections_this_step = 0usize;
loop {
let mut attempt_state = u.clone();
let (y5_field, error_field) =
dopri45_stages(domain, &chain, &mut attempt_state, t, local_dt)?;
let error_norm = controller.error_norm(&error_field, &y5_field);
if controller.accept(error_norm) {
u = ContextValue::ScalarField(y5_field);
t += local_dt;
accepted_steps += 1;
accepted_since_save += 1;
check_finite(&u, t)?;
dt = controller.next_dt(local_dt, error_norm);
if accepted_since_save >= save_every {
states.push(u.clone());
times.push(t);
accepted_since_save = 0;
}
break;
}
rejected_steps += 1;
rejections_this_step += 1;
let suggested = controller.next_dt(local_dt, error_norm);
if rejections_this_step >= MAX_REJECTIONS_PER_STEP
|| suggested <= controller.dt_min()
{
return Err(OxiflowError::SolverDivergence {
time: t,
reason: format!(
"step rejected {rejections_this_step} times in a row; cannot \
satisfy rtol/atol even at dt_min={}",
controller.dt_min()
),
});
}
local_dt = suggested;
}
}
let mut metadata = HashMap::new();
metadata.insert("solver.accepted_steps".to_string(), accepted_steps as f64);
metadata.insert("solver.rejected_steps".to_string(), rejected_steps as f64);
Ok(SimulationResult {
states,
times,
n_steps: accepted_steps,
metadata,
})
}
}
fn dopri45_stages(
domain: &Domain,
chain: &[&dyn ContextCalculator],
state: &mut ContextValue,
t: f64,
dt: f64,
) -> Result<(DVector<f64>, DVector<f64>), OxiflowError> {
let k1_val = evaluate_derivative(domain, chain, state, t, dt)?;
let u_field = state.as_scalar_field()?.clone();
let k1 = k1_val.as_scalar_field()?.clone();
let y2 = u_field.clone() + k1.clone() * (dt * A21);
let mut s2 = ContextValue::ScalarField(y2);
let k2_val = evaluate_derivative(domain, chain, &mut s2, t + C2 * dt, dt)?;
let k2 = k2_val.as_scalar_field()?.clone();
let y3 = u_field.clone() + k1.clone() * (dt * A31) + k2.clone() * (dt * A32);
let mut s3 = ContextValue::ScalarField(y3);
let k3_val = evaluate_derivative(domain, chain, &mut s3, t + C3 * dt, dt)?;
let k3 = k3_val.as_scalar_field()?.clone();
let y4 = u_field.clone()
+ k1.clone() * (dt * A41)
+ k2.clone() * (dt * A42)
+ k3.clone() * (dt * A43);
let mut s4 = ContextValue::ScalarField(y4);
let k4_val = evaluate_derivative(domain, chain, &mut s4, t + C4 * dt, dt)?;
let k4 = k4_val.as_scalar_field()?.clone();
let y5s = u_field.clone()
+ k1.clone() * (dt * A51)
+ k2.clone() * (dt * A52)
+ k3.clone() * (dt * A53)
+ k4.clone() * (dt * A54);
let mut s5 = ContextValue::ScalarField(y5s);
let k5_val = evaluate_derivative(domain, chain, &mut s5, t + C5 * dt, dt)?;
let k5 = k5_val.as_scalar_field()?.clone();
let y6 = u_field.clone()
+ k1.clone() * (dt * A61)
+ k2.clone() * (dt * A62)
+ k3.clone() * (dt * A63)
+ k4.clone() * (dt * A64)
+ k5.clone() * (dt * A65);
let mut s6 = ContextValue::ScalarField(y6);
let k6_val = evaluate_derivative(domain, chain, &mut s6, t + C6 * dt, dt)?;
let k6 = k6_val.as_scalar_field()?.clone();
let y7 = u_field.clone()
+ k1.clone() * (dt * A71)
+ k3.clone() * (dt * A73)
+ k4.clone() * (dt * A74)
+ k5.clone() * (dt * A75)
+ k6.clone() * (dt * A76);
let mut s7 = ContextValue::ScalarField(y7.clone());
let k7_val = evaluate_derivative(domain, chain, &mut s7, t + C7 * dt, dt)?;
let k7 = k7_val.as_scalar_field()?.clone();
let y_5th = y7;
let error: DVector<f64> = k1 * (dt * E1)
+ k2 * (dt * E2)
+ k3 * (dt * E3)
+ k4 * (dt * E4)
+ k5 * (dt * E5)
+ k6 * (dt * E6)
+ k7 * (dt * E7);
Ok((y_5th, error))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::context::compute::ComputeContext;
use crate::context::variable::ContextVariable;
use crate::mesh::{Mesh, UniformGrid1D};
use crate::model::traits::{PhysicalModel, RequiresContext};
use crate::solver::config::{IntegratorKind, SolverConfiguration, TimeConfiguration};
use crate::solver::scenario::Scenario;
#[derive(Debug)]
struct ExponentialDecay {
lambda: f64,
}
impl RequiresContext for ExponentialDecay {
fn required_variables(&self) -> Vec<ContextVariable> {
vec![]
}
}
impl PhysicalModel for ExponentialDecay {
fn compute_physics(
&self,
state: &ContextValue,
_ctx: &ComputeContext,
) -> Result<ContextValue, OxiflowError> {
let u = state.as_scalar_field()?;
Ok(ContextValue::ScalarField(u.map(|v| -self.lambda * v)))
}
fn initial_state(&self, mesh: &dyn Mesh) -> ContextValue {
ContextValue::ScalarField(DVector::from_element(mesh.n_dof(), 1.0))
}
fn name(&self) -> &str {
"exponential_decay"
}
}
#[derive(Debug)]
struct ZeroDerivative;
impl RequiresContext for ZeroDerivative {
fn required_variables(&self) -> Vec<ContextVariable> {
vec![]
}
}
impl PhysicalModel for ZeroDerivative {
fn compute_physics(
&self,
state: &ContextValue,
_ctx: &ComputeContext,
) -> Result<ContextValue, OxiflowError> {
let u = state.as_scalar_field()?;
Ok(ContextValue::ScalarField(DVector::zeros(u.len())))
}
fn initial_state(&self, mesh: &dyn Mesh) -> ContextValue {
ContextValue::ScalarField(DVector::from_element(mesh.n_dof(), 2.5))
}
fn name(&self) -> &str {
"zero_derivative"
}
}
fn make_mesh(n: usize) -> Box<dyn Mesh> {
Box::new(UniformGrid1D::new(n, 0.0, 1.0).unwrap())
}
fn make_config(t_end: f64, dt_init: f64, rtol: f64, atol: f64) -> SolverConfiguration {
SolverConfiguration::new(
TimeConfiguration::new(
t_end,
StepControl::Adaptive {
dt_init,
dt_min: 1e-8,
dt_max: 1.0,
rtol,
atol,
},
),
IntegratorKind::DoPri45,
)
}
#[test]
fn row_sums_match_c_nodes() {
assert!((A21 - C2).abs() < 1e-12);
assert!(((A31 + A32) - C3).abs() < 1e-12);
assert!(((A41 + A42 + A43) - C4).abs() < 1e-12);
assert!(((A51 + A52 + A53 + A54) - C5).abs() < 1e-12);
assert!(((A61 + A62 + A63 + A64 + A65) - C6).abs() < 1e-12);
assert!(((A71 + A73 + A74 + A75 + A76) - C7).abs() < 1e-12); }
#[test]
fn fifth_order_weights_sum_to_one() {
let sum = A71 + A73 + A74 + A75 + A76; assert!((sum - 1.0).abs() < 1e-12);
}
#[test]
fn error_weights_are_nonzero_where_expected() {
assert_eq!(E2, 0.0);
assert_ne!(E1, 0.0);
assert_ne!(E3, 0.0);
assert_ne!(E4, 0.0);
assert_ne!(E5, 0.0);
assert_ne!(E6, 0.0);
assert_ne!(E7, 0.0);
}
#[test]
fn zero_derivative_field_stays_constant() {
let scenario = Scenario::single(Box::new(ZeroDerivative), make_mesh(5));
let config = make_config(1.0, 0.1, 1e-6, 1e-9);
let result = DoPri45Solver.solve(&scenario, &config).unwrap();
for state in &result.states {
let field = state.as_scalar_field().unwrap();
for v in field.iter() {
assert!((v - 2.5).abs() < 1e-9);
}
}
}
#[test]
fn exponential_decay_within_tolerance() {
let lambda = 2.0;
let rtol = 1e-8;
let atol = 1e-10;
let scenario = Scenario::single(Box::new(ExponentialDecay { lambda }), make_mesh(2));
let config = make_config(1.0, 0.1, rtol, atol);
let result = DoPri45Solver.solve(&scenario, &config).unwrap();
let expected = (-lambda * 1.0_f64).exp();
let final_field = result.states.last().unwrap().as_scalar_field().unwrap();
for v in final_field.iter() {
assert!(
(v - expected).abs() < 1e-5,
"got {v}, expected {expected} (lambda={lambda})"
);
}
}
#[test]
fn t_final_reaches_t_end() {
let scenario = Scenario::single(Box::new(ExponentialDecay { lambda: 1.0 }), make_mesh(2));
let config = make_config(2.0, 0.1, 1e-6, 1e-9);
let result = DoPri45Solver.solve(&scenario, &config).unwrap();
assert!((result.t_final().unwrap() - 2.0).abs() < 1e-6);
}
#[test]
fn metadata_reports_accepted_and_rejected_steps() {
let scenario = Scenario::single(Box::new(ExponentialDecay { lambda: 5.0 }), make_mesh(2));
let config = make_config(1.0, 0.1, 1e-6, 1e-9);
let result = DoPri45Solver.solve(&scenario, &config).unwrap();
assert!(result.metadata.contains_key("solver.accepted_steps"));
assert!(result.metadata.contains_key("solver.rejected_steps"));
assert!(result.metadata["solver.accepted_steps"] > 0.0);
assert_eq!(
result.metadata["solver.accepted_steps"],
result.n_steps as f64
);
}
#[test]
fn dt_min_guard_raises_solver_divergence() {
let scenario = Scenario::single(Box::new(ExponentialDecay { lambda: 50.0 }), make_mesh(2));
let config = SolverConfiguration::new(
TimeConfiguration::new(
1.0,
StepControl::Adaptive {
dt_init: 0.5,
dt_min: 0.5,
dt_max: 0.5,
rtol: 1e-15,
atol: 1e-15,
},
),
IntegratorKind::DoPri45,
);
let err = DoPri45Solver.solve(&scenario, &config).unwrap_err();
assert!(matches!(err, OxiflowError::SolverDivergence { .. }));
}
#[test]
fn fixed_step_control_returns_error() {
let scenario = Scenario::single(Box::new(ZeroDerivative), make_mesh(2));
let config = SolverConfiguration::new(
TimeConfiguration::new(1.0, StepControl::Fixed { dt: 0.1 }),
IntegratorKind::DoPri45,
);
assert!(DoPri45Solver.solve(&scenario, &config).is_err());
}
#[test]
fn invalid_dt_bounds_return_error() {
let scenario = Scenario::single(Box::new(ZeroDerivative), make_mesh(2));
let config = SolverConfiguration::new(
TimeConfiguration::new(
1.0,
StepControl::Adaptive {
dt_init: 0.1,
dt_min: 0.5, dt_max: 0.2,
rtol: 1e-6,
atol: 1e-9,
},
),
IntegratorKind::DoPri45,
);
assert!(DoPri45Solver.solve(&scenario, &config).is_err());
}
#[test]
fn t_end_before_t_start_returns_error() {
let scenario = Scenario::single(Box::new(ZeroDerivative), make_mesh(2)).with_t_start(5.0);
let config = make_config(1.0, 0.1, 1e-6, 1e-9);
assert!(DoPri45Solver.solve(&scenario, &config).is_err());
}
#[test]
fn missing_calculator_returns_error() {
#[derive(Debug)]
struct NeedsExternal;
impl RequiresContext for NeedsExternal {
fn required_variables(&self) -> Vec<ContextVariable> {
vec![ContextVariable::External {
name: "missing".into(),
}]
}
}
impl PhysicalModel for NeedsExternal {
fn compute_physics(
&self,
s: &ContextValue,
_: &ComputeContext,
) -> Result<ContextValue, OxiflowError> {
Ok(s.clone())
}
fn initial_state(&self, mesh: &dyn Mesh) -> ContextValue {
ContextValue::ScalarField(DVector::from_element(mesh.n_dof(), 0.0))
}
fn name(&self) -> &str {
"needs_external"
}
}
let scenario = Scenario::single(Box::new(NeedsExternal), make_mesh(2));
let config = make_config(1.0, 0.1, 1e-6, 1e-9);
let err = DoPri45Solver.solve(&scenario, &config).unwrap_err();
assert!(matches!(err, OxiflowError::MissingCalculator(_)));
}
}