From 7bc7b6fb96003e43b5c7afa980e39c90703fbe16 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Juli=C3=A1n=20D=2E=20Ot=C3=A1lvaro?= Date: Thu, 20 Aug 2026 17:22:08 +0100 Subject: [PATCH] fix: accept solver stop times already reached --- src/dsl/native.rs | 35 ++----- src/simulator/equation/ode/mod.rs | 56 ++--------- tests/bolus_reinit_stop_time.rs | 152 ++++++++++++++++++++++++++++++ 3 files changed, 171 insertions(+), 72 deletions(-) create mode 100644 tests/bolus_reinit_stop_time.rs diff --git a/src/dsl/native.rs b/src/dsl/native.rs index f2fc1955..6ccfecf9 100644 --- a/src/dsl/native.rs +++ b/src/dsl/native.rs @@ -234,7 +234,7 @@ impl FunctionSession for NativeFunctionSession<'_> { )) })?; - function(time, states, params, covariates, routes, derived, out); + unsafe { function(time, states, params, covariates, routes, derived, out) }; Ok(()) } } @@ -1472,32 +1472,15 @@ impl NativeOdeModel { OdeSolverError::StopTimeAtCurrentTime, )) => { solver.problem().eqn.set_left_continuity_time(None); - let state_t = solver.state().t; - let stop_reached = crate::simulator::equation::ode::stop_time_reached( - stop_time, state_t, - ); - - if stop_reached { - if is_infusion_boundary { - pending_reinit = true; - } - // The requested stop is the current time within - // a small relative tolerance. If it is an - // infusion boundary before the next subject - // event, keep integrating toward the event; - // break only when the reached stop is the - // event time itself. - if stop_time < next_event_time { - continue; - } - break; + // diffsol has already applied the solver's stop-time + // tolerance; an actually earlier stop uses a distinct error. + if is_infusion_boundary { + pending_reinit = true; + } + if stop_time < next_event_time { + continue; } - return Err(PharmsolError::from_solver_error( - diffsol::error::DiffsolError::OdeSolverError( - OdeSolverError::StopTimeAtCurrentTime, - ), - stop_time, - )); + break; } Err(err) => { solver.problem().eqn.set_left_continuity_time(None); diff --git a/src/simulator/equation/ode/mod.rs b/src/simulator/equation/ode/mod.rs index c90128ab..fdae500b 100644 --- a/src/simulator/equation/ode/mod.rs +++ b/src/simulator/equation/ode/mod.rs @@ -585,24 +585,6 @@ where state.dy.copy_from(dy_scratch); } -/// Whether a requested solver stop is effectively at the current state time. -/// -/// diffsol reports `StopTimeAtCurrentTime` not only for a stop exactly at the -/// current time, but also when its internal state time has landed a few ULPs -/// past the requested stop (adaptive steps may end slightly beyond a stop). -/// Dense output grids built with floating-point arithmetic routinely place -/// requested times a few ULPs away from event times (e.g. a `t += dt` -/// accumulation puts a point ~16 ULPs after a bolus at `t = 12`), so accept a -/// stop within a small relative tolerance of the current time instead of -/// erroring. The tolerance stays far below any meaningful time difference: -/// ~64-128 ULPs of the current time, i.e. at most ~1e-13 at `t = 12`. -/// -/// Shared with the DSL/JIT ODE path ([`crate::dsl::native::NativeOdeModel`]). -pub(crate) fn stop_time_reached(stop_time: f64, state_t: f64) -> bool { - let tolerance = f64::EPSILON * state_t.abs().max(1.0) * 64.0; - (stop_time - state_t).abs() <= tolerance -} - impl ODE { /// Generic event-loop runner, parameterized over the concrete solver type. #[allow(clippy::too_many_arguments)] @@ -782,33 +764,15 @@ impl ODE { OdeSolverError::StopTimeAtCurrentTime, )) => { solver.problem().eqn.set_left_continuity_time(None); - let state_t = solver.state().t; - let stop_reached = stop_time_reached(stop_time, state_t); - - if stop_reached { - if is_infusion_boundary { - pending_reinit = true; - } - // The requested stop is the current time within - // a small relative tolerance. If it is an - // infusion boundary before the next subject - // event, keep integrating toward the event; - // break only when the reached stop is the - // event time itself. Breaking early would skip - // the remaining interval and leave the solver - // state at the boundary when the observation is - // evaluated. - if stop_time < next_event_time { - continue; - } - break; + // diffsol has already applied the solver's stop-time + // tolerance; an actually earlier stop uses a distinct error. + if is_infusion_boundary { + pending_reinit = true; + } + if stop_time < next_event_time { + continue; } - return Err(PharmsolError::from_solver_error( - diffsol::error::DiffsolError::OdeSolverError( - OdeSolverError::StopTimeAtCurrentTime, - ), - stop_time, - )); + break; } Err(err) => { solver.problem().eqn.set_left_continuity_time(None); @@ -1349,8 +1313,8 @@ mod tests { // The infusion ends exactly one ULP after the observation at t = 10. // After landing on the observation stop the solver is already within // diffsol's round-off of the end boundary, so `set_stop_time` reports - // `StopTimeAtCurrentTime` and the loop accepts it through the - // same-ULP check. The reached boundary is *before* the observation at + // `StopTimeAtCurrentTime`, which confirms the boundary is reached. + // The reached boundary is *before* the observation at // t = 20, so the event loop must keep integrating toward it; breaking // early would evaluate the observation with the state frozen at the // boundary and miss the exponential decay. diff --git a/tests/bolus_reinit_stop_time.rs b/tests/bolus_reinit_stop_time.rs new file mode 100644 index 00000000..bb6a078b --- /dev/null +++ b/tests/bolus_reinit_stop_time.rs @@ -0,0 +1,152 @@ +//! Regression coverage for solver restarts after bolus state changes. +//! +//! A TSIT45 restart can land a few ULPs short of an observation time while +//! diffsol still correctly reports that the requested stop was reached. The +//! event loop must accept diffsol's `StopTimeAtCurrentTime` response rather than +//! requesting the same stop again and turning it into a simulation error. + +use pharmsol::prelude::*; + +#[cfg(feature = "dsl-jit")] +use pharmsol::dsl::{ + compile_module_source_to_runtime, CompiledRuntimeModel, RuntimeCompilationTarget, +}; + +const OBSERVATION_TIMES: [f64; 15] = [ + 0.0, + 0.35, + 0.516666666666667, + 0.983333333333333, + 1.48333333333333, + 2.0, + 2.5, + 3.0, + 4.0, + 4.98333333333333, + 6.98333333333333, + 7.98333333333333, + 10.0, + 11.0, + 12.0, +]; + +const PARAMETERS: [(&str, f64); 5] = [ + ("ka", 3.6156922578811646), + ("cl0", 1.0289061069488525), + ("vc0", 187.13204860687256), + ("q0", 2.4602913856506348), + ("vp0", 58.32162380218506), +]; + +fn subject_with_bolus_history() -> Subject { + let mut builder = Subject::builder("g34"); + for dose_index in -16..=0 { + builder = builder.bolus(f64::from(dose_index) * 12.0, 750.0, "input_1"); + } + for time in OBSERVATION_TIMES { + builder = builder.missing_observation(time, "outeq_1"); + } + builder.build() +} + +fn closure_model(solver: OdeSolver) -> equation::ODE { + equation::ODE::new( + |x, p, _t, dx, b, rateiv, _cov| { + fetch_params!(p, ka, cl0, vc0, q0, vp0); + let ke = cl0 / vc0; + let k23 = q0 / vc0; + let k32 = q0 / vp0; + + dx[0] = b[0] - x[0] * ka; + dx[1] = rateiv[0] + x[0] * ka + x[2] * k32 - x[1] * (ke + k23); + dx[2] = x[1] * k23 - x[2] * k32; + }, + |_p, _t, _cov| lag! {}, + |_p, _t, _cov| fa! {}, + |_p, _t, _cov, _x| {}, + |x, p, _t, _cov, y| { + fetch_params!(p, _ka, _cl0, vc0, _q0, _vp0); + y[0] = x[1] / vc0; + }, + ) + .with_nstates(3) + .with_ndrugs(1) + .with_nout(1) + .with_solver(solver) + .with_metadata( + equation::metadata::new("bolus_reinit_stop_time") + .parameters(["ka", "cl0", "vc0", "q0", "vp0"]) + .states(["x1", "x2", "x3"]) + .outputs(["outeq_1"]) + .routes([equation::Route::bolus("input_1") + .to_state("x1") + .expect_explicit_input()]), + ) + .expect("regression model metadata should validate") +} + +#[test] +fn closure_tsit45_accepts_reached_stop_after_bolus_restarts( +) -> Result<(), Box> { + let model = closure_model(OdeSolver::ExplicitRk(ExplicitRkTableau::Tsit45)); + let parameters = Parameters::with_model(&model, PARAMETERS)?; + let predictions = + model.estimate_predictions_dense(&subject_with_bolus_history(), parameters.as_slice())?; + + assert_eq!(predictions.predictions().len(), OBSERVATION_TIMES.len()); + Ok(()) +} + +#[cfg(feature = "dsl-jit")] +const DSL_MODEL: &str = r#" +name = bolus_reinit_stop_time +kind = ode +params = ka, cl0, vc0, q0, vp0 +states = x1, x2, x3 +outputs = outeq_1 + +bolus(input_1) -> x1 +infusion(input_1) -> x2 + +cl = cl0 +vc = vc0 +q = q0 +vp = vp0 +ke = cl / vc +k23 = q / vc +k32 = q / vp + +dx(x1) = -(x1 * ka) +dx(x2) = x1 * ka + (x3 * k32) - (x2 * (ke + k23)) +dx(x3) = x2 * k23 - (x3 * k32) + +out(outeq_1) = x2 / vc +"#; + +#[test] +#[cfg(feature = "dsl-jit")] +fn jit_tsit45_accepts_reached_stop_after_bolus_restarts() -> Result<(), Box> +{ + let compiled = compile_module_source_to_runtime( + DSL_MODEL, + Some("bolus_reinit_stop_time"), + RuntimeCompilationTarget::Jit, + |_, _| {}, + )?; + let model = match compiled { + CompiledRuntimeModel::Ode(model) => CompiledRuntimeModel::Ode( + model.with_solver(OdeSolver::ExplicitRk(ExplicitRkTableau::Tsit45)), + ), + _ => return Err("expected an ODE model".into()), + }; + let parameters = Parameters::with_model(&model, PARAMETERS)?; + + let predictions = match &model { + CompiledRuntimeModel::Ode(model) => model + .estimate_predictions_dense(&subject_with_bolus_history(), parameters.as_slice())?, + _ => unreachable!(), + }; + + assert_eq!(predictions.predictions().len(), OBSERVATION_TIMES.len()); + Ok(()) +}