use std::sync::Arc;
use std::time::Instant;
use faer::Mat;
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::{
self, BugConfig, LEAF_PARENT, conservative_correction, k_step_leaf,
representative_accelerations, representative_velocities, sample_aug_displacements,
update_transfer,
};
pub struct ParallelBugConfig {
pub tolerance: f64,
pub max_rank: usize,
pub conservative: bool,
pub rank_increase: usize,
pub error_rejection: bool,
pub rejection_tolerance: f64,
}
impl Default for ParallelBugConfig {
fn default() -> Self {
Self {
tolerance: 1e-8,
max_rank: 50,
conservative: false,
rank_increase: 2,
error_rejection: false,
rejection_tolerance: 1e-4,
}
}
}
pub struct ParallelBugIntegrator {
pub config: ParallelBugConfig,
pub g: f64,
last_timings: StepTimings,
progress: Option<Arc<StepProgress>>,
}
impl ParallelBugIntegrator {
pub fn new(g: f64, config: ParallelBugConfig) -> 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 parallel_k_steps(
ht: &HtTensor,
accel: &AccelerationField,
dt_drift: f64,
dt_kick: f64,
config: &ParallelBugConfig,
) -> Vec<(Mat<f64>, Mat<f64>)> {
use rayon::prelude::*;
let dims: Vec<usize> = (0..6).collect();
dims.par_iter()
.map(|&d| {
if d < 3 {
let reps = representative_velocities(ht, d + 3);
let primary = if reps.is_empty() {
0.0
} else {
reps.iter().sum::<f64>() / reps.len() as f64 * dt_drift
};
let aug = sample_aug_displacements(&reps, dt_drift, config.rank_increase);
k_step_leaf(ht, d, primary, &aug, config.max_rank, config.tolerance)
} else {
let reps = representative_accelerations(ht, d - 3, accel);
let primary = if reps.is_empty() {
0.0
} else {
reps.iter().sum::<f64>() / reps.len() as f64 * dt_kick
};
let aug = sample_aug_displacements(&reps, dt_kick, config.rank_increase);
k_step_leaf(ht, d, primary, &aug, config.max_rank, config.tolerance)
}
})
.collect()
}
fn apply_k_steps(ht: &mut HtTensor, results: Vec<(Mat<f64>, Mat<f64>)>) {
let apply_leaf = |ht: &mut HtTensor, d: usize, result: &(Mat<f64>, Mat<f64>)| {
*ht.leaf_frame_mut(d) = result.0.clone();
update_transfer(ht, d, &result.1);
};
apply_leaf(ht, 0, &results[0]);
apply_leaf(ht, 3, &results[3]);
apply_leaf(ht, 1, &results[1]);
apply_leaf(ht, 2, &results[2]);
apply_leaf(ht, 4, &results[4]);
apply_leaf(ht, 5, &results[5]);
}
fn parallel_bug_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, "Parallel BUG requires HtTensor");
return;
};
let density_before = if self.config.conservative {
Some(ht.compute_density())
} else {
None
};
helpers::report_phase!(self.progress, StepPhase::BugKStep, 0, 5);
helpers::time_ms!(timings, drift_ms, {
let results = Self::parallel_drift_k_steps(ht, dt / 2.0, &self.config);
for (d, (new_frame, r_mat)) in results.into_iter().enumerate() {
*ht.leaf_frame_mut(d) = new_frame;
update_transfer(ht, d, &r_mat);
}
});
helpers::report_phase!(self.progress, StepPhase::BugLStep, 1, 5);
let (_, _, accel) = helpers::time_ms!(
timings,
poisson_ms,
helpers::solve_poisson(ht, solver, self.g)
);
helpers::report_phase!(self.progress, StepPhase::BugLStep, 2, 5);
helpers::time_ms!(timings, kick_ms, {
let results = Self::parallel_kick_k_steps(ht, &accel, dt, &self.config);
for (i, (new_frame, r_mat)) in results.into_iter().enumerate() {
let d = i + 3;
*ht.leaf_frame_mut(d) = new_frame;
update_transfer(ht, d, &r_mat);
}
});
helpers::report_phase!(self.progress, StepPhase::BugLStep, 3, 5);
helpers::time_ms!(timings, drift_ms, {
let results = Self::parallel_drift_k_steps(ht, dt / 2.0, &self.config);
for (d, (new_frame, r_mat)) in results.into_iter().enumerate() {
*ht.leaf_frame_mut(d) = new_frame;
update_transfer(ht, d, &r_mat);
}
});
helpers::report_phase!(self.progress, StepPhase::BugSStep, 4, 5);
ht.truncate(self.config.tolerance);
if let Some(ref dens) = density_before {
conservative_correction(ht, dens);
}
}
fn parallel_drift_k_steps(
ht: &HtTensor,
dt: f64,
config: &ParallelBugConfig,
) -> Vec<(Mat<f64>, Mat<f64>)> {
use rayon::prelude::*;
(0..3usize)
.into_par_iter()
.map(|d| {
let reps = representative_velocities(ht, d + 3);
let primary = if reps.is_empty() {
0.0
} else {
reps.iter().sum::<f64>() / reps.len() as f64 * dt
};
let aug = sample_aug_displacements(&reps, dt, config.rank_increase);
k_step_leaf(ht, d, primary, &aug, config.max_rank, config.tolerance)
})
.collect()
}
fn parallel_kick_k_steps(
ht: &HtTensor,
accel: &AccelerationField,
dt: f64,
config: &ParallelBugConfig,
) -> Vec<(Mat<f64>, Mat<f64>)> {
use rayon::prelude::*;
(3..6usize)
.into_par_iter()
.map(|d| {
let reps = representative_accelerations(ht, d - 3, accel);
let primary = if reps.is_empty() {
0.0
} else {
reps.iter().sum::<f64>() / reps.len() as f64 * dt
};
let aug = sample_aug_displacements(&reps, dt, config.rank_increase);
k_step_leaf(ht, d, primary, &aug, config.max_rank, config.tolerance)
})
.collect()
}
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 ParallelBugIntegrator {
fn advance(
&mut self,
repr: &mut dyn PhaseSpaceRepr,
solver: &dyn PoissonSolver,
advector: &dyn Advector,
dt: f64,
) -> Result<StepProducts, CausticError> {
let _span = tracing::info_span!("parallel_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.parallel_bug_step(repr, solver, dt, &mut timings);
} else {
self.strang_fallback(repr, solver, advector, dt, &mut timings);
}
helpers::report_phase!(self.progress, StepPhase::StepComplete, 5, 5);
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 parallel_bug_config_defaults() {
let cfg = ParallelBugConfig::default();
assert_eq!(cfg.max_rank, 50);
assert!(!cfg.conservative);
assert!(!cfg.error_rejection);
assert_eq!(cfg.rank_increase, 2);
}
}