use std::sync::Arc;
use rayon::prelude::*;
use rust_decimal::Decimal;
use crate::tooling::core::{
advecator::Advector,
conditions::{ExitReason, TimeLimitCondition},
conservation::lomac::LoMaC,
diagnostics::{Diagnostics, GlobalDiagnostics},
init::{domain::Domain, input::optional::OptionalParams},
integrator::{StepTimings, TimeIntegrator},
io::{IOManager, OutputFormat},
output::exit::{package::ExitPackage, standard::ExitEvaluator},
phasespace::PhaseSpaceRepr,
progress::{StepPhase, StepProgress},
solver::PoissonSolver,
types::*,
};
#[inline]
pub fn dec(v: f64) -> Decimal {
Decimal::from_f64_retain(v).unwrap_or(Decimal::ZERO)
}
pub fn f64_to_decimal(v: f64, field: &str) -> Decimal {
Decimal::from_f64_retain(v).unwrap_or_else(|| {
tracing::warn!(
"{field}({v}) is not representable as Decimal (NaN/Inf/subnormal); defaulting to 0. \
Use the _decimal() variant for exact values."
);
Decimal::ZERO
})
}
pub struct Simulation {
pub domain: Domain,
pub repr: Box<dyn PhaseSpaceRepr>,
pub poisson: Box<dyn PoissonSolver>,
pub advector: Box<dyn Advector>,
pub integrator: Box<dyn TimeIntegrator>,
pub diagnostics: Diagnostics,
pub io: IOManager,
pub exit_evaluator: ExitEvaluator,
pub opts: OptionalParams,
pub lomac: Option<LoMaC>,
pub g: f64,
pub time: f64,
pub step: u64,
pub start_time: std::time::Instant,
cached_rho_max: Option<f64>,
pub last_step_timings: StepTimings,
pub cached_density: Option<DensityField>,
pub cached_potential: Option<PotentialField>,
pub cached_acceleration: Option<AccelerationField>,
progress: Option<Arc<StepProgress>>,
}
impl Simulation {
pub fn builder() -> SimulationBuilder {
SimulationBuilder::new()
}
pub fn run(&mut self) -> anyhow::Result<ExitPackage> {
loop {
if let Some(reason) = self.step()? {
let snapshot =
self.repr
.to_snapshot(self.time)
.unwrap_or_else(|| PhaseSpaceSnapshot {
data: vec![],
shape: [0; 6],
time: self.time,
});
let history = self.diagnostics.history.clone();
let wall_secs = self.start_time.elapsed().as_secs_f64();
return Ok(ExitPackage::assemble(
snapshot,
history,
reason,
String::new(),
wall_secs,
self.step,
0,
));
}
}
}
pub fn step(&mut self) -> anyhow::Result<Option<ExitReason>> {
let cfl_factor = self.opts.cfl_factor_f64();
let mut dt = if let Some(rho_max) = self.cached_rho_max.take() {
if rho_max <= 0.0 || self.g <= 0.0 {
1e10
} else {
let t_dyn = 1.0 / (self.g * rho_max).sqrt();
cfl_factor * t_dyn
}
} else {
self.integrator.max_dt(&*self.repr, cfl_factor)
};
if let Some(adaptive_dt) = self.integrator.suggested_dt() {
dt = dt.min(adaptive_dt);
}
let t_final = {
use rust_decimal::prelude::ToPrimitive;
self.domain.time_range.t_final.to_f64().unwrap_or(f64::MAX)
};
if self.time + dt > t_final {
dt = (t_final - self.time).max(0.0);
}
let products =
self.integrator
.advance(&mut *self.repr, &*self.poisson, &*self.advector, dt)?;
let mut timings = self
.integrator
.last_step_timings()
.cloned()
.unwrap_or_default();
if let Some(ref p) = self.progress {
let snap = p.read();
let post_count: u8 = 2 + u8::from(self.lomac.is_some());
p.set_sub_step(snap.sub_step, snap.sub_step_total + post_count);
}
if self.lomac.is_some()
&& let Some(ref p) = self.progress
{
let sub = p.read().sub_step;
p.set_phase(StepPhase::LoMaC);
p.set_sub_step(sub + 1, p.read().sub_step_total);
}
if let Some(ref mut lomac) = self.lomac {
let t0 = std::time::Instant::now();
let gx = &products.acceleration.gx;
let gy = &products.acceleration.gy;
let gz = &products.acceleration.gz;
if let Some(ht) = self
.repr
.as_any()
.downcast_ref::<crate::tooling::core::algos::ht::HtTensor>()
{
if ht.can_materialize() {
lomac.advance_macroscopic(dt, gx, gy, gz);
let corrected = lomac.project_ht(ht);
let shape = ht.shape;
let corrected_snap = PhaseSpaceSnapshot {
data: corrected,
shape,
time: self.time,
};
self.repr = Box::new(
crate::tooling::core::algos::uniform::UniformGrid6D::from_snapshot(
corrected_snap,
self.domain.clone(),
),
);
if let Some(ref p) = self.progress {
self.repr.set_progress(p.clone());
}
} else {
tracing::warn!(
"LoMaC projection skipped: HT tensor too large to materialize ({} elements)",
ht.shape.iter().product::<usize>()
);
lomac.advance_macroscopic(dt, gx, gy, gz);
}
} else {
let snapshot = self.repr.to_snapshot(self.time).ok_or_else(|| {
anyhow::anyhow!("LoMaC dense path requires to_snapshot support")
})?;
let corrected = lomac.apply(dt, gx, gy, gz, &snapshot.data);
let corrected_snap = PhaseSpaceSnapshot {
data: corrected,
shape: snapshot.shape,
time: self.time,
};
self.repr = Box::new(
crate::tooling::core::algos::uniform::UniformGrid6D::from_snapshot(
corrected_snap,
self.domain.clone(),
),
);
if let Some(ref p) = self.progress {
self.repr.set_progress(p.clone());
}
}
timings.other_ms += t0.elapsed().as_secs_f64() * 1000.0;
}
self.time += dt;
self.step += 1;
if let Some(ref p) = self.progress {
let sub = p.read().sub_step;
p.set_phase(StepPhase::PostDensity);
p.set_sub_step(sub + 1, p.read().sub_step_total);
}
let t0 = std::time::Instant::now();
let (density, potential) = if let Some(ref lomac) = self.lomac {
let density = DensityField {
data: lomac.kfvs.state.iter().map(|m| m.density).collect(),
shape: lomac.spatial_shape,
};
let potential = self.poisson.solve(&density, self.g);
(density, potential)
} else {
(products.density, products.potential)
};
timings.density_ms += t0.elapsed().as_secs_f64() * 1000.0;
let dx = self.domain.dx();
let dx3 = dx[0] * dx[1] * dx[2];
self.cached_rho_max = Some(
density
.data
.par_iter()
.cloned()
.reduce(|| 0.0_f64, f64::max),
);
if let Some(ref p) = self.progress {
let sub = p.read().sub_step;
p.set_phase(StepPhase::Diagnostics);
p.set_sub_step(sub + 1, p.read().sub_step_total);
}
let t0 = std::time::Instant::now();
let diag = self.diagnostics.compute_with_density(
&*self.repr,
&density,
&potential,
self.time,
dx3,
);
timings.diagnostics_ms += t0.elapsed().as_secs_f64() * 1000.0;
self.last_step_timings = timings;
let accel = self.poisson.compute_acceleration(&potential);
self.cached_density = Some(density);
self.cached_potential = Some(potential);
self.cached_acceleration = Some(accel);
if let Some(ref p) = self.progress {
p.set_phase(StepPhase::StepComplete);
}
Ok(self.exit_evaluator.check(&diag))
}
pub fn set_progress(&mut self, p: Arc<StepProgress>) {
self.integrator.set_progress(p.clone());
self.repr.set_progress(p.clone());
self.poisson.set_progress(p.clone());
if let Some(ref mut lomac) = self.lomac {
lomac.set_progress(p.clone());
}
self.progress = Some(p);
}
pub fn current_time(&self) -> f64 {
self.time
}
}
pub struct SimulationBuilder {
domain: Option<Domain>,
repr: Option<Box<dyn PhaseSpaceRepr>>,
poisson: Option<Box<dyn PoissonSolver>>,
advector: Option<Box<dyn Advector>>,
integrator: Option<Box<dyn TimeIntegrator>>,
opts: Option<OptionalParams>,
ic: Option<PhaseSpaceSnapshot>,
t_final: Option<Decimal>,
output_interval: Option<Decimal>,
energy_tolerance: Option<Decimal>,
g: Option<Decimal>,
enable_lomac: bool,
}
impl SimulationBuilder {
pub fn new() -> Self {
Self {
domain: None,
repr: None,
poisson: None,
advector: None,
integrator: None,
opts: None,
ic: None,
t_final: None,
output_interval: None,
energy_tolerance: None,
g: None,
enable_lomac: false,
}
}
pub fn domain(mut self, d: Domain) -> Self {
self.domain = Some(d);
self
}
pub fn representation(mut self, r: impl PhaseSpaceRepr + 'static) -> Self {
self.repr = Some(Box::new(r));
self
}
pub fn representation_boxed(mut self, r: Box<dyn PhaseSpaceRepr>) -> Self {
self.repr = Some(r);
self
}
pub fn poisson_solver(mut self, p: impl PoissonSolver + 'static) -> Self {
self.poisson = Some(Box::new(p));
self
}
pub fn advector(mut self, a: impl Advector + 'static) -> Self {
self.advector = Some(Box::new(a));
self
}
pub fn integrator(mut self, i: impl TimeIntegrator + 'static) -> Self {
self.integrator = Some(Box::new(i));
self
}
pub fn poisson_solver_boxed(mut self, p: Box<dyn PoissonSolver>) -> Self {
self.poisson = Some(p);
self
}
pub fn integrator_boxed(mut self, i: Box<dyn TimeIntegrator>) -> Self {
self.integrator = Some(i);
self
}
pub fn initial_conditions(mut self, ic: PhaseSpaceSnapshot) -> Self {
self.ic = Some(ic);
self
}
pub fn time_final(mut self, t: f64) -> Self {
self.t_final = Some(f64_to_decimal(t, "time_final"));
self
}
pub fn time_final_decimal(mut self, t: Decimal) -> Self {
self.t_final = Some(t);
self
}
pub fn output_interval(mut self, dt: f64) -> Self {
self.output_interval = Some(f64_to_decimal(dt, "output_interval"));
self
}
pub fn output_interval_decimal(mut self, dt: Decimal) -> Self {
self.output_interval = Some(dt);
self
}
pub fn exit_on_energy_drift(mut self, tol: f64) -> Self {
self.energy_tolerance = Some(f64_to_decimal(tol, "exit_on_energy_drift"));
self
}
pub fn exit_on_energy_drift_decimal(mut self, tol: Decimal) -> Self {
self.energy_tolerance = Some(tol);
self
}
pub fn gravitational_constant(mut self, g: f64) -> Self {
self.g = Some(f64_to_decimal(g, "gravitational_constant"));
self
}
pub fn gravitational_constant_decimal(mut self, g: Decimal) -> Self {
self.g = Some(g);
self
}
pub fn lomac(mut self, enable: bool) -> Self {
self.enable_lomac = enable;
self
}
pub fn cfl_factor(mut self, cfl: f64) -> Self {
let mut opts = self.opts.unwrap_or_default();
opts.cfl_factor = f64_to_decimal(cfl, "cfl_factor");
self.opts = Some(opts);
self
}
pub fn cfl_factor_decimal(mut self, cfl: Decimal) -> Self {
let mut opts = self.opts.unwrap_or_default();
opts.cfl_factor = cfl;
self.opts = Some(opts);
self
}
pub fn build(self) -> anyhow::Result<Simulation> {
use crate::tooling::core::algos::uniform::UniformGrid6D;
use rust_decimal::prelude::ToPrimitive;
let domain = self
.domain
.ok_or_else(|| anyhow::anyhow!("domain not set"))?;
let poisson = self
.poisson
.ok_or_else(|| anyhow::anyhow!("poisson_solver not set"))?;
let advector = self
.advector
.ok_or_else(|| anyhow::anyhow!("advector not set"))?;
let integrator = self
.integrator
.ok_or_else(|| anyhow::anyhow!("integrator not set"))?;
let t_final_dec = self.t_final.unwrap_or(domain.time_range.t_final);
let t_final = t_final_dec.to_f64().unwrap_or(1.0);
let g = self.g.unwrap_or(Decimal::ONE).to_f64().unwrap_or(1.0);
let repr: Box<dyn PhaseSpaceRepr> = if let Some(ic) = self.ic {
Box::new(UniformGrid6D::from_snapshot(ic, domain.clone()))
} else if let Some(r) = self.repr {
r
} else {
anyhow::bail!("either initial_conditions or representation must be set");
};
let dx = domain.dx();
let dx3 = dx[0] * dx[1] * dx[2];
let density = repr.compute_density();
let potential = poisson.solve(&density, g);
let mut diagnostics = Diagnostics {
history: Vec::new(),
};
let initial_diag = diagnostics.compute(&*repr, &potential, 0.0, dx3);
let mut exit_evaluator = ExitEvaluator::new(initial_diag);
exit_evaluator.add_condition(Box::new(TimeLimitCondition { t_final }));
if let Some(tol_dec) = self.energy_tolerance {
use crate::tooling::core::conditions::EnergyDriftCondition;
let tol = tol_dec.to_f64().unwrap_or(1e-6);
exit_evaluator.add_condition(Box::new(EnergyDriftCondition { tolerance: tol }));
}
let lomac = if self.enable_lomac {
let dv = domain.dv();
let lv = domain.lv();
let v_min = [-lv[0], -lv[1], -lv[2]];
let spatial_shape = [
domain.spatial_res.x1 as usize,
domain.spatial_res.x2 as usize,
domain.spatial_res.x3 as usize,
];
let velocity_shape = [
domain.velocity_res.v1 as usize,
domain.velocity_res.v2 as usize,
domain.velocity_res.v3 as usize,
];
let mut lom = LoMaC::new(spatial_shape, velocity_shape, dx, dv, v_min);
if let Some(ht) = repr
.as_any()
.downcast_ref::<crate::tooling::core::algos::ht::HtTensor>()
{
lom.initialize_from_ht(ht);
} else {
let snapshot = repr.to_snapshot(0.0).ok_or_else(|| {
anyhow::anyhow!("LoMaC initialization requires to_snapshot support")
})?;
lom.initialize_from_kinetic(&snapshot.data);
}
Some(lom)
} else {
None
};
Ok(Simulation {
domain,
repr,
poisson,
advector,
integrator,
diagnostics,
io: IOManager::new("output", OutputFormat::Binary),
exit_evaluator,
opts: self.opts.unwrap_or_default(),
lomac,
g,
time: 0.0,
step: 0,
start_time: std::time::Instant::now(),
cached_rho_max: None,
last_step_timings: StepTimings::default(),
cached_density: None,
cached_potential: None,
cached_acceleration: None,
progress: None,
})
}
}