use std::sync::Arc;
use super::super::{
advecator::Advector,
integrator::{StepProducts, StepTimings, TimeIntegrator},
phasespace::PhaseSpaceRepr,
progress::{StepPhase, StepProgress},
solver::PoissonSolver,
types::*,
};
use super::helpers;
use crate::CausticError;
pub struct PIController {
pub tolerance: f64,
pub k_i: f64,
pub k_p: f64,
pub dt_min: f64,
pub dt_max: f64,
pub safety_factor: f64,
prev_err: Option<f64>,
}
impl PIController {
pub fn new(tolerance: f64, order: f64) -> Self {
Self {
tolerance,
k_i: 0.3 / order,
k_p: 0.4 / order,
dt_min: 1e-15,
dt_max: 1e10,
safety_factor: 0.9,
prev_err: None,
}
}
pub fn step(&mut self, dt: f64, err: f64) -> (f64, bool) {
let accepted = err <= self.tolerance;
let ratio = if err > 1e-30 {
self.tolerance / err
} else {
5.0
};
let dt_new = if let Some(prev) = self.prev_err {
let prev_ratio = if prev > 1e-30 {
self.tolerance / prev
} else {
5.0
};
self.safety_factor * dt * ratio.powf(self.k_i) * prev_ratio.powf(-self.k_p)
} else {
self.safety_factor * dt * ratio.powf(self.k_i)
};
let dt_new = dt_new.clamp(self.dt_min, self.dt_max);
let dt_new = dt_new.min(5.0 * dt).max(0.2 * dt);
if accepted {
self.prev_err = Some(err);
}
(dt_new, accepted)
}
}
pub struct AdaptiveStrangSplitting {
pub g: f64,
pub controller: PIController,
suggested_dt: Option<f64>,
last_timings: StepTimings,
progress: Option<Arc<StepProgress>>,
pub max_retries: usize,
}
impl AdaptiveStrangSplitting {
pub fn new(g: f64, tolerance: f64) -> Self {
Self {
g,
controller: PIController::new(tolerance, 2.0),
suggested_dt: None,
last_timings: StepTimings::default(),
progress: None,
max_retries: 5,
}
}
fn relative_error(a: &[f64], b: &[f64]) -> f64 {
let mut diff_sq = 0.0f64;
let mut norm_sq = 0.0f64;
for (&ai, &bi) in a.iter().zip(b.iter()) {
let d = ai - bi;
diff_sq += d * d;
norm_sq += ai * ai;
}
if norm_sq > 1e-30 {
(diff_sq / norm_sq).sqrt()
} else {
diff_sq.sqrt()
}
}
}
impl TimeIntegrator for AdaptiveStrangSplitting {
fn advance(
&mut self,
repr: &mut dyn PhaseSpaceRepr,
solver: &dyn PoissonSolver,
advector: &dyn Advector,
dt: f64,
) -> Result<StepProducts, CausticError> {
let _span = tracing::info_span!("adaptive_strang_advance").entered();
let mut timings = StepTimings::default();
if let Some(ref p) = self.progress {
p.start_step();
}
let mut dt_try = self.suggested_dt.unwrap_or(dt).min(dt);
for _retry in 0..self.max_retries {
let snap_for_lie = repr.to_snapshot(0.0).ok_or_else(|| {
CausticError::Solver("adaptive integrator requires to_snapshot support".into())
})?;
let snap_for_rollback = repr.to_snapshot(0.0).ok_or_else(|| {
CausticError::Solver("adaptive integrator requires to_snapshot support".into())
})?;
helpers::report_phase!(self.progress, StepPhase::DriftHalf1, 0, 7);
{
let _s = tracing::info_span!("strang_drift_half").entered();
helpers::time_ms!(timings, drift_ms, advector.drift(repr, dt_try / 2.0));
}
helpers::report_phase!(self.progress, StepPhase::PoissonSolve, 1, 7);
let accel = {
let _s = tracing::info_span!("strang_poisson").entered();
let (_density, _potential, accel) = helpers::time_ms!(
timings,
poisson_ms,
helpers::solve_poisson(repr, solver, self.g)
);
accel
};
helpers::report_phase!(self.progress, StepPhase::Kick, 2, 7);
{
let _s = tracing::info_span!("strang_kick").entered();
helpers::time_ms!(timings, kick_ms, advector.kick(repr, &accel, dt_try));
}
helpers::report_phase!(self.progress, StepPhase::DriftHalf2, 3, 7);
{
let _s = tracing::info_span!("strang_drift_half").entered();
helpers::time_ms!(timings, drift_ms, advector.drift(repr, dt_try / 2.0));
}
let strang_snap = repr.to_snapshot(0.0).ok_or_else(|| {
CausticError::Solver("adaptive integrator requires to_snapshot support".into())
})?;
repr.load_snapshot(snap_for_lie)?;
helpers::report_phase!(self.progress, StepPhase::DriftHalf1, 4, 7);
{
let _s = tracing::info_span!("lie_drift").entered();
helpers::time_ms!(timings, drift_ms, advector.drift(repr, dt_try));
}
helpers::report_phase!(self.progress, StepPhase::Kick, 5, 7);
{
let _s = tracing::info_span!("lie_kick").entered();
let (_density, _potential, accel) = helpers::time_ms!(
timings,
poisson_ms,
helpers::solve_poisson(repr, solver, self.g)
);
helpers::time_ms!(timings, kick_ms, advector.kick(repr, &accel, dt_try));
}
let lie_snap = repr.to_snapshot(0.0).ok_or_else(|| {
CausticError::Solver("adaptive integrator requires to_snapshot support".into())
})?;
let err = Self::relative_error(&strang_snap.data, &lie_snap.data);
helpers::report_phase!(self.progress, StepPhase::Diagnostics, 6, 7);
let (dt_new, accepted) = self.controller.step(dt_try, err);
self.suggested_dt = Some(dt_new);
if accepted {
repr.load_snapshot(strang_snap)?;
break;
} else {
repr.load_snapshot(snap_for_rollback)?;
dt_try = dt_new.min(dt);
}
}
helpers::report_phase!(self.progress, StepPhase::StepComplete, 0, 0);
let (density, potential, acceleration) = helpers::time_ms!(
timings,
density_ms,
helpers::solve_poisson(repr, solver, self.g)
);
self.last_timings = timings;
Ok(StepProducts {
density,
potential,
acceleration,
})
}
fn max_dt(&self, repr: &dyn PhaseSpaceRepr, cfl_factor: f64) -> f64 {
helpers::dynamical_timestep(repr, self.g, cfl_factor)
}
fn last_step_timings(&self) -> Option<&StepTimings> {
Some(&self.last_timings)
}
fn set_progress(&mut self, progress: Arc<StepProgress>) {
self.progress = Some(progress);
}
}