From ab43646b35a9d51248ca9b4ca91d6b5a85f5847a Mon Sep 17 00:00:00 2001 From: Markus <66058642+mhovd@users.noreply.github.com> Date: Fri, 28 Aug 2026 11:29:56 +0200 Subject: [PATCH 1/3] WIP --- src/bestdose/cost.rs | 515 +++++++++++++++-------------------- src/bestdose/optimization.rs | 11 +- src/bestdose/predictions.rs | 69 +++-- src/bestdose/types.rs | 7 +- src/lib.rs | 1 + tests/bestdose_tests.rs | 136 +++++++++ 6 files changed, 413 insertions(+), 326 deletions(-) diff --git a/src/bestdose/cost.rs b/src/bestdose/cost.rs index 6f59be9fb..79ce09667 100644 --- a/src/bestdose/cost.rs +++ b/src/bestdose/cost.rs @@ -68,7 +68,10 @@ use crate::bestdose::predictions::{ use crate::bestdose::types::{Achievement, BestDoseObjective, Target}; use pharmsol::prelude::*; use pharmsol::Equation; +use pharmsol::OutputLabel; use pharmsol::Predictions; +use rayon::prelude::{IntoParallelIterator, ParallelIterator}; +use std::collections::HashMap; /// Cost together with the per-observation target achievements at a candidate /// dose regimen. @@ -77,6 +80,181 @@ pub(crate) struct Evaluation { pub achievements: Vec, } +/// Dense-grid setup for the AUC targets. +/// +/// Holds one dense-sampling subject per distinct output label. Every prediction +/// a simulation returns then belongs to that label, so labels never have to be +/// resolved to the dense output indices carried by `Prediction`. +struct AucGrid { + /// Output labels of the target observations, in their original order. + obs_labels: Vec, + groups: Vec, +} + +struct AucGroup { + label: OutputLabel, + dense_subject: Subject, + /// Per occasion of `dense_subject`, in order. + occasions: Vec, +} + +struct AucOccasion { + dense_times: Vec, + obs_times: Vec, +} + +fn build_auc_grid(target_subject: &Subject, prediction_interval: f64) -> AucGrid { + let obs_labels: Vec = target_subject + .occasions() + .iter() + .flat_map(|occ| occ.events()) + .filter_map(|event| match event { + Event::Observation(obs) => Some(obs.outeq().clone()), + _ => None, + }) + .collect(); + + let mut unique_labels = obs_labels.clone(); + unique_labels.sort(); + unique_labels.dedup(); + + let groups = unique_labels + .into_iter() + .map(|label| build_auc_group(target_subject, label, prediction_interval)) + .collect(); + + AucGrid { obs_labels, groups } +} + +fn build_auc_group( + target_subject: &Subject, + label: OutputLabel, + prediction_interval: f64, +) -> AucGroup { + // Cloning preserves the covariates and occasion structure the simulation needs; + // only the observations are swapped for dense sampling times of `label`. + let mut dense_subject = target_subject.clone(); + let mut occasions = Vec::with_capacity(dense_subject.occasions().len()); + + for occasion in dense_subject.iter_mut() { + let obs_times: Vec = occasion + .events() + .iter() + .filter_map(|event| match event { + Event::Observation(obs) if obs.outeq() == &label => Some(obs.time()), + _ => None, + }) + .collect(); + + let dense_times = if obs_times.is_empty() { + Vec::new() + } else { + let end_time = obs_times.last().copied().unwrap_or(0.0); + calculate_dense_times(0.0, end_time, &obs_times, prediction_interval) + }; + + occasion + .events_mut() + .retain(|event| !matches!(event, Event::Observation(_))); + for &t in &dense_times { + occasion.add_missing_observation(t, &label); + } + + occasions.push(AucOccasion { + dense_times, + obs_times, + }); + } + + AucGroup { + label, + dense_subject, + occasions, + } +} + +fn concentration_predictions( + eq: &E, + target_subject: &Subject, + spp: &[f64], +) -> Result> { + let pred = eq.simulate_subject_dense(target_subject, spp, None)?; + Ok(pred + .0 + .get_predictions() + .iter() + .map(|p| p.prediction()) + .collect()) +} + +fn auc_predictions( + eq: &E, + grid: &AucGrid, + target_type: Target, + spp: &[f64], +) -> Result> { + let mut aucs_by_label: HashMap<&OutputLabel, Vec> = HashMap::new(); + + for group in &grid.groups { + let pred = eq.simulate_subject_dense(&group.dense_subject, spp, None)?; + let dense_predictions: Vec = pred + .0 + .get_predictions() + .iter() + .map(|p| p.prediction()) + .collect(); + + // State resets between occasions, so integrate within one, never across. + let mut aucs = Vec::with_capacity(group.occasions.len()); + let mut offset = 0; + for (occasion, plan) in group.dense_subject.occasions().iter().zip(&group.occasions) { + let end = offset + plan.dense_times.len(); + let preds = dense_predictions.get(offset..end).ok_or_else(|| { + anyhow::anyhow!( + "expected {} dense predictions for output `{}`, got {}", + end, + group.label, + dense_predictions.len() + ) + })?; + offset = end; + + aucs.extend(match target_type { + Target::AUCFromLastDose => calculate_interval_auc_per_observation( + occasion, + &plan.dense_times, + preds, + &plan.obs_times, + )?, + _ => calculate_auc_at_times(&plan.dense_times, preds, &plan.obs_times)?, + }); + } + + aucs_by_label.insert(&group.label, aucs); + } + + // Reassemble into the original observation order. + let mut taken: HashMap<&OutputLabel, usize> = HashMap::new(); + grid.obs_labels + .iter() + .map(|label| { + let aucs = aucs_by_label + .get(label) + .ok_or_else(|| anyhow::anyhow!("no AUC group for output `{}`", label))?; + let index = taken.entry(label).or_insert(0); + let auc = aucs.get(*index).copied().ok_or_else(|| { + anyhow::anyhow!( + "AUC could not be computed for observation {} of output `{}`", + index, + label + ) + })?; + *index += 1; + Ok(auc) + }) + .collect() +} + /// Calculate cost function for a candidate dose regimen /// /// This is the core objective function minimized by the Nelder-Mead optimizer. @@ -238,318 +416,55 @@ pub(crate) fn evaluate( }) .collect(); - let obs_outeqs: Vec = target_subject + let obs_labels: Vec = target_subject .occasions() .iter() .flat_map(|occ| occ.events()) .filter_map(|event| match event { - Event::Observation(obs) => Some(obs.outeq_index().unwrap_or(0)), + Event::Observation(obs) => Some(obs.outeq().clone()), _ => None, }) .collect(); let n_obs = obs_vec.len(); - // Accumulators - let mut variance = 0.0_f64; // Expected squared error E[(target - pred)²] - let mut y_bar = vec![0.0_f64; n_obs]; // Weighted-mean predictions - - // Both cost terms are computed from the single distribution weights. - for (row, prob) in problem - .theta - .matrix() - .row_iter() - .zip(problem.weights.iter()) - { - let spp = row.iter().copied().collect::>(); - - // Get predictions based on target type - let preds_i: Vec = match problem.target_type { - Target::Concentration => { - // Simulate at observation times only - let pred = problem - .eq - .simulate_subject_dense(&target_subject, &spp, None)?; - pred.0 - .get_predictions() - .iter() - .map(|p| p.prediction()) - .collect() - } - Target::AUCFromZero => { - // For AUC: simulate at dense time grid and calculate cumulative AUC - let idelta = problem.prediction_interval; - let start_time = 0.0; // Future starts at 0 - let end_time = obs_times.last().copied().unwrap_or(0.0); - - // Generate dense time grid - let dense_times = calculate_dense_times(start_time, end_time, &obs_times, idelta); - - // Create temporary subject with dense time points for simulation - let subject_id = target_subject.id().to_string(); - let mut builder = Subject::builder(&subject_id); - - // Add all doses from original subject - for occasion in target_subject.occasions() { - for event in occasion.events() { - match event { - Event::Bolus(bolus) => { - builder = - builder.bolus(bolus.time(), bolus.amount(), bolus.input()); - } - Event::Infusion(infusion) => { - builder = builder.infusion( - infusion.time(), - infusion.amount(), - infusion.input(), - infusion.duration(), - ); - } - Event::Observation(_) => {} // Skip original observations - } - } - } - - // Collect observations with (time, outeq) pairs to preserve original order - let obs_time_outeq: Vec<(f64, usize)> = target_subject - .occasions() - .iter() - .flat_map(|occ| occ.events()) - .filter_map(|event| match event { - Event::Observation(obs) => Some( - obs.outeq_index() - .map(|outeq| (obs.time(), outeq)) - .ok_or_else(|| { - anyhow::anyhow!( - "BestDose AUC calculations require numeric observation output labels; got `{}`", - obs.outeq() - ) - }), - ), - _ => None, - }) - .collect::>>()?; - - let mut unique_outeqs: Vec = - obs_time_outeq.iter().map(|(_, outeq)| *outeq).collect(); - unique_outeqs.sort(); - unique_outeqs.dedup(); - - // Add observations at dense times (with dummy values for timing only) - for outeq in unique_outeqs.iter() { - for &t in &dense_times { - builder = builder.missing_observation(t, *outeq); - } - } - - let dense_subject = builder.build(); - - // Simulate at dense times - let pred = problem - .eq - .simulate_subject_dense(&dense_subject, &spp, None)?; - let dense_predictions_with_outeq = pred.0.get_predictions(); - - // Group predictions by outeq using the Prediction struct - let mut outeq_predictions: std::collections::HashMap> = - std::collections::HashMap::new(); - - for prediction in dense_predictions_with_outeq { - outeq_predictions - .entry(prediction.outeq()) - .or_default() - .push(prediction.prediction()); - } - - // Calculate AUC for each outeq separately - let mut outeq_aucs: std::collections::HashMap> = - std::collections::HashMap::new(); - - for &outeq in unique_outeqs.iter() { - let outeq_preds = outeq_predictions.get(&outeq).ok_or_else(|| { - anyhow::anyhow!("Missing predictions for outeq {}", outeq) - })?; - - // Get observation times for this outeq only - let outeq_obs_times: Vec = obs_time_outeq - .iter() - .filter(|(_, o)| *o == outeq) - .map(|(t, _)| *t) - .collect(); - - // Calculate AUC at observation times for this outeq - let aucs = calculate_auc_at_times(&dense_times, outeq_preds, &outeq_obs_times); - outeq_aucs.insert(outeq, aucs); - } - - // Build final AUC vector in original observation order - let mut result_aucs = Vec::with_capacity(obs_time_outeq.len()); - let mut outeq_counters: std::collections::HashMap = - std::collections::HashMap::new(); - - for (_, outeq) in obs_time_outeq.iter() { - let aucs = outeq_aucs - .get(outeq) - .ok_or_else(|| anyhow::anyhow!("Missing AUC for outeq {}", outeq))?; - - let counter = outeq_counters.entry(*outeq).or_insert(0); - if *counter < aucs.len() { - result_aucs.push(aucs[*counter]); - *counter += 1; - } else { - return Err(anyhow::anyhow!( - "AUC index out of bounds for outeq {}", - outeq - )); - } - } + let auc_grid = match problem.target_type { + Target::Concentration => None, + Target::AUCFromZero | Target::AUCFromLastDose => { + Some(build_auc_grid(&target_subject, problem.prediction_interval)) + } + }; + + // Simulation dominates the cost of an evaluation, so support points — which + // are independent — are simulated in parallel. + let theta = problem.theta.matrix(); + let preds: Vec> = (0..theta.nrows()) + .into_par_iter() + .map(|i| { + let spp: Vec = theta.row(i).iter().copied().collect(); + let preds_i = match &auc_grid { + None => concentration_predictions(&problem.eq, &target_subject, &spp)?, + Some(grid) => auc_predictions(&problem.eq, grid, problem.target_type, &spp)?, + }; - result_aucs + if preds_i.len() != n_obs { + return Err(anyhow::anyhow!( + "prediction length ({}) != observation length ({})", + preds_i.len(), + n_obs + )); } - Target::AUCFromLastDose => { - // For interval AUC: simulate at dense time grid and calculate AUC from last dose - let idelta = problem.prediction_interval; - let end_time = obs_times.last().copied().unwrap_or(0.0); - - // Generate dense time grid from 0 to end_time (need full grid for intervals) - let dense_times = calculate_dense_times(0.0, end_time, &obs_times, idelta); - - // Create temporary subject with dense time points for simulation - let subject_id = target_subject.id().to_string(); - let mut builder = Subject::builder(&subject_id); - - // Add all doses from original subject - for occasion in target_subject.occasions() { - for event in occasion.events() { - match event { - Event::Bolus(bolus) => { - builder = - builder.bolus(bolus.time(), bolus.amount(), bolus.input()); - } - Event::Infusion(infusion) => { - builder = builder.infusion( - infusion.time(), - infusion.amount(), - infusion.input(), - infusion.duration(), - ); - } - Event::Observation(_) => {} // Skip original observations - } - } - } - // Collect observations with (time, outeq) pairs to preserve original order - let obs_time_outeq: Vec<(f64, usize)> = target_subject - .occasions() - .iter() - .flat_map(|occ| occ.events()) - .filter_map(|event| match event { - Event::Observation(obs) => Some( - obs.outeq_index() - .map(|outeq| (obs.time(), outeq)) - .ok_or_else(|| { - anyhow::anyhow!( - "BestDose AUC calculations require numeric observation output labels; got `{}`", - obs.outeq() - ) - }), - ), - _ => None, - }) - .collect::>>()?; - - let mut unique_outeqs: Vec = - obs_time_outeq.iter().map(|(_, outeq)| *outeq).collect(); - unique_outeqs.sort(); - unique_outeqs.dedup(); - - // Add observations at dense times - for outeq in unique_outeqs.iter() { - for &t in &dense_times { - builder = builder.missing_observation(t, *outeq); - } - } - - let dense_subject = builder.build(); - - // Simulate at dense times - let pred = problem - .eq - .simulate_subject_dense(&dense_subject, &spp, None)?; - let dense_predictions_with_outeq = pred.0.get_predictions(); - - // Group predictions by outeq - let mut outeq_predictions: std::collections::HashMap> = - std::collections::HashMap::new(); - - for prediction in dense_predictions_with_outeq { - outeq_predictions - .entry(prediction.outeq()) - .or_default() - .push(prediction.prediction()); - } - - // Calculate interval AUC for each outeq separately - let mut outeq_aucs: std::collections::HashMap> = - std::collections::HashMap::new(); - - for &outeq in unique_outeqs.iter() { - let outeq_preds = outeq_predictions.get(&outeq).ok_or_else(|| { - anyhow::anyhow!("Missing predictions for outeq {}", outeq) - })?; - - // Get observation times for this outeq only - let outeq_obs_times: Vec = obs_time_outeq - .iter() - .filter(|(_, o)| *o == outeq) - .map(|(t, _)| *t) - .collect(); - - // Calculate interval AUC at observation times for this outeq - let aucs = calculate_interval_auc_per_observation( - &target_subject, - &dense_times, - outeq_preds, - &outeq_obs_times, - ); - outeq_aucs.insert(outeq, aucs); - } - - // Build final AUC vector in original observation order - let mut result_aucs = Vec::with_capacity(obs_time_outeq.len()); - let mut outeq_counters: std::collections::HashMap = - std::collections::HashMap::new(); - - for (_, outeq) in obs_time_outeq.iter() { - let aucs = outeq_aucs - .get(outeq) - .ok_or_else(|| anyhow::anyhow!("Missing AUC for outeq {}", outeq))?; - - let counter = outeq_counters.entry(*outeq).or_insert(0); - if *counter < aucs.len() { - result_aucs.push(aucs[*counter]); - *counter += 1; - } else { - return Err(anyhow::anyhow!( - "AUC index out of bounds for outeq {}", - outeq - )); - } - } - - result_aucs - } - }; + Ok(preds_i) + }) + .collect::>>()?; - if preds_i.len() != n_obs { - return Err(anyhow::anyhow!( - "prediction length ({}) != observation length ({})", - preds_i.len(), - n_obs - )); - } + // Accumulators + let mut variance = 0.0_f64; // Expected squared error E[(target - pred)²] + let mut y_bar = vec![0.0_f64; n_obs]; // Weighted-mean predictions + // Both cost terms are computed from the single distribution weights. + for (preds_i, prob) in preds.iter().zip(problem.weights.iter()) { // Calculate variance term: weighted by the distribution probability let mut sumsq_i = 0.0_f64; for (j, &obs_val) in obs_vec.iter().enumerate() { @@ -577,12 +492,12 @@ pub(crate) fn evaluate( // Expected achieved value at each observation is the weighted-mean prediction. let achievements = obs_times .iter() - .zip(obs_outeqs.iter()) + .zip(obs_labels.iter()) .zip(obs_vec.iter()) .zip(y_bar.iter()) - .map(|(((&time, &outeq), &target), &achieved)| Achievement { + .map(|(((&time, outeq), &target), &achieved)| Achievement { time, - outeq, + outeq: outeq.clone(), target, achieved, }) diff --git a/src/bestdose/optimization.rs b/src/bestdose/optimization.rs index 400df6f5b..a4a398f43 100644 --- a/src/bestdose/optimization.rs +++ b/src/bestdose/optimization.rs @@ -85,6 +85,12 @@ pub(crate) fn optimize(objective: &BestDoseObjective) -> Result< let initial_point = vec![initial_guess; num_optimizable]; let initial_simplex = create_initial_simplex(&initial_point); + // Nelder-Mead unwraps cost errors while evaluating the initial simplex, + // so evaluate it here first to report any failure as an error. + for vertex in &initial_simplex { + calculate_cost(objective, vertex)?; + } + let solver: NelderMead, f64> = NelderMead::new(initial_simplex).with_sd_tolerance(1e-10)?; @@ -92,7 +98,10 @@ pub(crate) fn optimize(objective: &BestDoseObjective) -> Result< .configure(|state| state.max_iters(1000)) .run()?; - opt.state().best_param.clone().unwrap() + opt.state() + .best_param + .clone() + .ok_or_else(|| anyhow::anyhow!("Nelder-Mead returned no best dose vector"))? }; // Evaluate once at the optimum to recover the cost and target achievements. diff --git a/src/bestdose/predictions.rs b/src/bestdose/predictions.rs index e73948876..ae0b0e4e2 100644 --- a/src/bestdose/predictions.rs +++ b/src/bestdose/predictions.rs @@ -7,25 +7,27 @@ //! AUC(t) = Σᵢ (C[i] + C[i-1]) / 2 × (t[i] - t[i-1]) //! ``` +use anyhow::{anyhow, Result}; use pharmsol::prelude::*; /// Find the time of the last dose (bolus or infusion) before a given observation /// time. Returns `0.0` if no dose exists before `obs_time`. -pub fn find_last_dose_time_before(subject: &Subject, obs_time: f64) -> f64 { +/// +/// Scoped to one occasion: the model state resets between occasions, so a dose +/// in an earlier occasion is not the last dose for this one. +pub fn find_last_dose_time_before(occasion: &Occasion, obs_time: f64) -> f64 { let mut last_dose_time = 0.0; - for occasion in subject.occasions() { - for event in occasion.events() { - let event_time = match event { - Event::Bolus(b) => Some(b.time()), - Event::Infusion(i) => Some(i.time()), - Event::Observation(_) => None, - }; - - if let Some(t) = event_time { - if t < obs_time && t > last_dose_time { - last_dose_time = t; - } + for event in occasion.events() { + let event_time = match event { + Event::Bolus(b) => Some(b.time()), + Event::Infusion(i) => Some(i.time()), + Event::Observation(_) => None, + }; + + if let Some(t) = event_time { + if t < obs_time && t > last_dose_time { + last_dose_time = t; } } } @@ -60,7 +62,7 @@ pub fn calculate_dense_times( times.push(end_time); } - times.sort_by(|a, b| a.partial_cmp(b).unwrap()); + times.sort_by(f64::total_cmp); let tolerance = 1e-10; let mut unique_times = Vec::new(); @@ -84,8 +86,14 @@ pub fn calculate_auc_at_times( dense_times: &[f64], dense_predictions: &[f64], target_times: &[f64], -) -> Vec { - assert_eq!(dense_times.len(), dense_predictions.len()); +) -> Result> { + if dense_times.len() != dense_predictions.len() { + return Err(anyhow!( + "dense time grid ({}) and prediction ({}) lengths differ", + dense_times.len(), + dense_predictions.len() + )); + } let mut target_aucs = Vec::with_capacity(target_times.len()); let mut auc = 0.0; @@ -105,7 +113,7 @@ pub fn calculate_auc_at_times( } } - target_aucs + Ok(target_aucs) } /// Calculate interval AUC for each observation independently. @@ -113,18 +121,35 @@ pub fn calculate_auc_at_times( /// For each observation at time `t`, integrates from the last dose before `t` to /// `t` (e.g. dosing-interval AUCτ at steady state). pub fn calculate_interval_auc_per_observation( - subject: &Subject, + occasion: &Occasion, dense_times: &[f64], dense_predictions: &[f64], obs_times: &[f64], -) -> Vec { - assert_eq!(dense_times.len(), dense_predictions.len()); +) -> Result> { + if dense_times.len() != dense_predictions.len() { + return Err(anyhow!( + "dense time grid ({}) and prediction ({}) lengths differ", + dense_times.len(), + dense_predictions.len() + )); + } + + if obs_times.is_empty() { + return Ok(Vec::new()); + } + + if dense_times.is_empty() { + return Err(anyhow!( + "cannot compute interval AUC for {} observations on an empty time grid", + obs_times.len() + )); + } let mut interval_aucs = Vec::with_capacity(obs_times.len()); let tolerance = 1e-10; for &obs_time in obs_times { - let last_dose_time = find_last_dose_time_before(subject, obs_time); + let last_dose_time = find_last_dose_time_before(occasion, obs_time); let start_idx = dense_times .iter() @@ -146,5 +171,5 @@ pub fn calculate_interval_auc_per_observation( interval_aucs.push(auc); } - interval_aucs + Ok(interval_aucs) } diff --git a/src/bestdose/types.rs b/src/bestdose/types.rs index 67d38da3d..d60a442ee 100644 --- a/src/bestdose/types.rs +++ b/src/bestdose/types.rs @@ -9,6 +9,7 @@ use crate::estimation::nonparametric::{Theta, Weights}; use pharmsol::prelude::*; use pharmsol::Equation; +use pharmsol::OutputLabel; use serde::{Deserialize, Serialize}; /// Target type for dose optimization. @@ -210,12 +211,12 @@ pub(crate) struct BestDoseObjective { /// `achieved` is the expected value under the distribution (the weighted mean /// prediction across support points) — a concentration for /// [`Target::Concentration`] or an AUC for the AUC targets. -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone)] pub struct Achievement { /// Observation time. pub time: f64, - /// Output equation index of the observation. - pub outeq: usize, + /// Output label of the observation. + pub outeq: OutputLabel, /// The requested target value at this observation. pub target: f64, /// The expected achieved value at the optimal doses. diff --git a/src/lib.rs b/src/lib.rs index 824aae3f9..34f5266d5 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -89,6 +89,7 @@ pub mod prelude { // Items required by downstream code that are not part of `pharmsol::prelude`. pub use pharmsol::equation::{EquationTypes, Predictions}; pub use pharmsol::optimize::effect::get_e2; + pub use pharmsol::OutputLabel; pub use pharmsol::{ODE, SDE}; // Organized submodules mirroring pharmsol's grouping. diff --git a/tests/bestdose_tests.rs b/tests/bestdose_tests.rs index e0bb6ee26..fe6cebb55 100644 --- a/tests/bestdose_tests.rs +++ b/tests/bestdose_tests.rs @@ -49,6 +49,26 @@ fn infusion_model() -> ODE { } } +/// One-compartment bolus model whose volume scales with a `wt` covariate. +fn covariate_model() -> ODE { + ode! { + name: "one_compartment_bolus_wt", + params: [ke, v], + covariates: [wt], + states: [central], + outputs: [outeq_0], + routes: [ + bolus(input_0) -> central, + ], + diffeq: |x, _t, dx| { + dx[central] = -ke * x[central]; + }, + out: |x, _t, y| { + y[outeq_0] = x[central] / (v * wt / 70.0); + }, + } +} + fn parameter_space() -> ParameterSpace { ParameterSpace::::new() .add("ke", 0.001, 3.0) @@ -298,6 +318,122 @@ fn auc_from_zero_hits_target() -> Result<()> { Ok(()) } +/// AUC targets must accept named output labels, not just numeric ones. +#[test] +fn auc_target_accepts_named_output_labels() -> Result<()> { + let problem = BestDoseProblem::new(bolus_model(), theta(&[[0.3, 50.0]]), Weights::uniform(1))?; + + let target_auc = 100.0; + let target = Subject::builder("p") + .bolus(0.0, 0.0, 0) + .observation(12.0, target_auc, "outeq_0") + .build(); + + let result = problem.optimize( + target, + Target::AUCFromZero, + DoseRange::new(0.0, 5000.0), + 0.0, + BestDoseOptions { + prediction_interval: 0.05, + }, + )?; + + let achievement = &result.achievements()[0]; + assert_eq!(achievement.outeq.as_str(), "outeq_0"); + let rel_error = ((achievement.achieved - target_auc) / target_auc).abs(); + assert!( + rel_error < 0.02, + "achieved AUC {} vs target {} (rel error {})", + achievement.achieved, + target_auc, + rel_error + ); + Ok(()) +} + +/// The dense AUC grid must keep the target's covariates, not just its doses. +#[test] +fn auc_target_uses_subject_covariates() -> Result<()> { + let problem = BestDoseProblem::new( + covariate_model(), + theta(&[[0.3, 50.0]]), + Weights::uniform(1), + )?; + + let dose_for = |wt: f64| -> Result { + let target = Subject::builder("p") + .bolus(0.0, 0.0, 0) + .observation(12.0, 100.0, 0) + .covariate("wt", 0.0, wt) + .build(); + + let result = problem.optimize( + target, + Target::AUCFromZero, + DoseRange::new(0.0, 20_000.0), + 0.0, + BestDoseOptions { + prediction_interval: 0.05, + }, + )?; + Ok(result.doses()[0]) + }; + + // Volume is proportional to wt, so the dose hitting a fixed AUC must be too. + let light = dose_for(70.0)?; + let heavy = dose_for(140.0)?; + assert!( + (heavy / light - 2.0).abs() < 0.02, + "dose must scale with the wt covariate: {light} vs {heavy}" + ); + Ok(()) +} + +/// The dense AUC grid must keep the target's occasions, which reset model state. +#[test] +fn auc_target_preserves_occasions() -> Result<()> { + let problem = BestDoseProblem::new(bolus_model(), theta(&[[0.3, 50.0]]), Weights::uniform(1))?; + + let target_auc = 100.0; + let target = Subject::builder("p") + .bolus(0.0, 0.0, 0) + .observation(12.0, target_auc, 0) + .reset() + .bolus(0.0, 0.0, 0) + .observation(12.0, target_auc, 0) + .build(); + + let result = problem.optimize( + target, + Target::AUCFromZero, + DoseRange::new(0.0, 5000.0), + 0.0, + BestDoseOptions { + prediction_interval: 0.05, + }, + )?; + + let doses = result.doses(); + assert_eq!(doses.len(), 2); + assert!( + (doses[0] - doses[1]).abs() / doses[0] < 0.01, + "identical occasions must need identical doses: {doses:?}" + ); + + let achievements = result.achievements(); + assert_eq!(achievements.len(), 2); + for a in achievements { + let rel_error = ((a.achieved - target_auc) / target_auc).abs(); + assert!( + rel_error < 0.02, + "achieved AUC {} vs target {target_auc}", + a.achieved + ); + } + Ok(()) +} + #[test] fn auc_from_last_dose_optimizes_maintenance_dose() -> Result<()> { let problem = BestDoseProblem::new(bolus_model(), theta(&[[0.3, 50.0]]), Weights::uniform(1))?; From f1ceac98319f6d7db871c3d9958e250598b4db14 Mon Sep 17 00:00:00 2001 From: Markus <66058642+mhovd@users.noreply.github.com> Date: Fri, 28 Aug 2026 11:31:43 +0200 Subject: [PATCH 2/3] tests --- src/bestdose/cost.rs | 6 ++++++ tests/bestdose_tests.rs | 34 ++++++++++++++++++++++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/src/bestdose/cost.rs b/src/bestdose/cost.rs index 79ce09667..4021f6eaf 100644 --- a/src/bestdose/cost.rs +++ b/src/bestdose/cost.rs @@ -406,6 +406,12 @@ pub(crate) fn evaluate( )); } + if let Some(bad) = obs_times.iter().find(|t| !t.is_finite()) { + return Err(anyhow::anyhow!( + "Target observation times must be finite, got {bad}" + )); + } + let obs_vec: Vec = target_subject .occasions() .iter() diff --git a/tests/bestdose_tests.rs b/tests/bestdose_tests.rs index fe6cebb55..881eea523 100644 --- a/tests/bestdose_tests.rs +++ b/tests/bestdose_tests.rs @@ -229,6 +229,40 @@ fn all_fixed_doses_return_unchanged() -> Result<()> { Ok(()) } +/// A failing cost function must surface as an error, never a panic. +#[test] +fn optimize_errors_instead_of_panicking() -> Result<()> { + let problem = BestDoseProblem::new(bolus_model(), theta(&[[0.3, 50.0]]), Weights::uniform(1))?; + + let no_observations = Subject::builder("p").bolus(0.0, 0.0, 0).build(); + let err = problem + .optimize( + no_observations, + Target::Concentration, + DoseRange::new(0.0, 300.0), + 0.0, + BestDoseOptions::default(), + ) + .expect_err("a target without observations must error"); + assert!(err.to_string().contains("no observations"), "{err}"); + + let nan_time = Subject::builder("p") + .bolus(0.0, 0.0, 0) + .observation(f64::NAN, 5.0, 0) + .build(); + let err = problem + .optimize( + nan_time, + Target::AUCFromZero, + DoseRange::new(0.0, 300.0), + 0.0, + BestDoseOptions::default(), + ) + .expect_err("a non-finite observation time must error"); + assert!(err.to_string().contains("finite"), "{err}"); + Ok(()) +} + #[test] fn infusions_are_optimizable() -> Result<()> { let problem = From d51292f83bf964b435c89ae46e979700fe761aea Mon Sep 17 00:00:00 2001 From: Markus <66058642+mhovd@users.noreply.github.com> Date: Mon, 28 Sep 2026 13:53:23 +0200 Subject: [PATCH 3/3] fix tests --- Cargo.toml | 2 +- tests/bestdose_tests.rs | 5 +---- 2 files changed, 2 insertions(+), 5 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 975130b99..4fd52626c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -28,7 +28,7 @@ tracing-subscriber = { version = "0.3.19", features = [ "time", ] } faer = "0.24.0" -pharmsol = "0.29.2" +pharmsol = "0.29.3" anyhow = "1.0.100" rayon = "1.10.0" rand = "0.10.1" diff --git a/tests/bestdose_tests.rs b/tests/bestdose_tests.rs index 5553ebc4b..21a51c56f 100644 --- a/tests/bestdose_tests.rs +++ b/tests/bestdose_tests.rs @@ -51,11 +51,8 @@ fn covariate_model() -> ODE { covariates: [wt], states: [central], outputs: [outeq_0], - routes: [ - bolus(input_0) -> central, - ], diffeq: |x, _t, dx| { - dx[central] = -ke * x[central]; + dx[central] = -ke * x[central] + bolus(input_0) ; }, out: |x, _t, y| { y[outeq_0] = x[central] / (v * wt / 70.0);