diff --git a/benches/optimization.rs b/benches/optimization.rs index cfdaafe..3435781 100644 --- a/benches/optimization.rs +++ b/benches/optimization.rs @@ -23,7 +23,7 @@ fn bench_tpe_sphere(c: &mut Criterion) { b.iter(|| { let study = Study::minimize(TpeSampler::builder().seed(42).build().unwrap()); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -48,7 +48,7 @@ fn bench_tpe_rosenbrock(c: &mut Criterion) { b.iter(|| { let study = Study::minimize(TpeSampler::builder().seed(42).build().unwrap()); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -72,7 +72,7 @@ fn bench_random_vs_tpe(c: &mut Criterion) { b.iter(|| { let study = Study::minimize(RandomSampler::with_seed(42)); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -88,7 +88,7 @@ fn bench_random_vs_tpe(c: &mut Criterion) { b.iter(|| { let study = Study::minimize(TpeSampler::builder().seed(42).build().unwrap()); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x: Vec = params .iter() .map(|p| p.suggest(trial)) diff --git a/examples/basic_optimization.rs b/examples/basic_optimization.rs index 097853a..e868d75 100644 --- a/examples/basic_optimization.rs +++ b/examples/basic_optimization.rs @@ -17,7 +17,7 @@ fn main() { // Run 50 trials, each evaluating f(x) = (x - 3)² study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let x_val = x.suggest(trial)?; let value = (x_val - 3.0).powi(2); Ok::<_, Error>(value) diff --git a/examples/early_stopping.rs b/examples/early_stopping.rs index efbf78c..7bfe707 100644 --- a/examples/early_stopping.rs +++ b/examples/early_stopping.rs @@ -44,7 +44,7 @@ fn main() -> optimizer::Result<()> { target: 0.01, }; - study.optimize_with(100, objective)?; + study.optimize(100, objective)?; let best = study.best_trial()?; println!( diff --git a/examples/journal_storage.rs b/examples/journal_storage.rs index c85eba4..a3410b6 100644 --- a/examples/journal_storage.rs +++ b/examples/journal_storage.rs @@ -24,7 +24,7 @@ fn main() -> optimizer::Result<()> { .storage(storage) .build(); - study.optimize(20, |trial| { + study.optimize(20, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(xv * xv) })?; @@ -46,7 +46,7 @@ fn main() -> optimizer::Result<()> { .build(); let before = study.n_trials(); - study.optimize(10, |trial| { + study.optimize(10, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(xv * xv) })?; diff --git a/examples/multi_objective.rs b/examples/multi_objective.rs index 3d66a8a..804e47a 100644 --- a/examples/multi_objective.rs +++ b/examples/multi_objective.rs @@ -15,7 +15,7 @@ fn main() -> optimizer::Result<()> { // Classic bi-objective: f1(x) = x², f2(x) = (x-1)² // The Pareto front is the curve where improving f1 worsens f2. - study.optimize(50, |trial| { + study.optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let f1 = xv * xv; let f2 = (xv - 1.0) * (xv - 1.0); diff --git a/examples/parameter_types.rs b/examples/parameter_types.rs index c509c5b..5efc823 100644 --- a/examples/parameter_types.rs +++ b/examples/parameter_types.rs @@ -41,7 +41,7 @@ fn main() { // --- Run the optimization --- study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let lr_val = lr.suggest(trial)?; let layers = n_layers.suggest(trial)?; let opt = optimizer.suggest(trial)?; diff --git a/examples/pruning.rs b/examples/pruning.rs index afbfea2..284afe5 100644 --- a/examples/pruning.rs +++ b/examples/pruning.rs @@ -26,7 +26,7 @@ fn main() -> optimizer::Result<()> { let n_epochs: u64 = 20; - study.optimize(30, |trial| { + study.optimize(30, |trial: &mut optimizer::Trial| { let lr_val = lr.suggest(trial)?; let mom = momentum.suggest(trial)?; diff --git a/examples/sampler_comparison.rs b/examples/sampler_comparison.rs index 3034ab3..9e484e2 100644 --- a/examples/sampler_comparison.rs +++ b/examples/sampler_comparison.rs @@ -20,7 +20,7 @@ fn run_study(study: Study, n_trials: usize) -> f64 { let y = FloatParam::new(-3.0, 3.0).name("y"); study - .optimize(n_trials, |trial| { + .optimize(n_trials, |trial: &mut optimizer::Trial| { let x_val = x.suggest(trial)?; let y_val = y.suggest(trial)?; Ok::<_, Error>(sphere(x_val, y_val)) diff --git a/src/fanova.rs b/src/fanova.rs index 91b6af0..89cd452 100644 --- a/src/fanova.rs +++ b/src/fanova.rs @@ -41,7 +41,7 @@ //! let y = FloatParam::new(0.0, 10.0).name("y"); //! //! study -//! .optimize(50, |trial| { +//! .optimize(50, |trial: &mut optimizer::Trial| { //! let xv = x.suggest(trial)?; //! let yv = y.suggest(trial)?; //! // x matters much more than y diff --git a/src/lib.rs b/src/lib.rs index a696b6a..1424b33 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -26,7 +26,7 @@ //! let x = FloatParam::new(-10.0, 10.0).name("x"); //! //! study -//! .optimize(50, |trial| { +//! .optimize(50, |trial: &mut optimizer::Trial| { //! let v = x.suggest(trial)?; //! Ok::<_, Error>((v - 3.0).powi(2)) //! }) diff --git a/src/multi_objective.rs b/src/multi_objective.rs index 8dd144a..a82b1ad 100644 --- a/src/multi_objective.rs +++ b/src/multi_objective.rs @@ -35,7 +35,7 @@ //! let x = FloatParam::new(0.0, 1.0); //! //! study -//! .optimize(20, |trial| { +//! .optimize(20, |trial: &mut optimizer::Trial| { //! let xv = x.suggest(trial)?; //! Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) //! }) @@ -229,7 +229,7 @@ impl Sampler for MoSamplerBridge { /// let x = FloatParam::new(0.0, 1.0); /// /// study -/// .optimize(30, |trial| { +/// .optimize(30, |trial: &mut optimizer::Trial| { /// let xv = x.suggest(trial)?; /// Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) /// }) diff --git a/src/objective.rs b/src/objective.rs index d1ce710..c22babf 100644 --- a/src/objective.rs +++ b/src/objective.rs @@ -1,6 +1,9 @@ //! The [`Objective`] trait defines what gets optimized. //! -//! For simple closures, pass them directly to +//! # Closures work directly +//! +//! Any `Fn(&mut Trial) -> Result` closure automatically implements +//! [`Objective`], so you can pass closures straight to //! [`Study::optimize`](crate::Study::optimize): //! //! ``` @@ -10,16 +13,18 @@ //! let x = FloatParam::new(-10.0, 10.0).name("x"); //! //! study -//! .optimize(50, |trial| { +//! .optimize(50, |trial: &mut optimizer::Trial| { //! let v = x.suggest(trial)?; //! Ok::<_, Error>((v - 3.0).powi(2)) //! }) //! .unwrap(); //! ``` //! -//! For richer control — early stopping, retries, or per-trial logging — -//! implement [`Objective`] on a struct and pass it to -//! [`Study::optimize_with`](crate::Study::optimize_with): +//! # Structs for lifecycle hooks +//! +//! For richer control — early stopping or per-trial logging — implement +//! [`Objective`] on a struct and pass it to the same +//! [`Study::optimize`](crate::Study::optimize) method: //! //! ``` //! use std::ops::ControlFlow; @@ -54,7 +59,7 @@ //! x: FloatParam::new(-10.0, 10.0).name("x"), //! target: 1.0, //! }; -//! study.optimize_with(200, obj).unwrap(); +//! study.optimize(200, obj).unwrap(); //! assert!(study.best_value().unwrap() < 1.0); //! ``` @@ -69,15 +74,13 @@ use crate::trial::Trial; /// The only required method is [`evaluate`](Objective::evaluate), which /// computes the objective value for a given trial. Optional hooks provide /// early stopping ([`before_trial`](Objective::before_trial), -/// [`after_trial`](Objective::after_trial)) and automatic retries -/// ([`max_retries`](Objective::max_retries)). +/// [`after_trial`](Objective::after_trial)). /// -/// # When to use `Objective` vs a closure +/// # Closures implement `Objective` automatically /// -/// - **Closure** — pass directly to [`Study::optimize`](crate::Study::optimize) -/// for simple evaluate-only objectives. -/// - **`Objective` struct** — implement this trait when you need hooks -/// (`before_trial`, `after_trial`) or retries. +/// A blanket implementation covers all `Fn(&mut Trial) -> Result` +/// closures, so you can pass closures directly to +/// [`Study::optimize`](crate::Study::optimize) without wrapping them. /// /// # Thread safety /// @@ -120,13 +123,19 @@ pub trait Objective { fn after_trial(&self, _study: &Study, _trial: &CompletedTrial) -> ControlFlow<()> { ControlFlow::Continue(()) } +} - /// Maximum number of retries for a failed trial. - /// - /// When `evaluate` returns a non-pruning error and retries remain, - /// the same parameter configuration is re-evaluated. Set to `0` - /// (the default) to disable retries. - fn max_retries(&self) -> usize { - 0 +/// Blanket implementation: any `Fn(&mut Trial) -> Result` is an +/// `Objective` with no lifecycle hooks. +impl Objective for F +where + F: Fn(&mut Trial) -> Result, + V: PartialOrd, + E: ToString + 'static, +{ + type Error = E; + + fn evaluate(&self, trial: &mut Trial) -> Result { + self(trial) } } diff --git a/src/sampler/mod.rs b/src/sampler/mod.rs index b512c49..cfa5537 100644 --- a/src/sampler/mod.rs +++ b/src/sampler/mod.rs @@ -138,7 +138,7 @@ impl CompletedTrial { /// let x = FloatParam::new(-10.0, 10.0); /// /// study - /// .optimize(5, |trial| { + /// .optimize(5, |trial: &mut optimizer::Trial| { /// let val = x.suggest(trial)?; /// Ok::<_, optimizer::Error>(val * val) /// }) diff --git a/src/sampler/moead.rs b/src/sampler/moead.rs index 6f11d44..4e6bb78 100644 --- a/src/sampler/moead.rs +++ b/src/sampler/moead.rs @@ -68,7 +68,7 @@ //! //! let x = FloatParam::new(0.0, 1.0); //! study -//! .optimize(100, |trial| { +//! .optimize(100, |trial: &mut optimizer::Trial| { //! let xv = x.suggest(trial)?; //! Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) //! }) diff --git a/src/sampler/motpe.rs b/src/sampler/motpe.rs index 823b791..abffc78 100644 --- a/src/sampler/motpe.rs +++ b/src/sampler/motpe.rs @@ -48,7 +48,7 @@ //! //! let x = FloatParam::new(0.0, 1.0); //! study -//! .optimize(30, |trial| { +//! .optimize(30, |trial: &mut optimizer::Trial| { //! let xv = x.suggest(trial)?; //! Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) //! }) @@ -103,7 +103,7 @@ use crate::{pareto, rng_util}; /// /// let x = FloatParam::new(0.0, 1.0); /// study -/// .optimize(30, |trial| { +/// .optimize(30, |trial: &mut optimizer::Trial| { /// let xv = x.suggest(trial)?; /// Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) /// }) @@ -660,6 +660,7 @@ mod tests { use super::*; use crate::distribution::{CategoricalDistribution, FloatDistribution, IntDistribution}; use crate::parameter::ParamId; + use crate::trial::Trial; fn create_mo_trial( id: u64, @@ -930,7 +931,7 @@ mod tests { let x = FloatParam::new(0.0, 1.0); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut Trial| { let xv = x.suggest(trial)?; Ok::<_, crate::Error>(vec![xv, 1.0 - xv]) }) diff --git a/src/sampler/nsga2.rs b/src/sampler/nsga2.rs index 0a8f873..f806032 100644 --- a/src/sampler/nsga2.rs +++ b/src/sampler/nsga2.rs @@ -55,7 +55,7 @@ //! //! let x = FloatParam::new(0.0, 1.0); //! study -//! .optimize(50, |trial| { +//! .optimize(50, |trial: &mut optimizer::Trial| { //! let xv = x.suggest(trial)?; //! Ok::<_, optimizer::Error>(vec![xv * xv, (xv - 1.0).powi(2)]) //! }) diff --git a/src/sampler/nsga3.rs b/src/sampler/nsga3.rs index 44028ef..7ac840c 100644 --- a/src/sampler/nsga3.rs +++ b/src/sampler/nsga3.rs @@ -66,7 +66,7 @@ //! let x = FloatParam::new(0.0, 1.0); //! let y = FloatParam::new(0.0, 1.0); //! study -//! .optimize(100, |trial| { +//! .optimize(100, |trial: &mut optimizer::Trial| { //! let xv = x.suggest(trial)?; //! let yv = y.suggest(trial)?; //! Ok::<_, optimizer::Error>(vec![xv, yv, (1.0 - xv - yv).abs()]) diff --git a/src/sampler/tpe/multivariate.rs b/src/sampler/tpe/multivariate.rs index 72ea840..502125f 100644 --- a/src/sampler/tpe/multivariate.rs +++ b/src/sampler/tpe/multivariate.rs @@ -205,7 +205,7 @@ pub enum ConstantLiarStrategy { /// let y = FloatParam::new(-5.0, 5.0); /// /// study -/// .optimize(30, |trial| { +/// .optimize(30, |trial: &mut optimizer::Trial| { /// let xv = x.suggest(trial)?; /// let yv = y.suggest(trial)?; /// Ok::<_, optimizer::Error>(xv * xv + yv * yv) diff --git a/src/storage/journal.rs b/src/storage/journal.rs index 367adfb..1ee477b 100644 --- a/src/storage/journal.rs +++ b/src/storage/journal.rs @@ -42,7 +42,7 @@ //! let storage = JournalStorage::::new("trials.jsonl"); //! let mut study = Study::builder().minimize().storage(storage).build(); //! study -//! .optimize(50, |trial| { +//! .optimize(50, |trial: &mut optimizer::Trial| { //! let x = FloatParam::new(-5.0, 5.0).suggest(trial)?; //! Ok::<_, optimizer::Error>(x * x) //! }) @@ -52,7 +52,7 @@ //! let storage = JournalStorage::::open("trials.jsonl").unwrap(); //! let mut study = Study::builder().minimize().storage(storage).build(); //! study -//! .optimize(50, |trial| { +//! .optimize(50, |trial: &mut optimizer::Trial| { //! let x = FloatParam::new(-5.0, 5.0).suggest(trial)?; //! Ok::<_, optimizer::Error>(x * x) //! }) diff --git a/src/study.rs b/src/study.rs index fd9acb6..dc4c648 100644 --- a/src/study.rs +++ b/src/study.rs @@ -411,22 +411,6 @@ where .map(|t| t.id) } - /// Create a new trial with pre-set parameter values. - /// - /// The trial gets a new unique ID but reuses the given parameters. When - /// `suggest_param` is called on the resulting trial, fixed values are - /// returned instead of sampling. - fn create_trial_with_params(&self, params: HashMap) -> Trial { - let id = self.next_trial_id(); - let mut trial = if let Some(factory) = &self.trial_factory { - factory(id) - } else { - Trial::new(id) - }; - trial.set_fixed_params(params); - trial - } - /// Return the number of enqueued parameter configurations. /// /// See [`enqueue`](Self::enqueue) for how to add configurations. @@ -879,12 +863,15 @@ where completed } - /// Run optimization with a closure. + /// Run optimization with an objective. /// - /// Runs up to `n_trials` evaluations of `objective` sequentially. - /// For lifecycle hooks (early stopping, retries), implement the - /// [`Objective`](crate::Objective) trait and use - /// [`optimize_with`](Self::optimize_with) instead. + /// Accepts any [`Objective`](crate::Objective) implementation, including + /// plain closures (`Fn(&mut Trial) -> Result`) thanks to the + /// blanket impl. Struct-based objectives can override + /// [`before_trial`](crate::Objective::before_trial) and + /// [`after_trial`](crate::Objective::after_trial) for early stopping. + /// + /// Runs up to `n_trials` evaluations sequentially. /// /// # Errors /// @@ -911,10 +898,13 @@ where /// assert!(study.n_trials() > 0); /// assert!(study.best_value().unwrap() >= 0.0); /// ``` - pub fn optimize(&self, n_trials: usize, mut objective: F) -> crate::Result<()> + #[allow(clippy::needless_pass_by_value)] + pub fn optimize( + &self, + n_trials: usize, + objective: impl crate::objective::Objective, + ) -> crate::Result<()> where - F: FnMut(&mut Trial) -> Result, - E: ToString + 'static, V: Clone + Default, { #[cfg(feature = "tracing")] @@ -922,13 +912,44 @@ where tracing::info_span!("optimize", n_trials, direction = ?self.direction).entered(); for _ in 0..n_trials { + if let ControlFlow::Break(()) = objective.before_trial(self) { + break; + } + let mut trial = self.create_trial(); - match objective(&mut trial) { + match objective.evaluate(&mut trial) { Ok(value) => { #[cfg(feature = "tracing")] let trial_id = trial.id(); self.complete_trial(trial, value); - trace_info!(trial_id, "trial completed"); + + #[cfg(feature = "tracing")] + { + tracing::info!(trial_id, "trial completed"); + let trials = self.storage.trials_arc().read(); + if trials + .iter() + .filter(|t| t.state == TrialState::Complete) + .count() + == 1 + || trials.last().map(|t| t.id) == self.best_id(&trials) + { + tracing::info!(trial_id, "new best value found"); + } + } + + // Fire after_trial hook + let trials = self.storage.trials_arc().read(); + if let Some(completed) = trials.last() { + let completed_clone = completed.clone(); + drop(trials); + if let ControlFlow::Break(()) = + objective.after_trial(self, &completed_clone) + { + // Return early — at least one trial completed. + return Ok(()); + } + } } Err(e) if is_trial_pruned(&e) => { #[cfg(feature = "tracing")] @@ -945,144 +966,6 @@ where } } - let has_complete = self - .storage - .trials_arc() - .read() - .iter() - .any(|t| t.state == TrialState::Complete); - if !has_complete { - return Err(crate::Error::NoCompletedTrials); - } - - Ok(()) - } - - /// Run optimization with an [`Objective`](crate::Objective) implementation. - /// - /// Like [`optimize`](Self::optimize), but accepts a struct implementing - /// [`Objective`](crate::Objective) for lifecycle hooks - /// ([`before_trial`](crate::Objective::before_trial), - /// [`after_trial`](crate::Objective::after_trial)) and automatic retries - /// ([`max_retries`](crate::Objective::max_retries)). - /// - /// # Errors - /// - /// Returns `Error::NoCompletedTrials` if no trials completed successfully. - /// - /// # Examples - /// - /// ``` - /// use std::ops::ControlFlow; - /// - /// use optimizer::prelude::*; - /// - /// struct QuadraticObj { - /// x: FloatParam, - /// target: f64, - /// } - /// - /// impl Objective for QuadraticObj { - /// type Error = Error; - /// fn evaluate(&self, trial: &mut Trial) -> Result { - /// let v = self.x.suggest(trial)?; - /// Ok((v - 3.0).powi(2)) - /// } - /// fn after_trial(&self, _: &Study, t: &CompletedTrial) -> ControlFlow<()> { - /// if t.value < self.target { - /// ControlFlow::Break(()) - /// } else { - /// ControlFlow::Continue(()) - /// } - /// } - /// } - /// - /// let study: Study = Study::new(Direction::Minimize); - /// let obj = QuadraticObj { - /// x: FloatParam::new(-10.0, 10.0), - /// target: 1.0, - /// }; - /// study.optimize_with(200, obj).unwrap(); - /// assert!(study.best_value().unwrap() < 1.0); - /// ``` - #[allow(clippy::needless_pass_by_value)] - pub fn optimize_with( - &self, - n_trials: usize, - objective: impl crate::objective::Objective, - ) -> crate::Result<()> - where - V: Clone + Default, - { - #[cfg(feature = "tracing")] - let _span = - tracing::info_span!("optimize_with", n_trials, direction = ?self.direction).entered(); - - let max_retries = objective.max_retries(); - - for _ in 0..n_trials { - if let ControlFlow::Break(()) = objective.before_trial(self) { - break; - } - - let mut trial = self.create_trial(); - let mut retries = 0; - loop { - match objective.evaluate(&mut trial) { - Ok(value) => { - #[cfg(feature = "tracing")] - let trial_id = trial.id(); - self.complete_trial(trial, value); - - #[cfg(feature = "tracing")] - { - tracing::info!(trial_id, "trial completed"); - let trials = self.storage.trials_arc().read(); - if trials - .iter() - .filter(|t| t.state == TrialState::Complete) - .count() - == 1 - || trials.last().map(|t| t.id) == self.best_id(&trials) - { - tracing::info!(trial_id, "new best value found"); - } - } - - // Fire after_trial hook - let trials = self.storage.trials_arc().read(); - if let Some(completed) = trials.last() { - let completed_clone = completed.clone(); - drop(trials); - if let ControlFlow::Break(()) = - objective.after_trial(self, &completed_clone) - { - // Return early — at least one trial completed. - return Ok(()); - } - } - break; - } - Err(e) if !is_trial_pruned(&e) && retries < max_retries => { - retries += 1; - trial = self.create_trial_with_params(trial.params().clone()); - } - Err(e) => { - #[cfg(feature = "tracing")] - let trial_id = trial.id(); - if is_trial_pruned(&e) { - self.prune_trial(trial); - trace_info!(trial_id, "trial pruned"); - } else { - self.fail_trial(trial, e.to_string()); - trace_debug!(trial_id, "trial failed"); - } - break; - } - } - } - } - // Return error if no trials completed successfully let has_complete = self .storage @@ -1097,13 +980,14 @@ where Ok(()) } - /// Run async optimization with a closure. + /// Run async optimization with an objective. /// - /// Each evaluation is wrapped in + /// Like [`optimize`](Self::optimize), but each evaluation is wrapped in /// [`spawn_blocking`](tokio::task::spawn_blocking), keeping the async /// runtime responsive for CPU-bound objectives. Trials run sequentially. /// - /// For lifecycle hooks, use [`optimize_with_async`](Self::optimize_with_async). + /// Accepts any [`Objective`](crate::Objective) implementation, including + /// plain closures. Struct-based objectives can override lifecycle hooks. /// /// # Errors /// @@ -1135,10 +1019,10 @@ where /// # } /// ``` #[cfg(feature = "async")] - pub async fn optimize_async(&self, n_trials: usize, objective: F) -> crate::Result<()> + pub async fn optimize_async(&self, n_trials: usize, objective: O) -> crate::Result<()> where - F: Fn(&mut Trial) -> Result + Send + Sync + 'static, - E: ToString + Send + 'static, + O: crate::objective::Objective + Send + Sync + 'static, + O::Error: Send, V: Clone + Default + Send + 'static, { #[cfg(feature = "tracing")] @@ -1148,10 +1032,14 @@ where let objective = Arc::new(objective); for _ in 0..n_trials { + if let ControlFlow::Break(()) = objective.before_trial(self) { + break; + } + let obj = Arc::clone(&objective); let mut trial = self.create_trial(); let result = tokio::task::spawn_blocking(move || { - let res = obj(&mut trial); + let res = obj.evaluate(&mut trial); (trial, res) }) .await @@ -1163,6 +1051,18 @@ where let trial_id = t.id(); self.complete_trial(t, value); trace_info!(trial_id, "trial completed"); + + // Fire after_trial hook + let trials = self.storage.trials_arc().read(); + if let Some(completed) = trials.last() { + let completed_clone = completed.clone(); + drop(trials); + if let ControlFlow::Break(()) = + objective.after_trial(self, &completed_clone) + { + return Ok(()); + } + } } (t, Err(e)) if is_trial_pruned(&e) => { #[cfg(feature = "tracing")] @@ -1192,108 +1092,16 @@ where Ok(()) } - /// Run async optimization with an [`Objective`](crate::Objective) implementation. - /// - /// Like [`optimize_async`](Self::optimize_async), but accepts a struct - /// implementing [`Objective`](crate::Objective) for lifecycle hooks and - /// automatic retries. - /// - /// # Errors - /// - /// Returns `Error::NoCompletedTrials` if no trials completed successfully. - /// Returns `Error::TaskError` if a spawned blocking task panics. - #[cfg(feature = "async")] - pub async fn optimize_with_async(&self, n_trials: usize, objective: O) -> crate::Result<()> - where - O: crate::objective::Objective + Send + Sync + 'static, - O::Error: Send, - V: Clone + Default + Send + 'static, - { - #[cfg(feature = "tracing")] - let _span = - tracing::info_span!("optimize_with_async", n_trials, direction = ?self.direction) - .entered(); - - let objective = Arc::new(objective); - let max_retries = objective.max_retries(); - - for _ in 0..n_trials { - if let ControlFlow::Break(()) = objective.before_trial(self) { - break; - } - - let mut trial = self.create_trial(); - let mut retries = 0; - loop { - let obj = Arc::clone(&objective); - let result = tokio::task::spawn_blocking(move || { - let res = obj.evaluate(&mut trial); - (trial, res) - }) - .await - .map_err(|e| crate::Error::TaskError(e.to_string()))?; - - match result { - (t, Ok(value)) => { - #[cfg(feature = "tracing")] - let trial_id = t.id(); - self.complete_trial(t, value); - trace_info!(trial_id, "trial completed"); - - // Fire after_trial hook - let trials = self.storage.trials_arc().read(); - if let Some(completed) = trials.last() { - let completed_clone = completed.clone(); - drop(trials); - if let ControlFlow::Break(()) = - objective.after_trial(self, &completed_clone) - { - return Ok(()); - } - } - break; - } - (t, Err(e)) if !is_trial_pruned(&e) && retries < max_retries => { - retries += 1; - trial = self.create_trial_with_params(t.params().clone()); - } - (t, Err(e)) => { - #[cfg(feature = "tracing")] - let trial_id = t.id(); - if is_trial_pruned(&e) { - self.prune_trial(t); - trace_info!(trial_id, "trial pruned"); - } else { - self.fail_trial(t, e.to_string()); - trace_debug!(trial_id, "trial failed"); - } - break; - } - } - } - } - - let has_complete = self - .storage - .trials_arc() - .read() - .iter() - .any(|t| t.state == TrialState::Complete); - if !has_complete { - return Err(crate::Error::NoCompletedTrials); - } - - Ok(()) - } - - /// Run parallel optimization with a closure. + /// Run parallel optimization with an objective. /// /// Spawns up to `concurrency` evaluations concurrently using /// [`spawn_blocking`](tokio::task::spawn_blocking). Results are /// collected via a [`JoinSet`](tokio::task::JoinSet). /// - /// For lifecycle hooks, use - /// [`optimize_with_parallel`](Self::optimize_with_parallel). + /// Accepts any [`Objective`](crate::Objective) implementation, including + /// plain closures. The [`after_trial`](crate::Objective::after_trial) + /// hook fires as each result arrives — returning `Break` stops spawning + /// new trials while in-flight tasks drain. /// /// # Errors /// @@ -1325,131 +1133,8 @@ where /// # } /// ``` #[cfg(feature = "async")] - #[allow(clippy::missing_panics_doc)] - pub async fn optimize_parallel( - &self, - n_trials: usize, - concurrency: usize, - objective: F, - ) -> crate::Result<()> - where - F: Fn(&mut Trial) -> Result + Send + Sync + 'static, - E: ToString + Send + 'static, - V: Clone + Default + Send + 'static, - { - use tokio::sync::Semaphore; - use tokio::task::JoinSet; - - #[cfg(feature = "tracing")] - let _span = tracing::info_span!("optimize_parallel", n_trials, concurrency, direction = ?self.direction).entered(); - - let objective = Arc::new(objective); - let semaphore = Arc::new(Semaphore::new(concurrency)); - let mut join_set: JoinSet<(Trial, Result)> = JoinSet::new(); - let mut spawned = 0; - - while spawned < n_trials { - // If the join set is full, drain one result to free a slot. - while join_set.len() >= concurrency { - let result = join_set - .join_next() - .await - .expect("join_set should not be empty") - .map_err(|e| crate::Error::TaskError(e.to_string()))?; - match result { - (t, Ok(value)) => { - #[cfg(feature = "tracing")] - let trial_id = t.id(); - self.complete_trial(t, value); - trace_info!(trial_id, "trial completed"); - } - (t, Err(e)) => { - #[cfg(feature = "tracing")] - let trial_id = t.id(); - if is_trial_pruned(&e) { - self.prune_trial(t); - trace_info!(trial_id, "trial pruned"); - } else { - self.fail_trial(t, e.to_string()); - trace_debug!(trial_id, "trial failed"); - } - } - } - } - - let permit = semaphore - .clone() - .acquire_owned() - .await - .map_err(|e| crate::Error::TaskError(e.to_string()))?; - - let mut trial = self.create_trial(); - let obj = Arc::clone(&objective); - join_set.spawn(async move { - let result = tokio::task::spawn_blocking(move || { - let res = obj(&mut trial); - (trial, res) - }) - .await - .expect("spawn_blocking should not panic"); - drop(permit); - result - }); - spawned += 1; - } - - // Drain remaining in-flight tasks. - while let Some(result) = join_set.join_next().await { - let result = result.map_err(|e| crate::Error::TaskError(e.to_string()))?; - match result { - (t, Ok(value)) => { - #[cfg(feature = "tracing")] - let trial_id = t.id(); - self.complete_trial(t, value); - trace_info!(trial_id, "trial completed"); - } - (t, Err(e)) => { - #[cfg(feature = "tracing")] - let trial_id = t.id(); - if is_trial_pruned(&e) { - self.prune_trial(t); - trace_info!(trial_id, "trial pruned"); - } else { - self.fail_trial(t, e.to_string()); - trace_debug!(trial_id, "trial failed"); - } - } - } - } - - let has_complete = self - .storage - .trials_arc() - .read() - .iter() - .any(|t| t.state == TrialState::Complete); - if !has_complete { - return Err(crate::Error::NoCompletedTrials); - } - - Ok(()) - } - - /// Run parallel optimization with an [`Objective`](crate::Objective) implementation. - /// - /// Like [`optimize_parallel`](Self::optimize_parallel), but accepts a struct - /// implementing [`Objective`](crate::Objective) for lifecycle hooks and - /// automatic retries. The [`after_trial`](crate::Objective::after_trial) - /// hook fires as each result arrives — returning `Break` stops spawning - /// new trials while in-flight tasks drain. - /// - /// # Errors - /// - /// Returns `Error::NoCompletedTrials` if no trials completed successfully. - /// Returns `Error::TaskError` if the semaphore is closed or a spawned task panics. - #[cfg(feature = "async")] #[allow(clippy::missing_panics_doc, clippy::too_many_lines)] - pub async fn optimize_with_parallel( + pub async fn optimize_parallel( &self, n_trials: usize, concurrency: usize, @@ -1464,7 +1149,7 @@ where use tokio::task::JoinSet; #[cfg(feature = "tracing")] - let _span = tracing::info_span!("optimize_with_parallel", n_trials, concurrency, direction = ?self.direction).entered(); + let _span = tracing::info_span!("optimize_parallel", n_trials, concurrency, direction = ?self.direction).entered(); let objective = Arc::new(objective); let semaphore = Arc::new(Semaphore::new(concurrency)); @@ -1838,7 +1523,7 @@ where /// let x = FloatParam::new(0.0, 10.0).name("x"); /// /// study - /// .optimize(20, |trial| { + /// .optimize(20, |trial: &mut optimizer::Trial| { /// let xv = x.suggest(trial)?; /// Ok::<_, optimizer::Error>(xv * xv) /// }) @@ -1941,7 +1626,7 @@ where /// let y = FloatParam::new(0.0, 10.0).name("y"); /// /// study - /// .optimize(30, |trial| { + /// .optimize(30, |trial: &mut optimizer::Trial| { /// let xv = x.suggest(trial)?; /// let yv = y.suggest(trial)?; /// Ok::<_, optimizer::Error>(xv * xv + 0.1 * yv) diff --git a/src/visualization.rs b/src/visualization.rs index f607f25..3c9558d 100644 --- a/src/visualization.rs +++ b/src/visualization.rs @@ -26,7 +26,7 @@ //! //! let study: Study = Study::new(Direction::Minimize); //! # let x = FloatParam::new(0.0, 1.0); -//! # study.optimize(10, |trial| { +//! # study.optimize(10, |trial: &mut optimizer::Trial| { //! # let v = x.suggest(trial)?; //! # Ok::<_, optimizer::Error>(v * v) //! # }).unwrap(); diff --git a/tests/export_tests.rs b/tests/export_tests.rs index 9a43149..20b4a5c 100644 --- a/tests/export_tests.rs +++ b/tests/export_tests.rs @@ -18,7 +18,7 @@ fn csv_includes_all_trial_data() { let y = IntParam::new(1, 5).name("y"); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, optimizer::Error>(xv + yv as f64) @@ -124,7 +124,7 @@ fn csv_output_is_parseable() { let layers = IntParam::new(1, 5).name("n_layers"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let l = lr.suggest(trial)?; let n = layers.suggest(trial)?; Ok::<_, optimizer::Error>(l * n as f64) @@ -153,7 +153,7 @@ fn export_csv_writes_file() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(42)); let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(xv * xv) }) @@ -179,7 +179,7 @@ fn export_json_writes_file() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(42)); let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(xv * xv) }) @@ -214,7 +214,7 @@ fn csv_includes_user_attributes() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(2, |trial| { + .optimize(2, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; trial.set_user_attr("training_time_secs", 45.2); Ok::<_, optimizer::Error>(xv * xv) diff --git a/tests/fanova_tests.rs b/tests/fanova_tests.rs index 88eb3c8..cb69950 100644 --- a/tests/fanova_tests.rs +++ b/tests/fanova_tests.rs @@ -10,7 +10,7 @@ fn fanova_dominant_parameter() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(42)); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let _yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv) @@ -34,7 +34,7 @@ fn fanova_interaction() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(7)); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * yv) @@ -62,7 +62,7 @@ fn fanova_consistent_with_correlation() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(99)); study - .optimize(80, |trial| { + .optimize(80, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(3.0 * xv + 0.5 * yv) diff --git a/tests/journal_tests.rs b/tests/journal_tests.rs index fab5790..4eef94e 100644 --- a/tests/journal_tests.rs +++ b/tests/journal_tests.rs @@ -119,7 +119,7 @@ fn study_with_journal_integration() { let study = Study::with_journal(Direction::Minimize, RandomSampler::with_seed(1), &path).unwrap(); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let val = x.suggest(trial)?; Ok::<_, optimizer::Error>(val * val) }) @@ -134,7 +134,7 @@ fn study_with_journal_integration() { // Continue optimizing study2 - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let val = x.suggest(trial)?; Ok::<_, optimizer::Error>(val * val) }) @@ -158,7 +158,7 @@ fn ids_are_unique_after_reload() { let study = Study::with_journal(Direction::Minimize, RandomSampler::with_seed(1), &path).unwrap(); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let _ = FloatParam::new(0.0, 1.0).suggest(trial)?; Ok::<_, optimizer::Error>(1.0) }) @@ -169,7 +169,7 @@ fn ids_are_unique_after_reload() { let study = Study::with_journal(Direction::Minimize, RandomSampler::with_seed(2), &path).unwrap(); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let _ = FloatParam::new(0.0, 1.0).suggest(trial)?; Ok::<_, optimizer::Error>(1.0) }) @@ -194,7 +194,7 @@ fn pruned_trials_are_stored() { // Complete one, prune one let x = FloatParam::new(0.0, 1.0); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; if trial.id() == 1 { Err(optimizer::TrialPruned)?; diff --git a/tests/multi_objective_tests.rs b/tests/multi_objective_tests.rs index dc30af9..c66e9df 100644 --- a/tests/multi_objective_tests.rs +++ b/tests/multi_objective_tests.rs @@ -18,7 +18,7 @@ fn test_basic_two_objective_random() { let x = FloatParam::new(0.0, 1.0); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) }) @@ -50,7 +50,7 @@ fn test_dimension_mismatch_error() { let study = MultiObjectiveStudy::new(vec![Direction::Minimize, Direction::Minimize]); let x = FloatParam::new(0.0, 1.0); - let result = study.optimize(1, |trial| { + let result = study.optimize(1, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; // Return wrong number of values Ok::<_, optimizer::Error>(vec![xv]) @@ -104,7 +104,7 @@ fn test_n_trials_counting() { let x = FloatParam::new(0.0, 1.0); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) }) @@ -124,7 +124,7 @@ fn test_three_objectives() { let y = FloatParam::new(0.0, 1.0); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, yv, 1.0 - xv - yv]) @@ -150,7 +150,7 @@ fn test_trials_accessor() { let x = FloatParam::new(0.0, 1.0); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) }) @@ -178,7 +178,7 @@ fn test_nsga2_zdt1() { MultiObjectiveStudy::with_sampler(vec![Direction::Minimize, Direction::Minimize], sampler); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xs: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -224,7 +224,7 @@ fn test_nsga2_with_seed_reproducible() { sampler, ); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, yv]) @@ -256,7 +256,7 @@ fn test_nsga2_builder() { let x = FloatParam::new(0.0, 1.0); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) }) @@ -275,7 +275,7 @@ fn test_nsga2_categorical_params() { let cat = CategoricalParam::new(vec!["a", "b", "c"]); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let cv = cat.suggest(trial)?; let bonus = match cv { @@ -301,7 +301,7 @@ fn test_nsga2_constraints() { let x = FloatParam::new(0.0, 1.0); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; // Constraint: x >= 0.3 (i.e. 0.3 - x <= 0) trial.set_constraints(vec![0.3 - xv]); @@ -326,7 +326,7 @@ fn test_multi_objective_trial_get() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, 10.0 - xv]) }) @@ -345,7 +345,7 @@ fn test_multi_objective_trial_is_feasible() { let x = FloatParam::new(0.0, 1.0); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; trial.set_constraints(vec![0.5 - xv]); // feasible if x >= 0.5 Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) @@ -369,7 +369,7 @@ fn test_multi_objective_trial_user_attrs() { let x = FloatParam::new(0.0, 1.0); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; trial.set_user_attr("iteration", 42_i64); Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) @@ -409,7 +409,7 @@ fn test_nsga3_zdt1() { MultiObjectiveStudy::with_sampler(vec![Direction::Minimize, Direction::Minimize], sampler); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xs: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -458,7 +458,7 @@ fn test_nsga3_four_objectives() { let study = MultiObjectiveStudy::with_sampler(directions, sampler); study - .optimize(500, |trial| { + .optimize(500, |trial: &mut optimizer::Trial| { let xs: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -506,7 +506,7 @@ fn test_nsga3_reproducible() { sampler, ); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, yv]) @@ -539,7 +539,7 @@ fn test_nsga3_builder() { let x = FloatParam::new(0.0, 1.0); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) }) @@ -557,7 +557,7 @@ fn test_nsga3_constraints() { let x = FloatParam::new(0.0, 1.0); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; trial.set_constraints(vec![0.3 - xv]); Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) @@ -588,7 +588,7 @@ fn test_moead_zdt1_tchebycheff() { MultiObjectiveStudy::with_sampler(vec![Direction::Minimize, Direction::Minimize], sampler); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xs: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -635,7 +635,7 @@ fn test_moead_zdt1_weighted_sum() { MultiObjectiveStudy::with_sampler(vec![Direction::Minimize, Direction::Minimize], sampler); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xs: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -666,7 +666,7 @@ fn test_moead_zdt1_pbi() { MultiObjectiveStudy::with_sampler(vec![Direction::Minimize, Direction::Minimize], sampler); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xs: Vec = params .iter() .map(|p| p.suggest(trial)) @@ -695,7 +695,7 @@ fn test_moead_reproducible() { sampler, ); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, yv]) @@ -729,7 +729,7 @@ fn test_moead_builder() { let x = FloatParam::new(0.0, 1.0); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(vec![xv, 1.0 - xv]) }) diff --git a/tests/parameter_tests.rs b/tests/parameter_tests.rs index 37e205f..3d4f308 100644 --- a/tests/parameter_tests.rs +++ b/tests/parameter_tests.rs @@ -180,7 +180,7 @@ fn parameter_api_with_study() { let study: Study = Study::new(Direction::Minimize); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let n = n_param.suggest(trial)?; let dropout = dropout_param.suggest(trial)?; diff --git a/tests/sampler/bohb.rs b/tests/sampler/bohb.rs index 980de7c..f9c24de 100644 --- a/tests/sampler/bohb.rs +++ b/tests/sampler/bohb.rs @@ -27,7 +27,7 @@ fn bohb_converges_on_quadratic() { let x_param = FloatParam::new(-10.0, 10.0); study - .optimize(60, |trial| { + .optimize(60, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; // Report intermediate values at budget steps 1, 3, 9 @@ -66,7 +66,7 @@ fn bohb_with_pruning() { let x_param = FloatParam::new(-5.0, 5.0); study - .optimize(40, |trial| { + .optimize(40, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let obj = x * x; @@ -112,7 +112,7 @@ fn bohb_uses_budget_conditioned_history() { let x_param = FloatParam::new(0.0, 10.0); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; // Intermediate values that guide optimization toward x=2 trial.report(1, (x - 2.0).powi(2) + 1.0); diff --git a/tests/sampler/cma_es.rs b/tests/sampler/cma_es.rs index 2c3f992..711cccd 100644 --- a/tests/sampler/cma_es.rs +++ b/tests/sampler/cma_es.rs @@ -10,7 +10,7 @@ fn sphere_function() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -34,7 +34,7 @@ fn rosenbrock_function() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(300, |trial| { + .optimize(300, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; let val = (1.0 - xv).powi(2) + 100.0 * (yv - xv * xv).powi(2); @@ -60,7 +60,7 @@ fn bounds_respected() { let y = FloatParam::new(0.0, 10.0).name("y"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv + yv) @@ -84,7 +84,7 @@ fn mixed_params_float_and_categorical() { let cat = CategoricalParam::new(vec!["a", "b", "c"]).name("cat"); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let cv = cat.suggest(trial)?; let penalty = match cv { @@ -114,7 +114,7 @@ fn seeded_reproducibility() { let sampler = CmaEsSampler::with_seed(seed); let study: Study = Study::with_sampler(Direction::Minimize, sampler); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -137,7 +137,7 @@ fn different_seeds_different_results() { let sampler = CmaEsSampler::with_seed(seed); let study: Study = Study::with_sampler(Direction::Minimize, sampler); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -162,7 +162,7 @@ fn single_dimension() { let x = FloatParam::new(-10.0, 10.0).name("x"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, Error>((xv - 3.0).powi(2)) }) @@ -184,7 +184,7 @@ fn integer_params() { let n = IntParam::new(1, 20).name("n"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let nv = n.suggest(trial)?; // Minimum at n = 10 Ok::<_, Error>(((nv - 10) * (nv - 10)) as f64) @@ -212,7 +212,7 @@ fn log_scale_params() { let lr = FloatParam::new(1e-5, 1.0).log_scale().name("lr"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let lrv = lr.suggest(trial)?; // Minimum at lr = 0.01 Ok::<_, Error>((lrv.ln() - 0.01_f64.ln()).powi(2)) @@ -241,7 +241,7 @@ fn custom_population_size_and_sigma() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) diff --git a/tests/sampler/differential_evolution.rs b/tests/sampler/differential_evolution.rs index eb7a927..4c75d5f 100644 --- a/tests/sampler/differential_evolution.rs +++ b/tests/sampler/differential_evolution.rs @@ -10,7 +10,7 @@ fn sphere_function() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -37,7 +37,7 @@ fn rosenbrock_function() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(400, |trial| { + .optimize(400, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; let val = (1.0 - xv).powi(2) + 100.0 * (yv - xv * xv).powi(2); @@ -67,7 +67,7 @@ fn rastrigin_function() { let y = FloatParam::new(-5.12, 5.12).name("y"); study - .optimize(500, |trial| { + .optimize(500, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; let val = 20.0 @@ -95,7 +95,7 @@ fn bounds_respected() { let y = FloatParam::new(0.0, 10.0).name("y"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv + yv) @@ -123,7 +123,7 @@ fn strategy_best1() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -151,7 +151,7 @@ fn strategy_current_to_best1() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(200, |trial| { + .optimize(200, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -175,7 +175,7 @@ fn mixed_params_float_and_categorical() { let cat = CategoricalParam::new(vec!["a", "b", "c"]).name("cat"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let cv = cat.suggest(trial)?; let penalty = match cv { @@ -204,7 +204,7 @@ fn seeded_reproducibility() { let sampler = DifferentialEvolutionSampler::with_seed(seed); let study: Study = Study::with_sampler(Direction::Minimize, sampler); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -227,7 +227,7 @@ fn different_seeds_different_results() { let sampler = DifferentialEvolutionSampler::with_seed(seed); let study: Study = Study::with_sampler(Direction::Minimize, sampler); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -252,7 +252,7 @@ fn single_dimension() { let x = FloatParam::new(-10.0, 10.0).name("x"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, Error>((xv - 3.0).powi(2)) }) @@ -274,7 +274,7 @@ fn integer_params() { let n = IntParam::new(1, 20).name("n"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let nv = n.suggest(trial)?; // Minimum at n = 10 Ok::<_, Error>(((nv - 10) * (nv - 10)) as f64) @@ -302,7 +302,7 @@ fn log_scale_params() { let lr = FloatParam::new(1e-5, 1.0).log_scale().name("lr"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let lrv = lr.suggest(trial)?; // Minimum at lr = 0.01 Ok::<_, Error>((lrv.ln() - 0.01_f64.ln()).powi(2)) @@ -332,7 +332,7 @@ fn custom_mutation_and_crossover() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) diff --git a/tests/sampler/gp.rs b/tests/sampler/gp.rs index 6ddf3b5..24b051d 100644 --- a/tests/sampler/gp.rs +++ b/tests/sampler/gp.rs @@ -10,7 +10,7 @@ fn sphere_function() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(80, |trial| { + .optimize(80, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -34,7 +34,7 @@ fn rosenbrock_function() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; let val = (1.0 - xv).powi(2) + 100.0 * (yv - xv * xv).powi(2); @@ -59,7 +59,7 @@ fn bounds_respected() { let y = FloatParam::new(0.0, 10.0).name("y"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv + yv) @@ -83,7 +83,7 @@ fn mixed_params_float_and_categorical() { let cat = CategoricalParam::new(vec!["a", "b", "c"]).name("cat"); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let cv = cat.suggest(trial)?; let penalty = match cv { @@ -112,7 +112,7 @@ fn seeded_reproducibility() { let sampler = GpSampler::with_seed(seed); let study: Study = Study::with_sampler(Direction::Minimize, sampler); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -135,7 +135,7 @@ fn different_seeds_different_results() { let sampler = GpSampler::with_seed(seed); let study: Study = Study::with_sampler(Direction::Minimize, sampler); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) @@ -160,7 +160,7 @@ fn single_dimension() { let x = FloatParam::new(-10.0, 10.0).name("x"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, Error>((xv - 3.0).powi(2)) }) @@ -182,7 +182,7 @@ fn integer_params() { let n = IntParam::new(1, 20).name("n"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let nv = n.suggest(trial)?; Ok::<_, Error>(((nv - 10) * (nv - 10)) as f64) }) @@ -209,7 +209,7 @@ fn log_scale_params() { let lr = FloatParam::new(1e-5, 1.0).log_scale().name("lr"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let lrv = lr.suggest(trial)?; Ok::<_, Error>((lrv.ln() - 0.01_f64.ln()).powi(2)) }) @@ -238,7 +238,7 @@ fn builder_configuration() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, Error>(xv * xv + yv * yv) diff --git a/tests/sampler/multivariate_tpe.rs b/tests/sampler/multivariate_tpe.rs index db84523..a84ca34 100644 --- a/tests/sampler/multivariate_tpe.rs +++ b/tests/sampler/multivariate_tpe.rs @@ -55,7 +55,7 @@ fn test_multivariate_tpe_rosenbrock_finds_good_solution() { let y_param = FloatParam::new(-2.0, 4.0); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(rosenbrock(x, y)) @@ -90,7 +90,7 @@ fn test_independent_tpe_rosenbrock() { let y_param = FloatParam::new(-2.0, 4.0); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(rosenbrock(x, y)) @@ -133,7 +133,7 @@ fn test_multivariate_tpe_outperforms_on_correlated_problem() { let y_param = FloatParam::new(-2.0, 4.0); study - .optimize(n_trials, |trial| { + .optimize(n_trials, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(rosenbrock(x, y)) @@ -156,7 +156,7 @@ fn test_multivariate_tpe_outperforms_on_correlated_problem() { let y_param = FloatParam::new(-2.0, 4.0); study - .optimize(n_trials, |trial| { + .optimize(n_trials, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(rosenbrock(x, y)) @@ -218,7 +218,7 @@ fn test_multivariate_tpe_independent_problem() { let y_param = FloatParam::new(-5.0, 5.0); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(sphere(x, y)) @@ -250,7 +250,7 @@ fn test_independent_tpe_independent_problem() { let y_param = FloatParam::new(-5.0, 5.0); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(sphere(x, y)) @@ -290,7 +290,7 @@ fn test_both_samplers_work_on_independent_problem() { let y_param = FloatParam::new(-5.0, 5.0); study - .optimize(n_trials, |trial| { + .optimize(n_trials, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(sphere(x, y)) @@ -312,7 +312,7 @@ fn test_both_samplers_work_on_independent_problem() { let y_param = FloatParam::new(-5.0, 5.0); study - .optimize(n_trials, |trial| { + .optimize(n_trials, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(sphere(x, y)) @@ -360,7 +360,7 @@ fn test_multivariate_tpe_with_group_decomposition() { let y_param = FloatParam::new(-5.0, 5.0); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(sphere(x, y)) @@ -396,7 +396,7 @@ fn test_multivariate_tpe_mixed_parameter_types() { let mode_param = CategoricalParam::new(vec!["a", "b", "c"]); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let n = n_param.suggest(trial)?; let mode = mode_param.suggest(trial)?; diff --git a/tests/sampler/random.rs b/tests/sampler/random.rs index 911f192..66a9264 100644 --- a/tests/sampler/random.rs +++ b/tests/sampler/random.rs @@ -1,3 +1,5 @@ +use std::cell::RefCell; + use optimizer::parameter::{CategoricalParam, FloatParam, IntParam, Parameter}; use optimizer::sampler::random::RandomSampler; use optimizer::{Direction, Error, Study}; @@ -7,18 +9,20 @@ fn test_random_sampler_uniform_float_distribution() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(42)); let n_samples = 1000; - let mut samples = Vec::with_capacity(n_samples); + let samples = RefCell::new(Vec::with_capacity(n_samples)); let x_param = FloatParam::new(0.0, 1.0); study - .optimize(n_samples, |trial| { + .optimize(n_samples, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; - samples.push(x); + samples.borrow_mut().push(x); Ok::<_, Error>(x) }) .unwrap(); + let mut samples = samples.into_inner(); + // All samples should be in range for &s in &samples { assert!((0.0..=1.0).contains(&s), "sample {s} out of range [0, 1]"); @@ -44,19 +48,20 @@ fn test_random_sampler_uniform_int_distribution() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(123)); let n_samples = 5000; - let mut counts = [0u32; 10]; // counts for values 1-10 + let counts = RefCell::new([0u32; 10]); // counts for values 1-10 let n_param = IntParam::new(1, 10); study - .optimize(n_samples, |trial| { + .optimize(n_samples, |trial: &mut optimizer::Trial| { let n = n_param.suggest(trial)?; assert!((1..=10).contains(&n), "sample {n} out of range [1, 10]"); - counts[(n - 1) as usize] += 1; + counts.borrow_mut()[(n - 1) as usize] += 1; Ok::<_, Error>(n as f64) }) .unwrap(); + let counts = counts.into_inner(); let expected = n_samples as f64 / 10.0; for (i, &count) in counts.iter().enumerate() { let diff = (count as f64 - expected).abs() / expected; @@ -76,20 +81,21 @@ fn test_random_sampler_uniform_categorical_distribution() { let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(456)); let n_samples = 2000; - let mut counts = [0u32; 4]; + let counts = RefCell::new([0u32; 4]); let choices = ["a", "b", "c", "d"]; let cat_param = CategoricalParam::new(choices.to_vec()); study - .optimize(n_samples, |trial| { + .optimize(n_samples, |trial: &mut optimizer::Trial| { let choice = cat_param.suggest(trial)?; let idx = choices.iter().position(|&c| c == choice).unwrap(); - counts[idx] += 1; + counts.borrow_mut()[idx] += 1; Ok::<_, Error>(idx as f64) }) .unwrap(); + let counts = counts.into_inner(); let expected = n_samples as f64 / 4.0; for (i, &count) in counts.iter().enumerate() { let diff = (count as f64 - expected).abs() / expected; @@ -111,28 +117,30 @@ fn test_random_sampler_reproducibility() { let study2: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(999)); - let mut values1 = Vec::new(); - let mut values2 = Vec::new(); + let values1 = RefCell::new(Vec::new()); + let values2 = RefCell::new(Vec::new()); let x_param1 = FloatParam::new(0.0, 100.0); let x_param2 = FloatParam::new(0.0, 100.0); study1 - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x = x_param1.suggest(trial)?; - values1.push(x); + values1.borrow_mut().push(x); Ok::<_, Error>(x) }) .unwrap(); study2 - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x = x_param2.suggest(trial)?; - values2.push(x); + values2.borrow_mut().push(x); Ok::<_, Error>(x) }) .unwrap(); + let values1 = values1.into_inner(); + let values2 = values2.into_inner(); for (i, (v1, v2)) in values1.iter().zip(values2.iter()).enumerate() { assert_eq!( v1, v2, diff --git a/tests/sampler/tpe.rs b/tests/sampler/tpe.rs index 8091077..3e37899 100644 --- a/tests/sampler/tpe.rs +++ b/tests/sampler/tpe.rs @@ -18,7 +18,7 @@ fn test_tpe_optimizes_quadratic_function() { let x_param = FloatParam::new(-10.0, 10.0); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>((x - 3.0).powi(2)) }) @@ -51,7 +51,7 @@ fn test_tpe_optimizes_multivariate_function() { let y_param = FloatParam::new(-5.0, 5.0); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let y = y_param.suggest(trial)?; Ok::<_, Error>(x * x + y * y) @@ -83,7 +83,7 @@ fn test_tpe_maximization() { let x_param = FloatParam::new(-10.0, 10.0); study - .optimize(50, |trial| { + .optimize(50, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(-(x - 2.0).powi(2) + 10.0) }) @@ -113,7 +113,7 @@ fn test_tpe_with_categorical_parameter() { // Optimization where the best choice depends on the categorical study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let choice = model_param.suggest(trial)?; let x = x_param.suggest(trial)?; @@ -150,7 +150,7 @@ fn test_tpe_with_integer_parameters() { // Minimize (n - 7)^2 where n in [1, 10] study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let n = n_param.suggest(trial)?; Ok::<_, Error>(((n - 7) as f64).powi(2)) }) @@ -177,7 +177,7 @@ fn test_tpe_with_log_scale_int() { let batch_param = IntParam::new(1, 1024).log_scale(); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let batch_size = batch_param.suggest(trial)?; Ok::<_, Error>(((batch_size as f64).log2() - 5.0).powi(2)) }) @@ -200,7 +200,7 @@ fn test_tpe_with_step_distributions() { let n_param = IntParam::new(0, 100).step(10); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let n = n_param.suggest(trial)?; Ok::<_, Error>((x - 5.0).powi(2) + ((n - 50) as f64).powi(2)) @@ -224,7 +224,7 @@ fn test_tpe_with_fixed_kde_bandwidth() { let x_param = FloatParam::new(-5.0, 5.0); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x * x) }) @@ -252,7 +252,7 @@ fn test_tpe_split_trials_with_two_trials() { let x_param = FloatParam::new(0.0, 10.0); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x) }) @@ -276,7 +276,7 @@ fn test_tpe_empty_good_or_bad_values_fallback() { // First optimize with one parameter study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x) }) @@ -284,7 +284,7 @@ fn test_tpe_empty_good_or_bad_values_fallback() { // Now try with a different parameter - TPE won't have history for "y" study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let y = y_param.suggest(trial)?; Ok::<_, Error>(y) }) @@ -304,7 +304,7 @@ fn test_tpe_sampler_builder_default_trait() { let x_param = FloatParam::new(0.0, 1.0); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x) }) @@ -321,7 +321,7 @@ fn test_tpe_sampler_default_trait() { let x_param = FloatParam::new(0.0, 1.0); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x) }) @@ -343,7 +343,7 @@ fn test_suggest_bool_with_tpe() { let x_param = FloatParam::new(0.0, 10.0); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let use_large = use_large_param.suggest(trial)?; let x = x_param.suggest(trial)?; // The value depends on use_large flag @@ -369,7 +369,7 @@ fn test_params_with_tpe() { let n_param = IntParam::new(1, 10); study - .optimize(30, |trial| { + .optimize(30, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let n = n_param.suggest(trial)?; Ok::<_, Error>(x * x + (n as f64 - 5.0).powi(2)) diff --git a/tests/serde_tests.rs b/tests/serde_tests.rs index b734152..1495710 100644 --- a/tests/serde_tests.rs +++ b/tests/serde_tests.rs @@ -13,7 +13,7 @@ fn round_trip_save_load() { let n = IntParam::new(1, 100).name("n"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x_val = x.suggest(trial)?; let n_val = n.suggest(trial)?; Ok::<_, optimizer::Error>(x_val * x_val + n_val as f64) @@ -51,7 +51,7 @@ fn json_output_is_human_readable() { let x = FloatParam::new(0.0, 1.0).name("x"); study - .optimize(2, |trial| { + .optimize(2, |trial: &mut optimizer::Trial| { let v = x.suggest(trial)?; Ok::<_, optimizer::Error>(v) }) @@ -153,7 +153,7 @@ fn round_trip_preserves_trial_id_counter() { let x = FloatParam::new(0.0, 1.0); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let v = x.suggest(trial)?; Ok::<_, optimizer::Error>(v) }) @@ -182,7 +182,7 @@ fn save_and_resume_continues_trial_ids() { // Run 10 trials study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let v = x.suggest(trial)?; Ok::<_, optimizer::Error>(v * v) }) @@ -196,7 +196,7 @@ fn save_and_resume_continues_trial_ids() { // Continue with 5 more trials let remaining = 15 - loaded.n_trials(); loaded - .optimize(remaining, |trial| { + .optimize(remaining, |trial: &mut optimizer::Trial| { let v = x.suggest(trial)?; Ok::<_, optimizer::Error>(v * v) }) @@ -223,7 +223,7 @@ fn save_uses_atomic_write() { let save_path = dir.join("atomic.json"); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let v = x.suggest(trial)?; Ok::<_, optimizer::Error>(v) }) diff --git a/tests/study/builder.rs b/tests/study/builder.rs index 8b421ff..fc25a33 100644 --- a/tests/study/builder.rs +++ b/tests/study/builder.rs @@ -33,7 +33,7 @@ fn test_builder_with_sampler() { let study: Study = Study::builder().sampler(TpeSampler::new()).build(); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let val = x.suggest(trial)?; Ok::<_, Error>(val * val) }) @@ -77,7 +77,7 @@ fn test_builder_optimizes_correctly() { .build(); study - .optimize(100, |trial| { + .optimize(100, |trial: &mut optimizer::Trial| { let val = x.suggest(trial)?; Ok::<_, Error>((val - 3.0) * (val - 3.0)) }) diff --git a/tests/study/enqueue.rs b/tests/study/enqueue.rs index c84e008..8e20f37 100644 --- a/tests/study/enqueue.rs +++ b/tests/study/enqueue.rs @@ -1,3 +1,4 @@ +use std::cell::RefCell; use std::collections::HashMap; use optimizer::parameter::{FloatParam, IntParam, ParamValue, Parameter}; @@ -73,16 +74,17 @@ fn test_enqueue_with_optimize() { study.enqueue(HashMap::from([(x.id(), ParamValue::Float(1.0))])); study.enqueue(HashMap::from([(x.id(), ParamValue::Float(2.0))])); - let mut values = Vec::new(); + let values = RefCell::new(Vec::new()); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x_val = x.suggest(trial)?; - values.push(x_val); + values.borrow_mut().push(x_val); Ok::<_, Error>(x_val * x_val) }) .unwrap(); + let values = values.into_inner(); // First two trials should use enqueued values assert_eq!(values[0], 1.0); assert_eq!(values[1], 2.0); @@ -117,7 +119,7 @@ fn test_enqueue_trials_appear_in_completed_trials() { study.enqueue(HashMap::from([(x.id(), ParamValue::Float(7.0))])); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let x_val = x.suggest(trial)?; Ok::<_, Error>(x_val) }) @@ -178,7 +180,7 @@ fn test_enqueue_counted_in_n_trials() { study.enqueue(HashMap::from([(x.id(), ParamValue::Float(2.0))])); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x_val = x.suggest(trial)?; Ok::<_, Error>(x_val) }) diff --git a/tests/study/objective.rs b/tests/study/objective.rs index 6bb1944..8d28bf1 100644 --- a/tests/study/objective.rs +++ b/tests/study/objective.rs @@ -30,7 +30,7 @@ fn test_callback_early_stopping() { let study: Study = Study::new(Direction::Minimize); study - .optimize_with( + .optimize( 100, EarlyStopAfter5 { x_param: FloatParam::new(0.0, 10.0), @@ -69,7 +69,7 @@ fn test_callback_early_stopping_on_first_trial() { let study: Study = Study::new(Direction::Minimize); study - .optimize_with( + .optimize( 100, StopImmediately { x_param: FloatParam::new(0.0, 10.0), @@ -109,7 +109,7 @@ fn test_callback_sampler_early_stopping() { let sampler = RandomSampler::with_seed(42); let study: Study = Study::with_sampler(Direction::Minimize, sampler); study - .optimize_with( + .optimize( 100, StopAfter3 { x_param: FloatParam::new(0.0, 10.0), @@ -121,226 +121,42 @@ fn test_callback_sampler_early_stopping() { } #[test] -fn test_retries_successful_trials_not_retried() { - use std::sync::Arc; - use std::sync::atomic::{AtomicU32, Ordering}; - +fn test_objective_struct_basic() { use optimizer::Objective; - struct SuccessObj { + struct SquareObj { x_param: FloatParam, - call_count: Arc, } - impl Objective for SuccessObj { + impl Objective for SquareObj { type Error = Error; fn evaluate(&self, trial: &mut Trial) -> Result { let x = self.x_param.suggest(trial)?; - self.call_count.fetch_add(1, Ordering::Relaxed); Ok(x * x) } - fn max_retries(&self) -> usize { - 3 - } } let study: Study = Study::new(Direction::Minimize); - let call_count = Arc::new(AtomicU32::new(0)); - let obj = SuccessObj { + let obj = SquareObj { x_param: FloatParam::new(0.0, 10.0), - call_count: Arc::clone(&call_count), }; - study.optimize_with(5, obj).unwrap(); + study.optimize(5, obj).unwrap(); - // All trials succeed on first try — exactly 5 calls - assert_eq!(call_count.load(Ordering::Relaxed), 5); assert_eq!(study.n_trials(), 5); } #[test] -fn test_retries_failed_trials_retried_up_to_max() { - use std::sync::Arc; - use std::sync::atomic::{AtomicU32, Ordering}; - - use optimizer::Objective; - - struct AlwaysFailObj { - x_param: FloatParam, - call_count: Arc, - } - - impl Objective for AlwaysFailObj { - type Error = String; - fn evaluate(&self, trial: &mut Trial) -> Result { - let _ = self.x_param.suggest(trial).map_err(|e| e.to_string())?; - self.call_count.fetch_add(1, Ordering::Relaxed); - Err("always fails".to_string()) - } - fn max_retries(&self) -> usize { - 3 - } - } - - let study: Study = Study::new(Direction::Minimize); - let call_count = Arc::new(AtomicU32::new(0)); - let obj = AlwaysFailObj { - x_param: FloatParam::new(0.0, 10.0), - call_count: Arc::clone(&call_count), - }; - - let result = study.optimize_with(1, obj); - - // 1 initial attempt + 3 retries = 4 total calls - assert_eq!(call_count.load(Ordering::Relaxed), 4); - // No trials completed - assert!(matches!(result, Err(Error::NoCompletedTrials))); -} - -#[test] -fn test_retries_permanently_failed_after_exhaustion() { - use optimizer::Objective; - - struct AlwaysFailObj { - x_param: FloatParam, - } - - impl Objective for AlwaysFailObj { - type Error = String; - fn evaluate(&self, trial: &mut Trial) -> Result { - let _ = self.x_param.suggest(trial).map_err(|e| e.to_string())?; - Err("transient error".to_string()) - } - fn max_retries(&self) -> usize { - 2 - } - } - - let study: Study = Study::new(Direction::Minimize); - let obj = AlwaysFailObj { - x_param: FloatParam::new(0.0, 10.0), - }; - - let result = study.optimize_with(3, obj); - - assert!( - matches!(result, Err(Error::NoCompletedTrials)), - "all trials should permanently fail" - ); - assert_eq!( - study.n_trials(), - 0, - "no completed trials should be recorded" - ); -} - -#[test] -fn test_retries_uses_same_parameters() { - use std::sync::atomic::{AtomicU32, Ordering}; - use std::sync::{Arc, Mutex}; - - use optimizer::Objective; - - struct RetryObj { - x_param: FloatParam, - seen_values: Arc>>, - call_count: Arc, - } - - impl Objective for RetryObj { - type Error = String; - fn evaluate(&self, trial: &mut Trial) -> Result { - let x = self.x_param.suggest(trial).map_err(|e| e.to_string())?; - self.seen_values.lock().unwrap().push(x); - let count = self.call_count.fetch_add(1, Ordering::Relaxed) + 1; - // Fail first two attempts, succeed on third - if count < 3 { - Err("transient".to_string()) - } else { - Ok(x * x) - } - } - fn max_retries(&self) -> usize { - 2 - } - } - - let study: Study = Study::new(Direction::Minimize); - let seen_values = Arc::new(Mutex::new(Vec::new())); - let call_count = Arc::new(AtomicU32::new(0)); - let obj = RetryObj { - x_param: FloatParam::new(0.0, 10.0), - seen_values: Arc::clone(&seen_values), - call_count: Arc::clone(&call_count), - }; - - study.optimize_with(1, obj).unwrap(); - - let values = seen_values.lock().unwrap(); - assert_eq!(values.len(), 3, "should be called 3 times (1 + 2 retries)"); - // All three calls should have gotten the same parameter value - assert_eq!(values[0], values[1]); - assert_eq!(values[1], values[2]); -} - -#[test] -fn test_retries_n_trials_counts_unique_configs() { - use std::sync::Arc; - use std::sync::atomic::{AtomicU32, Ordering}; - - use optimizer::Objective; - - struct FailFirstObj { - x_param: FloatParam, - call_count: Arc, - } - - impl Objective for FailFirstObj { - type Error = String; - fn evaluate(&self, trial: &mut Trial) -> Result { - let x = self.x_param.suggest(trial).map_err(|e| e.to_string())?; - let count = self.call_count.fetch_add(1, Ordering::Relaxed) + 1; - // Fail first attempt of each config, succeed on retry - if count % 2 == 1 { - Err("transient".to_string()) - } else { - Ok(x * x) - } - } - fn max_retries(&self) -> usize { - 2 - } - } - - let study: Study = Study::new(Direction::Minimize); - let call_count = Arc::new(AtomicU32::new(0)); - let obj = FailFirstObj { - x_param: FloatParam::new(0.0, 10.0), - call_count: Arc::clone(&call_count), - }; - - study.optimize_with(3, obj).unwrap(); - - // 3 unique configs, each needing 2 calls = 6 total calls - assert_eq!(call_count.load(Ordering::Relaxed), 6); - // But only 3 completed trials - assert_eq!(study.n_trials(), 3); -} - -#[test] -fn test_retries_with_zero_max_retries_same_as_optimize() { - let study: Study = Study::new(Direction::Minimize); +fn test_closure_and_objective_produce_same_results() { let x_param = FloatParam::new(0.0, 10.0); - let call_count = std::cell::Cell::new(0u32); + let study: Study = Study::new(Direction::Minimize); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut Trial| { let x = x_param.suggest(trial)?; - call_count.set(call_count.get() + 1); Ok::<_, Error>(x * x) }) .unwrap(); - assert_eq!(call_count.get(), 5); assert_eq!(study.n_trials(), 5); } diff --git a/tests/study/summary.rs b/tests/study/summary.rs index 8f688af..f868f0d 100644 --- a/tests/study/summary.rs +++ b/tests/study/summary.rs @@ -8,7 +8,7 @@ fn test_summary_with_completed_trials() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let val = x.suggest(trial)?; Ok::<_, Error>(val * val) }) @@ -61,7 +61,7 @@ fn test_display_matches_summary() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(3, |trial| { + .optimize(3, |trial: &mut optimizer::Trial| { let val = x.suggest(trial)?; Ok::<_, Error>(val) }) diff --git a/tests/study/workflow.rs b/tests/study/workflow.rs index f41bff8..4b998ec 100644 --- a/tests/study/workflow.rs +++ b/tests/study/workflow.rs @@ -8,7 +8,7 @@ fn test_study_basic_workflow() { let x_param = FloatParam::new(-5.0, 5.0); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x * x) }) @@ -25,11 +25,11 @@ fn test_study_with_failures() { let x_param = FloatParam::new(-5.0, 5.0); // Every other trial fails - let mut counter = 0; + let counter = std::cell::Cell::new(0u32); study - .optimize(10, |trial| { - counter += 1; - if counter % 2 == 0 { + .optimize(10, |trial: &mut optimizer::Trial| { + counter.set(counter.get() + 1); + if counter.get().is_multiple_of(2) { return Err::("intentional failure"); } let x = x_param.suggest(trial).map_err(|_| "param error")?; @@ -64,7 +64,7 @@ fn test_study_trials_iteration() { let x_param = FloatParam::new(0.0, 1.0); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x) }) @@ -95,7 +95,7 @@ fn test_study_set_sampler() { let x_param = FloatParam::new(-5.0, 5.0); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x * x) }) @@ -110,7 +110,7 @@ fn test_study_with_i32_value_type() { let x_param = IntParam::new(-10, 10); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x.abs() as i32) }) @@ -125,7 +125,9 @@ fn test_study_with_i32_value_type() { fn test_optimize_all_trials_fail() { let study: Study = Study::new(Direction::Minimize); - let result = study.optimize(5, |_trial| Err::("always fails")); + let result = study.optimize(5, |_trial: &mut optimizer::Trial| { + Err::("always fails") + }); assert!( matches!(result, Err(Error::NoCompletedTrials)), @@ -139,7 +141,7 @@ fn test_best_value() { let x_param = FloatParam::new(0.0, 10.0); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x) }) @@ -160,7 +162,7 @@ fn test_best_trial_with_nan_values() { let x_param = FloatParam::new(0.0, 10.0); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; Ok::<_, Error>(x) }) @@ -199,7 +201,7 @@ fn test_multiple_params_in_optimization() { let n_param = IntParam::new(1, 5); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let n = n_param.suggest(trial)?; Ok::<_, Error>(x * x + n as f64) @@ -216,7 +218,7 @@ fn test_suggest_bool_in_optimization() { let x_param = FloatParam::new(0.0, 10.0); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let use_feature = use_feature_param.suggest(trial)?; let x = x_param.suggest(trial)?; @@ -235,7 +237,7 @@ fn test_completed_trial_get() { let n_param = IntParam::new(1, 10).name("n"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let x = x_param.suggest(trial)?; let n = n_param.suggest(trial)?; Ok::<_, Error>(x * x + n as f64) diff --git a/tests/user_attr_tests.rs b/tests/user_attr_tests.rs index 866560b..bf544d9 100644 --- a/tests/user_attr_tests.rs +++ b/tests/user_attr_tests.rs @@ -7,7 +7,7 @@ fn set_and_get_float_attr() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; trial.set_user_attr("score", 42.5); assert_eq!(trial.user_attr("score"), Some(&AttrValue::Float(42.5))); @@ -22,7 +22,7 @@ fn set_and_get_int_attr() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; trial.set_user_attr("epoch", 42_i64); assert_eq!(trial.user_attr("epoch"), Some(&AttrValue::Int(42))); @@ -37,7 +37,7 @@ fn set_and_get_string_attr() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; trial.set_user_attr("model", "resnet50"); assert_eq!( @@ -55,7 +55,7 @@ fn set_and_get_bool_attr() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; trial.set_user_attr("converged", true); assert_eq!(trial.user_attr("converged"), Some(&AttrValue::Bool(true))); @@ -70,7 +70,7 @@ fn attrs_propagate_to_completed_trial() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; trial.set_user_attr("time_secs", 1.5); trial.set_user_attr("tag", "baseline"); @@ -92,7 +92,7 @@ fn overwrite_attr_replaces_value() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; trial.set_user_attr("key", "old"); trial.set_user_attr("key", "new"); @@ -117,7 +117,7 @@ fn missing_attr_returns_none() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; assert_eq!(trial.user_attr("nonexistent"), None); Ok::<_, optimizer::Error>(1.0) @@ -134,7 +134,7 @@ fn user_attrs_map_returns_all() { let x = FloatParam::new(0.0, 1.0); study - .optimize(1, |trial| { + .optimize(1, |trial: &mut optimizer::Trial| { let _ = x.suggest(trial)?; trial.set_user_attr("a", 1.0); trial.set_user_attr("b", true); diff --git a/tests/visualization_tests.rs b/tests/visualization_tests.rs index 19a3b3a..349ec47 100644 --- a/tests/visualization_tests.rs +++ b/tests/visualization_tests.rs @@ -9,7 +9,7 @@ fn html_report_creates_file() { let y = IntParam::new(1, 5).name("y"); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, optimizer::Error>(xv + yv as f64) @@ -32,7 +32,7 @@ fn html_report_contains_all_chart_sections() { let y = FloatParam::new(-5.0, 5.0).name("y"); study - .optimize(20, |trial| { + .optimize(20, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; let yv = y.suggest(trial)?; Ok::<_, optimizer::Error>(xv * xv + yv * yv) @@ -85,7 +85,7 @@ fn html_report_single_param_no_parcoords() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(xv * xv) }) @@ -109,7 +109,7 @@ fn html_report_maximize_direction() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(xv) }) @@ -130,7 +130,7 @@ fn export_html_convenience_method() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(5, |trial| { + .optimize(5, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; Ok::<_, optimizer::Error>(xv * xv) }) @@ -156,7 +156,7 @@ fn html_report_with_intermediate_values() { let x = FloatParam::new(0.0, 10.0).name("x"); study - .optimize(10, |trial| { + .optimize(10, |trial: &mut optimizer::Trial| { let xv = x.suggest(trial)?; for step in 0..5 { let val = xv * xv + step as f64;