use crate::{
error::DiffsolError,
ode_equations::augmented_channel,
ode_solver::method::write_state_out,
ode_solver::solution::{Solution, SolutionMode},
ode_solver_error, AugmentedOdeSolverMethod, Context, DefaultDenseMatrix, DefaultSolver,
DenseMatrix, Matrix, MatrixCommon, NonLinearOp, NonLinearOpJacobian, NonLinearOpSens,
OdeEquationsImplicitSens, OdeSolverProblem, OdeSolverStopReason, Op, SensEquations, StateRef,
Vector, VectorViewMut,
};
pub trait SensitivitiesOdeSolverMethod<'a, Eqn>:
AugmentedOdeSolverMethod<'a, Eqn, SensEquations<'a, Eqn>>
where
Eqn: OdeEquationsImplicitSens + 'a,
{
fn solve_soln_sensitivities(mut self, soln: &mut Solution<Eqn::V>) -> Result<Self, DiffsolError>
where
Eqn::V: DefaultDenseMatrix,
Self: Sized,
{
if self.problem().integrate_out {
return Err(ode_solver_error!(
Other,
"Cannot integrate out when solving for sensitivities"
));
}
let start_col = match soln.mode {
SolutionMode::Tevals(start_col) => start_col,
SolutionMode::Tfinal(_) => {
return Err(ode_solver_error!(
Other,
"solve_soln_sensitivities requires Solution::new_dense"
));
}
};
let ctx = self.problem().context().clone();
let nrows = self
.problem()
.eqn
.out()
.map(|out| out.nout())
.unwrap_or_else(|| self.problem().eqn.rhs().nout());
let nstates = self.problem().eqn.rhs().nstates();
let nparams = self.problem().eqn.rhs().nparams();
let nout = self.problem().eqn.out().map(|out| out.nout()).unwrap_or(0);
soln.ensure_sens_allocation(&ctx, nrows, nout, nstates, nparams)?;
let mut out_sens = match self.problem().eqn.out() {
Some(out) => <Eqn as Op>::M::new_from_sparsity(nout, nparams, out.sens_sparsity(), ctx),
None => <Eqn as Op>::M::zeros(0, 0, ctx),
};
let (stop_reason, col) = solve_dense_sensitivities(
&mut soln.ys,
&mut soln.y_sens,
&soln.ts,
&mut soln.tmp_nout,
&mut soln.tmp_nstates,
&mut soln.tmp_nsens,
&mut soln.tmp_nsens_out,
&mut out_sens,
&mut self,
start_col,
)?;
soln.stop_reason = Some(stop_reason);
soln.mode = SolutionMode::Tevals(col);
Ok(self)
}
#[allow(clippy::type_complexity)]
fn solve_dense_sensitivities(
&mut self,
t_eval: &[Eqn::T],
) -> Result<
(
<Eqn::V as DefaultDenseMatrix>::M,
Vec<<Eqn::V as DefaultDenseMatrix>::M>,
OdeSolverStopReason<Eqn::T>,
),
DiffsolError,
>
where
Eqn: OdeEquationsImplicitSens,
Eqn::V: DefaultDenseMatrix,
Eqn::M: DefaultSolver,
Self: Sized,
{
if self.problem().integrate_out {
return Err(ode_solver_error!(
Other,
"Cannot integrate out when solving for sensitivities"
));
}
let nrows = if let Some(out) = self.problem().eqn.out() {
out.nout()
} else {
self.problem().eqn.rhs().nout()
};
let nstates = self.problem().eqn.rhs().nstates();
let nparams = self.problem().eqn.rhs().nparams();
let ctx = self.problem().context().clone();
let mut ret = ctx.dense_mat_zeros::<Eqn::V>(nrows, t_eval.len());
let mut ret_sens = vec![ctx.dense_mat_zeros::<Eqn::V>(nrows, t_eval.len()); nparams];
let mut tmp_nout = Eqn::V::zeros(
self.problem().eqn.out().map(|out| out.nout()).unwrap_or(0),
ctx.clone(),
);
let mut tmp_nstates = Eqn::V::zeros(nstates, ctx.clone());
let aug_ctx = ctx.clone_with_nbatch(ctx.nbatch() * nparams)?;
let mut tmp_nsens = Eqn::V::zeros(nstates, aug_ctx.clone());
let mut tmp_nsens_out = Eqn::V::zeros(
self.problem().eqn.out().map(|out| out.nout()).unwrap_or(0),
aug_ctx.clone(),
);
let nout_sens = self.problem().eqn.out().map(|out| out.nout()).unwrap_or(0);
let mut tmp_out_sens = match self.problem().eqn.out() {
Some(out) => {
<Eqn as Op>::M::new_from_sparsity(nout_sens, nparams, out.sens_sparsity(), ctx)
}
None => <Eqn as Op>::M::zeros(0, 0, ctx),
};
let t0 = self.state().t;
if t_eval.windows(2).any(|w| w[0] > w[1] || w[0] < t0) {
return Err(ode_solver_error!(InvalidTEval));
}
let (stop_reason, col) = solve_dense_sensitivities_auto_reset(
&mut ret,
&mut ret_sens,
t_eval,
&mut tmp_nout,
&mut tmp_nstates,
&mut tmp_nsens,
&mut tmp_nsens_out,
&mut tmp_out_sens,
self,
0,
)?;
if let OdeSolverStopReason::RootFound(_, _) = stop_reason {
if col < t_eval.len() {
write_state_out(self.problem(), &self.state(), &mut ret, col, &mut tmp_nout);
write_state_sens_out(
self.problem(),
&self.state(),
&mut ret_sens,
col,
&mut tmp_nout,
&mut tmp_nsens_out,
&mut tmp_out_sens,
);
if col + 1 < ret.ncols() {
ret.resize_cols(col + 1);
for rs in &mut ret_sens {
rs.resize_cols(col + 1);
}
}
}
}
Ok((ret, ret_sens, stop_reason))
}
}
#[allow(clippy::too_many_arguments)]
fn solve_dense_sensitivities<'a, Eqn, S>(
ret: &mut <Eqn::V as DefaultDenseMatrix>::M,
ret_sens: &mut [<Eqn::V as DefaultDenseMatrix>::M],
t_eval: &[Eqn::T],
tmp_nout: &mut Eqn::V,
tmp_nstates: &mut Eqn::V,
tmp_nsens: &mut Eqn::V,
tmp_nsens_out: &mut Eqn::V,
tmp_out_sens: &mut <Eqn as Op>::M,
s: &mut S,
start_col: usize,
) -> Result<(OdeSolverStopReason<Eqn::T>, usize), DiffsolError>
where
Eqn: OdeEquationsImplicitSens + 'a,
Eqn::V: DefaultDenseMatrix,
S: SensitivitiesOdeSolverMethod<'a, Eqn>,
{
s.set_stop_time(t_eval[t_eval.len() - 1])?;
let mut stop_reason: OdeSolverStopReason<Eqn::T>;
let mut col = start_col;
loop {
stop_reason = s.step()?;
let t_current = if let OdeSolverStopReason::RootFound(t, _) = stop_reason {
t
} else {
s.state().t
};
while col < t_eval.len() && t_eval[col] <= t_current {
dense_write_out_sensitivities(
s,
ret,
ret_sens,
t_eval[col],
col,
tmp_nout,
tmp_nstates,
tmp_nsens,
tmp_nsens_out,
tmp_out_sens,
)?;
col += 1;
}
match stop_reason {
OdeSolverStopReason::InternalTimestep => {}
OdeSolverStopReason::TstopReached => {
assert!(
col == t_eval.len(),
"Solver reached stop time before consuming all t_eval points, this should not happen"
);
break;
}
OdeSolverStopReason::RootFound(t_root, _) => {
s.state_mut_back(t_root)?;
break;
}
}
}
Ok((stop_reason, col))
}
#[allow(clippy::too_many_arguments)]
fn solve_dense_sensitivities_auto_reset<'a, Eqn, S>(
ret: &mut <Eqn::V as DefaultDenseMatrix>::M,
ret_sens: &mut [<Eqn::V as DefaultDenseMatrix>::M],
t_eval: &[Eqn::T],
tmp_nout: &mut Eqn::V,
tmp_nstates: &mut Eqn::V,
tmp_nsens: &mut Eqn::V,
tmp_nsens_out: &mut Eqn::V,
tmp_out_sens: &mut <Eqn as Op>::M,
s: &mut S,
start_col: usize,
) -> Result<(OdeSolverStopReason<Eqn::T>, usize), DiffsolError>
where
Eqn: OdeEquationsImplicitSens + 'a,
Eqn::V: DefaultDenseMatrix,
Eqn::M: DefaultSolver,
S: SensitivitiesOdeSolverMethod<'a, Eqn>,
{
s.set_stop_time(t_eval[t_eval.len() - 1])?;
let has_reset = s.problem().eqn.reset().is_some();
let mut stop_reason: OdeSolverStopReason<Eqn::T>;
let mut col = start_col;
loop {
stop_reason = s.step()?;
match stop_reason {
OdeSolverStopReason::InternalTimestep => {
while col < t_eval.len() && t_eval[col] <= s.state().t {
dense_write_out_sensitivities(
s,
ret,
ret_sens,
t_eval[col],
col,
tmp_nout,
tmp_nstates,
tmp_nsens,
tmp_nsens_out,
tmp_out_sens,
)?;
col += 1;
}
}
OdeSolverStopReason::TstopReached => {
while col < t_eval.len() && t_eval[col] <= s.state().t {
dense_write_out_sensitivities(
s,
ret,
ret_sens,
t_eval[col],
col,
tmp_nout,
tmp_nstates,
tmp_nsens,
tmp_nsens_out,
tmp_out_sens,
)?;
col += 1;
}
assert!(
col == t_eval.len(),
"Solver reached stop time before consuming all t_eval points, this should not happen"
);
break;
}
OdeSolverStopReason::RootFound(t_root, root_idx) => {
while col < t_eval.len() && t_eval[col] <= t_root {
dense_write_out_sensitivities(
s,
ret,
ret_sens,
t_eval[col],
col,
tmp_nout,
tmp_nstates,
tmp_nsens,
tmp_nsens_out,
tmp_out_sens,
)?;
col += 1;
}
s.state_mut_back(t_root)?;
if has_reset {
s.apply_reset_with_sens(root_idx)?;
if s.state().t < t_eval[t_eval.len() - 1] {
s.set_stop_time(t_eval[t_eval.len() - 1])?;
} else {
stop_reason = OdeSolverStopReason::TstopReached;
break;
}
} else {
stop_reason = OdeSolverStopReason::RootFound(t_root, root_idx);
break;
}
}
}
}
Ok((stop_reason, col))
}
#[allow(clippy::too_many_arguments)]
fn dense_write_out_sensitivities<'a, Eqn, S>(
s: &S,
ret: &mut <Eqn::V as DefaultDenseMatrix>::M,
ret_sens: &mut [<Eqn::V as DefaultDenseMatrix>::M],
t: Eqn::T,
col: usize,
tmp_nout: &mut Eqn::V,
tmp_nstates: &mut Eqn::V,
tmp_nsens: &mut Eqn::V,
tmp_nsens_out: &mut Eqn::V,
tmp_out_sens: &mut <Eqn as Op>::M,
) -> Result<(), DiffsolError>
where
Eqn: OdeEquationsImplicitSens + 'a,
Eqn::V: DefaultDenseMatrix,
S: SensitivitiesOdeSolverMethod<'a, Eqn>,
{
s.interpolate_inplace(t, tmp_nstates)?;
s.interpolate_sens_inplace(t, tmp_nsens)?;
let nparams = ret_sens.len();
if let Some(out) = s.problem().eqn.out() {
out.call_inplace(tmp_nstates, t, tmp_nout);
ret.column_mut(col).copy_from(tmp_nout);
out.jac_mul_inplace(tmp_nstates, t, tmp_nsens, tmp_nsens_out);
out.sens_inplace(tmp_nstates, t, tmp_out_sens);
for (j, sens) in ret_sens.iter_mut().enumerate() {
augmented_channel(tmp_nsens_out, nparams, j, tmp_nout);
tmp_out_sens.add_column_to_vector(j, tmp_nout);
sens.column_mut(col).copy_from(&*tmp_nout);
}
} else {
ret.column_mut(col).copy_from(tmp_nstates);
for (j, sens) in ret_sens.iter_mut().enumerate() {
augmented_channel(tmp_nsens, nparams, j, tmp_nstates);
sens.column_mut(col).copy_from(&*tmp_nstates);
}
}
Ok(())
}
pub(crate) fn write_state_sens_out<Eqn>(
problem: &OdeSolverProblem<Eqn>,
state: &StateRef<'_, Eqn::V>,
ret_sens: &mut [<Eqn::V as DefaultDenseMatrix>::M],
col: usize,
tmp_nout: &mut Eqn::V,
tmp_nsens_out: &mut Eqn::V,
tmp_out_sens: &mut <Eqn as Op>::M,
) where
Eqn: OdeEquationsImplicitSens,
Eqn::V: DefaultDenseMatrix,
{
let nparams = ret_sens.len();
if let Some(out) = problem.eqn.out() {
out.jac_mul_inplace(state.y, state.t, state.s, tmp_nsens_out);
out.sens_inplace(state.y, state.t, tmp_out_sens);
for (j, sens) in ret_sens.iter_mut().enumerate() {
augmented_channel(tmp_nsens_out, nparams, j, tmp_nout);
tmp_out_sens.add_column_to_vector(j, tmp_nout);
sens.column_mut(col).copy_from(&*tmp_nout);
}
} else {
let mut tmp = Eqn::V::zeros(state.s.len(), problem.context().clone());
for (j, sens) in ret_sens.iter_mut().enumerate() {
augmented_channel(state.s, nparams, j, &mut tmp);
sens.column_mut(col).copy_from(&tmp);
}
}
}