use std::sync::Arc;
use super::super::{
advecator::Advector,
algos::ht::HtTensor,
integrator::{StepProducts, StepTimings, TimeIntegrator},
phasespace::PhaseSpaceRepr,
progress::{StepPhase, StepProgress},
solver::PoissonSolver,
types::*,
};
use super::helpers;
use crate::CausticError;
use super::bug::{BugConfig, bug_drift_substep, bug_kick_substep, conservative_correction};
pub struct RkBugConfig {
pub tolerance: f64,
pub max_rank: usize,
pub rk_order: usize,
pub rank_increase: usize,
pub conservative: bool,
}
impl Default for RkBugConfig {
fn default() -> Self {
Self {
tolerance: 1e-8,
max_rank: 50,
rk_order: 3,
rank_increase: 2,
conservative: false,
}
}
}
pub struct RkBugIntegrator {
pub config: RkBugConfig,
pub g: f64,
last_timings: StepTimings,
progress: Option<Arc<StepProgress>>,
}
impl RkBugIntegrator {
pub fn new(g: f64, config: RkBugConfig) -> Self {
Self {
config,
g,
last_timings: StepTimings::default(),
progress: None,
}
}
fn bug_config(&self) -> BugConfig {
BugConfig {
tolerance: self.config.tolerance,
max_rank: self.config.max_rank,
midpoint: false,
conservative: false,
rank_increase: self.config.rank_increase,
}
}
fn bug_strang_step(
ht: &mut HtTensor,
solver: &dyn PoissonSolver,
g: f64,
dt: f64,
config: &BugConfig,
timings: &mut StepTimings,
) {
helpers::time_ms!(timings, drift_ms, bug_drift_substep(ht, dt / 2.0, config));
let (_, _, accel) =
helpers::time_ms!(timings, poisson_ms, helpers::solve_poisson(ht, solver, g));
helpers::time_ms!(timings, kick_ms, bug_kick_substep(ht, &accel, dt, config));
helpers::time_ms!(timings, drift_ms, bug_drift_substep(ht, dt / 2.0, config));
}
fn ssp_rk3_step(
&self,
repr: &mut dyn PhaseSpaceRepr,
solver: &dyn PoissonSolver,
dt: f64,
timings: &mut StepTimings,
) {
let Some(ht) = repr.as_any_mut().downcast_mut::<HtTensor>() else {
debug_assert!(false, "RK-BUG requires HtTensor");
return;
};
let density_before = if self.config.conservative {
Some(ht.compute_density())
} else {
None
};
let config = self.bug_config();
let tol = self.config.tolerance;
let y0 = ht.clone();
helpers::report_phase!(self.progress, StepPhase::BugKStep, 0, 4);
Self::bug_strang_step(ht, solver, self.g, dt, &config, timings);
helpers::report_phase!(self.progress, StepPhase::BugKStep, 1, 4);
let mut z2 = ht.clone();
Self::bug_strang_step(&mut z2, solver, self.g, dt, &config, timings);
let y2 = y0.scaled_add(3.0 / 4.0, &z2, 1.0 / 4.0, tol);
*ht = y2;
helpers::report_phase!(self.progress, StepPhase::BugLStep, 2, 4);
let mut z3 = ht.clone();
Self::bug_strang_step(&mut z3, solver, self.g, dt, &config, timings);
let result = y0.scaled_add(1.0 / 3.0, &z3, 2.0 / 3.0, tol);
*ht = result;
helpers::report_phase!(self.progress, StepPhase::BugSStep, 3, 4);
if let Some(ref dens) = density_before {
conservative_correction(ht, dens);
}
}
fn strang_fallback(
&self,
repr: &mut dyn PhaseSpaceRepr,
solver: &dyn PoissonSolver,
advector: &dyn Advector,
dt: f64,
timings: &mut StepTimings,
) {
helpers::time_ms!(timings, drift_ms, advector.drift(repr, dt / 2.0));
let (_, _, 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));
helpers::time_ms!(timings, drift_ms, advector.drift(repr, dt / 2.0));
}
}
impl TimeIntegrator for RkBugIntegrator {
fn advance(
&mut self,
repr: &mut dyn PhaseSpaceRepr,
solver: &dyn PoissonSolver,
advector: &dyn Advector,
dt: f64,
) -> Result<StepProducts, CausticError> {
let _span = tracing::info_span!("rk_bug_advance").entered();
let mut timings = StepTimings::default();
if let Some(ref p) = self.progress {
p.start_step();
}
let is_ht = repr.as_any().downcast_ref::<HtTensor>().is_some();
if is_ht {
self.ssp_rk3_step(repr, solver, dt, &mut timings);
} else {
self.strang_fallback(repr, solver, advector, dt, &mut timings);
}
helpers::report_phase!(self.progress, StepPhase::StepComplete, 4, 4);
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);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rk_bug_config_defaults() {
let cfg = RkBugConfig::default();
assert_eq!(cfg.rk_order, 3);
assert_eq!(cfg.max_rank, 50);
assert_eq!(cfg.rank_increase, 2);
assert!(!cfg.conservative);
}
}