Files
rust-optimizer/src/study.rs
T

1739 lines
58 KiB
Rust
Raw Normal View History

2026-01-30 16:02:42 +01:00
//! Study implementation for managing optimization trials.
2026-02-06 18:54:55 +01:00
use core::any::Any;
2026-02-11 17:27:43 +01:00
use core::fmt;
2026-01-30 16:02:42 +01:00
#[cfg(feature = "async")]
2026-01-30 19:21:35 +01:00
use core::future::Future;
use core::ops::ControlFlow;
use core::sync::atomic::{AtomicU64, Ordering};
use core::time::Duration;
use std::collections::{HashMap, VecDeque};
2026-01-30 16:02:42 +01:00
use std::sync::Arc;
use std::time::Instant;
2026-01-30 16:02:42 +01:00
use parking_lot::{Mutex, RwLock};
2026-01-30 16:02:42 +01:00
use crate::param::ParamValue;
use crate::parameter::ParamId;
use crate::pruner::{NopPruner, Pruner};
2026-01-30 19:21:35 +01:00
use crate::sampler::random::RandomSampler;
use crate::sampler::{CompletedTrial, Sampler};
2026-01-30 16:02:42 +01:00
use crate::trial::Trial;
use crate::types::{Direction, TrialState};
2026-01-30 16:02:42 +01:00
/// A study manages the optimization process, tracking trials and their results.
///
/// The study is parameterized by the objective value type `V`, which defaults to `f64`.
/// The only constraint on `V` is `PartialOrd`, allowing comparison of objective values
/// to determine which trial is best.
///
/// When `V = f64`, the study passes trial history to the sampler for informed
/// parameter suggestions (e.g., TPE sampler uses history to guide sampling).
///
/// # Examples
///
/// ```
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// // Create a study to minimize an objective function
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// assert_eq!(study.direction(), Direction::Minimize);
/// ```
pub struct Study<V = f64>
where
V: PartialOrd,
{
/// The optimization direction.
direction: Direction,
/// The sampler used to generate parameter values.
sampler: Arc<dyn Sampler>,
/// The pruner used to decide whether to stop trials early.
pruner: Arc<dyn Pruner>,
2026-01-30 16:02:42 +01:00
/// Completed trials (wrapped in Arc for sharing with Trial).
completed_trials: Arc<RwLock<Vec<CompletedTrial<V>>>>,
/// Counter for generating unique trial IDs.
next_trial_id: AtomicU64,
2026-02-06 18:54:55 +01:00
/// Optional factory for creating sampler-aware trials.
/// Set automatically for `Study<f64>` so that `create_trial()` and all
/// optimization methods use the sampler without requiring `_with_sampler` suffixes.
trial_factory: Option<Arc<dyn Fn(u64) -> Trial + Send + Sync>>,
/// Queue of parameter configurations to evaluate next.
enqueued_params: Arc<Mutex<VecDeque<HashMap<ParamId, ParamValue>>>>,
2026-01-30 16:02:42 +01:00
}
impl<V> Study<V>
where
V: PartialOrd,
{
/// Creates a new study with the given optimization direction.
///
/// Uses the default `RandomSampler` for parameter sampling.
///
/// # Arguments
///
/// * `direction` - Whether to minimize or maximize the objective function.
///
/// # Examples
///
/// ```
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// assert_eq!(study.direction(), Direction::Minimize);
/// ```
2026-01-30 19:21:35 +01:00
#[must_use]
2026-02-06 18:54:55 +01:00
pub fn new(direction: Direction) -> Self
where
V: 'static,
{
2026-01-30 16:02:42 +01:00
Self::with_sampler(direction, RandomSampler::new())
}
/// Creates a study that minimizes the objective value.
///
/// This is a shorthand for `Study::with_sampler(Direction::Minimize, sampler)`.
///
/// # Arguments
///
/// * `sampler` - The sampler to use for parameter sampling.
///
/// # Examples
///
/// ```
/// use optimizer::Study;
/// use optimizer::sampler::tpe::TpeSampler;
///
/// let study: Study<f64> = Study::minimize(TpeSampler::new());
/// assert_eq!(study.direction(), optimizer::Direction::Minimize);
/// ```
#[must_use]
pub fn minimize(sampler: impl Sampler + 'static) -> Self
where
V: 'static,
{
Self::with_sampler(Direction::Minimize, sampler)
}
/// Creates a study that maximizes the objective value.
///
/// This is a shorthand for `Study::with_sampler(Direction::Maximize, sampler)`.
///
/// # Arguments
///
/// * `sampler` - The sampler to use for parameter sampling.
///
/// # Examples
///
/// ```
/// use optimizer::Study;
/// use optimizer::sampler::tpe::TpeSampler;
///
/// let study: Study<f64> = Study::maximize(TpeSampler::new());
/// assert_eq!(study.direction(), optimizer::Direction::Maximize);
/// ```
#[must_use]
pub fn maximize(sampler: impl Sampler + 'static) -> Self
where
V: 'static,
{
Self::with_sampler(Direction::Maximize, sampler)
}
2026-01-30 16:02:42 +01:00
/// Creates a new study with a custom sampler.
///
/// # Arguments
///
/// * `direction` - Whether to minimize or maximize the objective function.
/// * `sampler` - The sampler to use for parameter sampling.
///
/// # Examples
///
/// ```
2026-01-30 19:21:35 +01:00
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Maximize, sampler);
/// assert_eq!(study.direction(), Direction::Maximize);
/// ```
2026-02-06 18:54:55 +01:00
pub fn with_sampler(direction: Direction, sampler: impl Sampler + 'static) -> Self
where
V: 'static,
{
let sampler: Arc<dyn Sampler> = Arc::new(sampler);
let completed_trials = Arc::new(RwLock::new(Vec::new()));
let pruner: Arc<dyn Pruner> = Arc::new(NopPruner);
2026-02-06 18:54:55 +01:00
// For Study<f64>, set up a trial factory that provides sampler integration.
// This uses Any downcasting to check at runtime whether V = f64.
let trial_factory = Self::make_trial_factory(&sampler, &completed_trials, &pruner);
2026-02-06 18:54:55 +01:00
2026-01-30 16:02:42 +01:00
Self {
direction,
2026-02-06 18:54:55 +01:00
sampler,
pruner,
2026-02-06 18:54:55 +01:00
completed_trials,
2026-01-30 16:02:42 +01:00
next_trial_id: AtomicU64::new(0),
2026-02-06 18:54:55 +01:00
trial_factory,
enqueued_params: Arc::new(Mutex::new(VecDeque::new())),
2026-01-30 16:02:42 +01:00
}
}
2026-02-06 18:54:55 +01:00
/// Builds a trial factory for sampler integration when `V = f64`.
fn make_trial_factory(
sampler: &Arc<dyn Sampler>,
completed_trials: &Arc<RwLock<Vec<CompletedTrial<V>>>>,
pruner: &Arc<dyn Pruner>,
2026-02-06 18:54:55 +01:00
) -> Option<Arc<dyn Fn(u64) -> Trial + Send + Sync>>
where
V: 'static,
{
// Try to downcast the completed_trials Arc to the f64 specialization.
// This succeeds only when V = f64, enabling automatic sampler integration.
let any_ref: &dyn Any = completed_trials;
let f64_trials: Option<&Arc<RwLock<Vec<CompletedTrial<f64>>>>> = any_ref.downcast_ref();
f64_trials.map(|trials| {
let sampler = Arc::clone(sampler);
let trials = Arc::clone(trials);
let pruner = Arc::clone(pruner);
2026-02-06 18:54:55 +01:00
let factory: Arc<dyn Fn(u64) -> Trial + Send + Sync> = Arc::new(move |id| {
Trial::with_sampler(
id,
Arc::clone(&sampler),
Arc::clone(&trials),
Arc::clone(&pruner),
)
2026-02-06 18:54:55 +01:00
});
factory
})
}
2026-01-30 16:02:42 +01:00
/// Returns the optimization direction.
pub fn direction(&self) -> Direction {
self.direction
}
/// Sets a new sampler for the study.
///
/// # Arguments
///
/// * `sampler` - The sampler to use for parameter sampling.
///
/// # Examples
///
/// ```
2026-01-30 19:21:35 +01:00
/// use optimizer::sampler::tpe::TpeSampler;
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let mut study: Study<f64> = Study::new(Direction::Minimize);
/// study.set_sampler(TpeSampler::new());
/// ```
/// Creates a new study with a custom sampler and pruner.
///
/// # Arguments
///
/// * `direction` - Whether to minimize or maximize the objective function.
/// * `sampler` - The sampler to use for parameter sampling.
/// * `pruner` - The pruner to use for trial pruning.
///
/// # Examples
///
/// ```
/// use optimizer::pruner::NopPruner;
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler_and_pruner(Direction::Minimize, sampler, NopPruner);
/// ```
pub fn with_sampler_and_pruner(
direction: Direction,
sampler: impl Sampler + 'static,
pruner: impl Pruner + 'static,
) -> Self
where
V: 'static,
{
let sampler: Arc<dyn Sampler> = Arc::new(sampler);
let pruner: Arc<dyn Pruner> = Arc::new(pruner);
let completed_trials = Arc::new(RwLock::new(Vec::new()));
let trial_factory = Self::make_trial_factory(&sampler, &completed_trials, &pruner);
Self {
direction,
sampler,
pruner,
completed_trials,
next_trial_id: AtomicU64::new(0),
trial_factory,
enqueued_params: Arc::new(Mutex::new(VecDeque::new())),
}
}
2026-02-06 18:54:55 +01:00
pub fn set_sampler(&mut self, sampler: impl Sampler + 'static)
where
V: 'static,
{
2026-01-30 16:02:42 +01:00
self.sampler = Arc::new(sampler);
self.trial_factory =
Self::make_trial_factory(&self.sampler, &self.completed_trials, &self.pruner);
2026-01-30 16:02:42 +01:00
}
/// Sets a new pruner for the study.
///
/// # Arguments
///
/// * `pruner` - The pruner to use for trial pruning.
pub fn set_pruner(&mut self, pruner: impl Pruner + 'static)
where
V: 'static,
{
self.pruner = Arc::new(pruner);
self.trial_factory =
Self::make_trial_factory(&self.sampler, &self.completed_trials, &self.pruner);
}
/// Returns a reference to the study's pruner.
pub fn pruner(&self) -> &dyn Pruner {
&*self.pruner
}
/// Enqueues a specific parameter configuration to be evaluated next.
///
/// The next call to [`ask()`](Self::ask) or the next trial in [`optimize()`](Self::optimize)
/// will use these exact parameters instead of sampling from the sampler.
///
/// Multiple configurations can be enqueued; they are evaluated in FIFO order.
/// If an enqueued configuration is missing a parameter that the objective calls
/// `suggest()` on, that parameter falls back to normal sampling.
///
/// # Arguments
///
/// * `params` - A map from parameter IDs to the values to use.
///
/// # Examples
///
/// ```
/// use std::collections::HashMap;
///
/// use optimizer::parameter::{FloatParam, IntParam, Parameter};
/// use optimizer::{Direction, ParamValue, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let x = FloatParam::new(0.0, 10.0);
/// let y = IntParam::new(1, 100);
///
/// // Evaluate these specific configurations first
/// study.enqueue(HashMap::from([
/// (x.id(), ParamValue::Float(0.001)),
/// (y.id(), ParamValue::Int(3)),
/// ]));
///
/// // Next trial will use x=0.001, y=3
/// let mut trial = study.ask();
/// assert_eq!(x.suggest(&mut trial).unwrap(), 0.001);
/// assert_eq!(y.suggest(&mut trial).unwrap(), 3);
/// ```
pub fn enqueue(&self, params: HashMap<ParamId, ParamValue>) {
self.enqueued_params.lock().push_back(params);
}
/// Returns the number of enqueued parameter configurations.
#[must_use]
pub fn n_enqueued(&self) -> usize {
self.enqueued_params.lock().len()
}
2026-01-30 16:02:42 +01:00
/// Generates the next unique trial ID.
pub(crate) fn next_trial_id(&self) -> u64 {
self.next_trial_id.fetch_add(1, Ordering::SeqCst)
}
/// Creates a new trial with a unique ID.
///
/// The trial starts in the `Running` state and can be used to suggest
/// parameter values. After the objective function is evaluated, call
/// `complete_trial` or `fail_trial` to record the result.
///
2026-02-06 18:54:55 +01:00
/// For `Study<f64>`, this method automatically integrates with the study's
/// sampler and trial history, so there is no need to call a separate
/// `create_trial_with_sampler()` method.
2026-01-30 16:02:42 +01:00
///
/// # Examples
///
/// ```
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let trial = study.create_trial();
/// assert_eq!(trial.id(), 0);
///
/// let trial2 = study.create_trial();
/// assert_eq!(trial2.id(), 1);
/// ```
pub fn create_trial(&self) -> Trial {
let id = self.next_trial_id();
let mut trial = if let Some(factory) = &self.trial_factory {
2026-02-06 18:54:55 +01:00
factory(id)
} else {
Trial::new(id)
};
// If there are enqueued params, inject them into this trial
if let Some(fixed_params) = self.enqueued_params.lock().pop_front() {
trial.set_fixed_params(fixed_params);
2026-02-06 18:54:55 +01:00
}
trial
2026-01-30 16:02:42 +01:00
}
/// Records a completed trial with its objective value.
///
/// This method stores the trial's parameters, distributions, and objective
/// value in the study's history. The stored data is used by samplers to
/// inform future parameter suggestions.
///
/// # Arguments
///
/// * `trial` - The trial that was evaluated.
/// * `value` - The objective value returned by the objective function.
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(0.0, 1.0);
2026-01-30 16:02:42 +01:00
/// let mut trial = study.create_trial();
2026-02-06 17:15:30 +01:00
/// let x = x_param.suggest(&mut trial).unwrap();
2026-01-30 16:02:42 +01:00
/// let objective_value = x * x;
/// study.complete_trial(trial, objective_value);
///
/// assert_eq!(study.n_trials(), 1);
/// ```
pub fn complete_trial(&self, mut trial: Trial, value: V) {
trial.set_complete();
let mut completed = CompletedTrial::with_intermediate_values(
2026-01-30 16:02:42 +01:00
trial.id(),
trial.params().clone(),
trial.distributions().clone(),
2026-02-06 17:15:30 +01:00
trial.param_labels().clone(),
2026-01-30 16:02:42 +01:00
value,
trial.intermediate_values().to_vec(),
trial.user_attrs().clone(),
2026-01-30 16:02:42 +01:00
);
completed.state = TrialState::Complete;
2026-01-30 16:02:42 +01:00
self.completed_trials.write().push(completed);
}
/// Records a failed trial with an error message.
///
/// Failed trials are not stored in the study's history and do not
/// contribute to future sampling decisions. This method is useful
/// when the objective function raises an error that should not stop
/// the optimization process.
///
/// # Arguments
///
/// * `trial` - The trial that failed.
/// * `_error` - An error message describing why the trial failed.
///
/// # Examples
///
/// ```
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let trial = study.create_trial();
/// study.fail_trial(trial, "objective function raised an exception");
///
/// // Failed trials are not counted
/// assert_eq!(study.n_trials(), 0);
/// ```
pub fn fail_trial(&self, mut trial: Trial, _error: impl ToString) {
trial.set_failed();
// Failed trials are not stored in completed_trials
// They could be stored in a separate list for debugging if needed
}
/// Request a new trial with suggested parameters.
///
/// This is the first half of the ask-and-tell interface. After calling
/// `ask()`, use parameter types to suggest values on the returned trial,
/// evaluate your objective externally, then pass the trial back to
/// [`tell()`](Self::tell) with the result.
///
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let x = FloatParam::new(0.0, 10.0);
///
/// let mut trial = study.ask();
/// let x_val = x.suggest(&mut trial).unwrap();
/// let value = x_val * x_val;
/// study.tell(trial, Ok::<_, &str>(value));
/// ```
pub fn ask(&self) -> Trial {
self.create_trial()
}
/// Report the result of a trial obtained from [`ask()`](Self::ask).
///
/// Pass `Ok(value)` for a successful evaluation or `Err(reason)` for a
/// failure. Failed trials are not stored in the study's history.
///
/// # Examples
///
/// ```
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
///
/// let trial = study.ask();
/// study.tell(trial, Ok::<_, &str>(42.0));
/// assert_eq!(study.n_trials(), 1);
///
/// let trial = study.ask();
/// study.tell(trial, Err::<f64, _>("evaluation failed"));
/// assert_eq!(study.n_trials(), 1); // failed trials not counted
/// ```
pub fn tell(&self, trial: Trial, value: core::result::Result<V, impl ToString>) {
match value {
Ok(v) => self.complete_trial(trial, v),
Err(e) => self.fail_trial(trial, e),
}
}
/// Records a pruned trial, preserving its intermediate values.
///
/// Pruned trials are stored alongside completed trials so that samplers
/// can optionally learn from partial evaluations. The trial's state is
/// set to `Pruned`.
///
/// # Arguments
///
/// * `trial` - The trial that was pruned.
pub fn prune_trial(&self, mut trial: Trial)
where
V: Default,
{
trial.set_pruned();
let mut completed = CompletedTrial::with_intermediate_values(
trial.id(),
trial.params().clone(),
trial.distributions().clone(),
trial.param_labels().clone(),
V::default(),
trial.intermediate_values().to_vec(),
trial.user_attrs().clone(),
);
completed.state = TrialState::Pruned;
self.completed_trials.write().push(completed);
}
2026-01-30 16:02:42 +01:00
/// Returns an iterator over all completed trials.
///
/// The iterator yields references to `CompletedTrial` values, which contain
/// the trial's parameters, distributions, and objective value.
///
/// Note: This method acquires a read lock on the completed trials, so the
/// returned vector is a clone of the internal storage.
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(0.0, 1.0);
2026-01-30 16:02:42 +01:00
/// let mut trial = study.create_trial();
2026-02-06 17:15:30 +01:00
/// let _ = x_param.suggest(&mut trial);
2026-01-30 16:02:42 +01:00
/// study.complete_trial(trial, 0.5);
///
/// for completed in study.trials() {
/// println!("Trial {} has value {:?}", completed.id, completed.value);
/// }
/// ```
pub fn trials(&self) -> Vec<CompletedTrial<V>>
where
V: Clone,
{
self.completed_trials.read().clone()
}
/// Returns the number of completed trials.
///
/// Failed trials are not counted.
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// assert_eq!(study.n_trials(), 0);
///
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(0.0, 1.0);
2026-01-30 16:02:42 +01:00
/// let mut trial = study.create_trial();
2026-02-06 17:15:30 +01:00
/// let _ = x_param.suggest(&mut trial);
2026-01-30 16:02:42 +01:00
/// study.complete_trial(trial, 0.5);
/// assert_eq!(study.n_trials(), 1);
/// ```
pub fn n_trials(&self) -> usize {
self.completed_trials.read().len()
}
/// Returns the number of pruned trials.
pub fn n_pruned_trials(&self) -> usize {
self.completed_trials
.read()
.iter()
.filter(|t| t.state == TrialState::Pruned)
.count()
}
2026-01-30 16:02:42 +01:00
/// Returns the trial with the best objective value.
///
/// The "best" trial depends on the optimization direction:
/// - `Direction::Minimize`: Returns the trial with the lowest objective value.
/// - `Direction::Maximize`: Returns the trial with the highest objective value.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if no trials have been completed.
2026-01-30 16:02:42 +01:00
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
///
/// // Error when no trials completed
/// assert!(study.best_trial().is_err());
///
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(0.0, 1.0);
///
2026-01-30 16:02:42 +01:00
/// let mut trial1 = study.create_trial();
2026-02-06 17:15:30 +01:00
/// let _ = x_param.suggest(&mut trial1);
2026-01-30 16:02:42 +01:00
/// study.complete_trial(trial1, 0.8);
///
/// let mut trial2 = study.create_trial();
2026-02-06 17:15:30 +01:00
/// let _ = x_param.suggest(&mut trial2);
2026-01-30 16:02:42 +01:00
/// study.complete_trial(trial2, 0.3);
///
/// let best = study.best_trial().unwrap();
/// assert_eq!(best.value, 0.3); // Minimize: lower is better
/// ```
pub fn best_trial(&self) -> crate::Result<CompletedTrial<V>>
where
V: Clone,
{
let trials = self.completed_trials.read();
let best = trials
.iter()
.filter(|t| t.state == TrialState::Complete)
2026-01-30 16:02:42 +01:00
.max_by(|a, b| {
// For Minimize, we want the smallest value to be "max" in ordering
// For Maximize, we want the largest value to be "max" in ordering
let ordering = a.value.partial_cmp(&b.value);
match self.direction {
Direction::Minimize => {
// Reverse ordering: smaller values are "greater" for max_by
2026-01-30 19:21:35 +01:00
ordering.map_or(core::cmp::Ordering::Equal, core::cmp::Ordering::reverse)
2026-01-30 16:02:42 +01:00
}
Direction::Maximize => {
// Normal ordering: larger values are "greater" for max_by
2026-01-30 19:21:35 +01:00
ordering.unwrap_or(core::cmp::Ordering::Equal)
2026-01-30 16:02:42 +01:00
}
}
})
.ok_or(crate::Error::NoCompletedTrials)?;
2026-01-30 16:02:42 +01:00
Ok(best.clone())
}
/// Returns the best objective value found so far.
///
/// The "best" value depends on the optimization direction:
/// - `Direction::Minimize`: Returns the lowest objective value.
/// - `Direction::Maximize`: Returns the highest objective value.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if no trials have been completed.
2026-01-30 16:02:42 +01:00
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// let study: Study<f64> = Study::new(Direction::Maximize);
///
/// // Error when no trials completed
/// assert!(study.best_value().is_err());
///
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(0.0, 1.0);
///
2026-01-30 16:02:42 +01:00
/// let mut trial1 = study.create_trial();
2026-02-06 17:15:30 +01:00
/// let _ = x_param.suggest(&mut trial1);
2026-01-30 16:02:42 +01:00
/// study.complete_trial(trial1, 0.3);
///
/// let mut trial2 = study.create_trial();
2026-02-06 17:15:30 +01:00
/// let _ = x_param.suggest(&mut trial2);
2026-01-30 16:02:42 +01:00
/// study.complete_trial(trial2, 0.8);
///
/// let best = study.best_value().unwrap();
/// assert_eq!(best, 0.8); // Maximize: higher is better
/// ```
pub fn best_value(&self) -> crate::Result<V>
where
V: Clone,
{
self.best_trial().map(|trial| trial.value)
}
/// Returns the top `n` trials sorted by objective value.
///
/// For `Direction::Minimize`, returns trials with the lowest values.
/// For `Direction::Maximize`, returns trials with the highest values.
/// Only includes completed trials (not failed or pruned).
///
/// If fewer than `n` completed trials exist, returns all of them.
pub fn top_trials(&self, n: usize) -> Vec<CompletedTrial<V>>
where
V: Clone,
{
let trials = self.completed_trials.read();
let mut completed: Vec<_> = trials
.iter()
.filter(|t| t.state == TrialState::Complete)
.cloned()
.collect();
completed.sort_by(|a, b| match self.direction {
Direction::Minimize => a
.value
.partial_cmp(&b.value)
.unwrap_or(core::cmp::Ordering::Equal),
Direction::Maximize => b
.value
.partial_cmp(&a.value)
.unwrap_or(core::cmp::Ordering::Equal),
});
completed.truncate(n);
completed
}
2026-01-30 16:02:42 +01:00
/// Runs optimization with the given objective function.
///
/// This method runs `n_trials` evaluations sequentially. For each trial:
/// 1. A new trial is created
/// 2. The objective function is called with the trial
/// 3. If successful, the trial is recorded as completed
/// 4. If the objective returns an error, the trial is recorded as failed
///
/// Failed trials do not stop the optimization; the process continues with
/// the next trial.
///
/// # Arguments
///
/// * `n_trials` - The number of trials to run.
/// * `objective` - A closure that takes a mutable reference to a `Trial` and
/// returns the objective value or an error.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
2026-01-30 16:02:42 +01:00
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
2026-01-30 19:21:35 +01:00
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// // Minimize x^2
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(-10.0, 10.0);
///
2026-01-30 16:02:42 +01:00
/// study
/// .optimize(10, |trial| {
2026-02-06 17:15:30 +01:00
/// let x = x_param.suggest(trial)?;
/// Ok::<_, optimizer::Error>(x * x)
2026-01-30 16:02:42 +01:00
/// })
/// .unwrap();
///
/// // At least one trial should have completed
/// assert!(study.n_trials() > 0);
/// let best = study.best_value().unwrap();
/// assert!(best >= 0.0);
/// ```
pub fn optimize<F, E>(&self, n_trials: usize, mut objective: F) -> crate::Result<()>
where
2026-01-30 19:21:35 +01:00
F: FnMut(&mut Trial) -> core::result::Result<V, E>,
E: ToString + 'static,
V: Default,
2026-01-30 16:02:42 +01:00
{
for _ in 0..n_trials {
let mut trial = self.create_trial();
match objective(&mut trial) {
Ok(value) => {
self.complete_trial(trial, value);
}
Err(e) => {
if is_trial_pruned(&e) {
self.prune_trial(trial);
} else {
self.fail_trial(trial, e.to_string());
}
2026-01-30 16:02:42 +01:00
}
}
}
// Return error if no trials completed successfully
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
2026-01-30 16:02:42 +01:00
}
Ok(())
}
/// Runs optimization asynchronously with the given objective function.
///
/// This method runs `n_trials` evaluations sequentially, but the objective
/// function can be async (e.g., for I/O-bound operations like network requests
/// or file operations).
///
/// The objective function takes ownership of the `Trial` and must return it
/// along with the result. This allows async operations to use the trial
/// across await points.
///
/// # Arguments
///
/// * `n_trials` - The number of trials to run.
/// * `objective` - A function that takes a `Trial` and returns a `Future`
/// that resolves to a tuple of `(Trial, Result<V, E>)`.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
2026-01-30 16:02:42 +01:00
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
2026-01-30 19:21:35 +01:00
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// # #[cfg(feature = "async")]
/// # async fn example() -> optimizer::Result<()> {
2026-01-30 16:02:42 +01:00
/// // Minimize x^2 with async objective
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(-10.0, 10.0);
///
2026-01-30 16:02:42 +01:00
/// study
2026-02-06 17:15:30 +01:00
/// .optimize_async(10, |mut trial| {
/// let x_param = x_param.clone();
/// async move {
/// let x = x_param.suggest(&mut trial)?;
/// // Simulate async work (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// }
2026-01-30 16:02:42 +01:00
/// })
/// .await?;
///
/// // At least one trial should have completed
/// assert!(study.n_trials() > 0);
/// # Ok(())
/// # }
/// ```
#[cfg(feature = "async")]
pub async fn optimize_async<F, Fut, E>(
&self,
n_trials: usize,
objective: F,
) -> crate::Result<()>
where
F: Fn(Trial) -> Fut,
2026-01-30 19:21:35 +01:00
Fut: Future<Output = core::result::Result<(Trial, V), E>>,
2026-01-30 16:02:42 +01:00
E: ToString,
{
for _ in 0..n_trials {
let trial = self.create_trial();
match objective(trial).await {
Ok((trial, value)) => {
self.complete_trial(trial, value);
}
Err(e) => {
// For async, we don't have the trial back on error
// We'll just count this as a failed trial without recording it
let _ = e.to_string();
}
}
}
// Return error if no trials completed successfully
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
2026-01-30 16:02:42 +01:00
}
Ok(())
}
/// Runs optimization with bounded parallelism for concurrent trial evaluation.
///
/// This method runs up to `concurrency` trials simultaneously, allowing
/// efficient use of async I/O-bound objective functions. A semaphore limits
/// the number of concurrent evaluations.
///
/// The objective function takes ownership of the `Trial` and must return it
/// along with the result. This allows async operations to use the trial
/// across await points.
///
/// # Arguments
///
/// * `n_trials` - The total number of trials to run.
/// * `concurrency` - The maximum number of trials to run simultaneously.
/// * `objective` - A function that takes a `Trial` and returns a `Future`
/// that resolves to a tuple of `(Trial, V)` or an error.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
/// Returns `Error::TaskError` if the semaphore is closed or a spawned task panics.
2026-01-30 16:02:42 +01:00
///
/// # Examples
///
/// ```
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
2026-01-30 19:21:35 +01:00
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// # #[cfg(feature = "async")]
/// # async fn example() -> optimizer::Result<()> {
2026-01-30 16:02:42 +01:00
/// // Minimize x^2 with parallel async evaluation
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(-10.0, 10.0);
///
2026-01-30 16:02:42 +01:00
/// study
2026-02-06 17:15:30 +01:00
/// .optimize_parallel(10, 4, move |mut trial| {
/// let x_param = x_param.clone();
/// async move {
/// let x = x_param.suggest(&mut trial)?;
/// // Async objective function (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// }
2026-01-30 16:02:42 +01:00
/// })
/// .await?;
///
/// // All trials should have completed
/// assert_eq!(study.n_trials(), 10);
/// # Ok(())
/// # }
/// ```
#[cfg(feature = "async")]
pub async fn optimize_parallel<F, Fut, E>(
&self,
n_trials: usize,
concurrency: usize,
objective: F,
) -> crate::Result<()>
where
F: Fn(Trial) -> Fut + Send + Sync + 'static,
2026-01-30 19:21:35 +01:00
Fut: Future<Output = core::result::Result<(Trial, V), E>> + Send,
2026-01-30 16:02:42 +01:00
E: ToString + Send + 'static,
V: Send + 'static,
{
use tokio::sync::Semaphore;
let semaphore = Arc::new(Semaphore::new(concurrency));
let objective = Arc::new(objective);
let mut handles = Vec::with_capacity(n_trials);
for _ in 0..n_trials {
2026-01-30 19:21:35 +01:00
let permit = semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| crate::Error::TaskError(e.to_string()))?;
2026-01-30 16:02:42 +01:00
let trial = self.create_trial();
let objective = Arc::clone(&objective);
let handle = tokio::spawn(async move {
let result = objective(trial).await;
drop(permit); // Release semaphore permit when done
result
});
handles.push(handle);
}
// Wait for all tasks and record results
for handle in handles {
2026-01-30 19:21:35 +01:00
match handle
.await
.map_err(|e| crate::Error::TaskError(e.to_string()))?
2026-01-30 19:21:35 +01:00
{
2026-01-30 16:02:42 +01:00
Ok((trial, value)) => {
self.complete_trial(trial, value);
}
Err(e) => {
let _ = e.to_string();
}
}
}
// Return error if no trials completed successfully
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
2026-01-30 16:02:42 +01:00
}
Ok(())
}
/// Runs optimization with a callback for monitoring progress.
///
/// This method is similar to `optimize`, but calls a callback function after
/// each completed trial. The callback can inspect the study state and the
/// completed trial, and can optionally stop optimization early by returning
/// `ControlFlow::Break(())`.
///
/// # Arguments
///
/// * `n_trials` - The maximum number of trials to run.
/// * `objective` - A closure that takes a mutable reference to a `Trial` and
/// returns the objective value or an error.
/// * `callback` - A closure called after each successful trial. Returns
/// `ControlFlow::Continue(())` to proceed or `ControlFlow::Break(())` to stop.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if no trials completed successfully
2026-01-30 16:02:42 +01:00
/// before optimization stopped (either by completing all trials or early stopping).
/// Returns `Error::Internal` if a completed trial is not found after adding (internal invariant violation).
2026-01-30 16:02:42 +01:00
///
/// # Examples
///
/// ```
/// use std::ops::ControlFlow;
///
2026-02-06 17:15:30 +01:00
/// use optimizer::parameter::{FloatParam, Parameter};
2026-01-30 19:21:35 +01:00
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
2026-01-30 16:02:42 +01:00
///
/// // Stop early when we find a good enough value
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
2026-02-06 17:15:30 +01:00
/// let x_param = FloatParam::new(-10.0, 10.0);
///
2026-01-30 16:02:42 +01:00
/// study
/// .optimize_with_callback(
/// 100,
/// |trial| {
2026-02-06 17:15:30 +01:00
/// let x = x_param.suggest(trial)?;
/// Ok::<_, optimizer::Error>(x * x)
2026-01-30 16:02:42 +01:00
/// },
/// |_study, completed_trial| {
/// // Stop early if we find a value less than 1.0
/// if completed_trial.value < 1.0 {
/// ControlFlow::Break(())
/// } else {
/// ControlFlow::Continue(())
/// }
/// },
/// )
/// .unwrap();
///
/// // May have stopped early, but should have at least one trial
/// assert!(study.n_trials() > 0);
/// ```
pub fn optimize_with_callback<F, C, E>(
&self,
n_trials: usize,
mut objective: F,
mut callback: C,
) -> crate::Result<()>
where
V: Clone + Default,
2026-01-30 19:21:35 +01:00
F: FnMut(&mut Trial) -> core::result::Result<V, E>,
2026-01-30 16:02:42 +01:00
C: FnMut(&Study<V>, &CompletedTrial<V>) -> ControlFlow<()>,
E: ToString + 'static,
2026-01-30 16:02:42 +01:00
{
for _ in 0..n_trials {
let mut trial = self.create_trial();
match objective(&mut trial) {
Ok(value) => {
self.complete_trial(trial, value);
// Get the just-completed trial for the callback
let trials = self.completed_trials.read();
2026-01-30 19:21:35 +01:00
let Some(completed) = trials.last() else {
return Err(crate::Error::Internal(
2026-01-30 19:21:35 +01:00
"completed trial not found after adding",
));
};
2026-01-30 16:02:42 +01:00
// Call the callback and check if we should stop
// Note: We need to drop the read lock before calling callback
// to avoid potential deadlock if callback accesses the study
let completed_clone = completed.clone();
drop(trials);
if let ControlFlow::Break(()) = callback(self, &completed_clone) {
break;
}
}
Err(e) => {
if is_trial_pruned(&e) {
self.prune_trial(trial);
} else {
self.fail_trial(trial, e.to_string());
}
2026-01-30 16:02:42 +01:00
}
}
}
// Return error if no trials completed successfully
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
2026-01-30 16:02:42 +01:00
}
Ok(())
}
/// Runs optimization until the given duration has elapsed.
///
/// Trials that are already running when the timeout is reached will
/// complete — we never interrupt mid-trial. The actual elapsed time
/// may therefore slightly exceed the specified duration.
///
/// # Arguments
///
/// * `duration` - The maximum wall-clock time to spend on optimization.
/// * `objective` - A closure that takes a mutable reference to a `Trial` and
/// returns the objective value or an error.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if no trials completed successfully
/// before the timeout.
///
/// # Examples
///
/// ```
/// use std::time::Duration;
///
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
/// let x_param = FloatParam::new(-10.0, 10.0);
///
/// study
/// .optimize_until(Duration::from_millis(100), |trial| {
/// let x = x_param.suggest(trial)?;
/// Ok::<_, optimizer::Error>(x * x)
/// })
/// .unwrap();
///
/// assert!(study.n_trials() > 0);
/// ```
pub fn optimize_until<F, E>(&self, duration: Duration, mut objective: F) -> crate::Result<()>
where
F: FnMut(&mut Trial) -> core::result::Result<V, E>,
E: ToString + 'static,
V: Default,
{
let deadline = Instant::now() + duration;
while Instant::now() < deadline {
let mut trial = self.create_trial();
match objective(&mut trial) {
Ok(value) => {
self.complete_trial(trial, value);
}
Err(e) => {
if is_trial_pruned(&e) {
self.prune_trial(trial);
} else {
self.fail_trial(trial, e.to_string());
}
}
}
}
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
}
Ok(())
}
/// Runs optimization until the given duration has elapsed, with a callback.
///
/// Like [`optimize_until`](Self::optimize_until), but calls a callback after
/// each completed trial. The callback can stop optimization early by returning
/// `ControlFlow::Break(())`.
///
/// # Arguments
///
/// * `duration` - The maximum wall-clock time to spend on optimization.
/// * `objective` - A closure that takes a mutable reference to a `Trial` and
/// returns the objective value or an error.
/// * `callback` - A closure called after each successful trial. Returns
/// `ControlFlow::Continue(())` to proceed or `ControlFlow::Break(())` to stop.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if no trials completed successfully.
/// Returns `Error::Internal` if a completed trial is not found after adding.
///
/// # Examples
///
/// ```
/// use std::ops::ControlFlow;
/// use std::time::Duration;
///
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
/// let x_param = FloatParam::new(-10.0, 10.0);
///
/// study
/// .optimize_until_with_callback(
/// Duration::from_secs(1),
/// |trial| {
/// let x = x_param.suggest(trial)?;
/// Ok::<_, optimizer::Error>(x * x)
/// },
/// |_study, completed_trial| {
/// if completed_trial.value < 1.0 {
/// ControlFlow::Break(())
/// } else {
/// ControlFlow::Continue(())
/// }
/// },
/// )
/// .unwrap();
///
/// assert!(study.n_trials() > 0);
/// ```
pub fn optimize_until_with_callback<F, C, E>(
&self,
duration: Duration,
mut objective: F,
mut callback: C,
) -> crate::Result<()>
where
V: Clone + Default,
F: FnMut(&mut Trial) -> core::result::Result<V, E>,
C: FnMut(&Study<V>, &CompletedTrial<V>) -> ControlFlow<()>,
E: ToString + 'static,
{
let deadline = Instant::now() + duration;
while Instant::now() < deadline {
let mut trial = self.create_trial();
match objective(&mut trial) {
Ok(value) => {
self.complete_trial(trial, value);
let trials = self.completed_trials.read();
let Some(completed) = trials.last() else {
return Err(crate::Error::Internal(
"completed trial not found after adding",
));
};
let completed_clone = completed.clone();
drop(trials);
if let ControlFlow::Break(()) = callback(self, &completed_clone) {
break;
}
}
Err(e) => {
if is_trial_pruned(&e) {
self.prune_trial(trial);
} else {
self.fail_trial(trial, e.to_string());
}
}
}
}
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
}
Ok(())
}
/// Runs optimization asynchronously until the given duration has elapsed.
///
/// The async variant of [`optimize_until`](Self::optimize_until). Trials are
/// run sequentially, but the objective function can be async.
///
/// # Arguments
///
/// * `duration` - The maximum wall-clock time to spend on optimization.
/// * `objective` - A function that takes a `Trial` and returns a `Future`
/// that resolves to a tuple of `(Trial, Result<V, E>)`.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if no trials completed successfully.
#[cfg(feature = "async")]
pub async fn optimize_until_async<F, Fut, E>(
&self,
duration: Duration,
objective: F,
) -> crate::Result<()>
where
F: Fn(Trial) -> Fut,
Fut: Future<Output = core::result::Result<(Trial, V), E>>,
E: ToString,
{
let deadline = Instant::now() + duration;
while Instant::now() < deadline {
let trial = self.create_trial();
match objective(trial).await {
Ok((trial, value)) => {
self.complete_trial(trial, value);
}
Err(e) => {
let _ = e.to_string();
}
}
}
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
}
Ok(())
}
/// Runs optimization with bounded parallelism until the given duration has elapsed.
///
/// The parallel variant of [`optimize_until`](Self::optimize_until). Runs up to
/// `concurrency` trials simultaneously using async tasks. New trials are spawned
/// as long as the deadline has not been reached; trials already running when the
/// deadline passes will complete.
///
/// # Arguments
///
/// * `duration` - The maximum wall-clock time to spend spawning new trials.
/// * `concurrency` - The maximum number of trials to run simultaneously.
/// * `objective` - A function that takes a `Trial` and returns a `Future`
/// that resolves to a tuple of `(Trial, V)` or an error.
///
/// # 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")]
pub async fn optimize_until_parallel<F, Fut, E>(
&self,
duration: Duration,
concurrency: usize,
objective: F,
) -> crate::Result<()>
where
F: Fn(Trial) -> Fut + Send + Sync + 'static,
Fut: Future<Output = core::result::Result<(Trial, V), E>> + Send,
E: ToString + Send + 'static,
V: Send + 'static,
{
use tokio::sync::Semaphore;
let deadline = Instant::now() + duration;
let semaphore = Arc::new(Semaphore::new(concurrency));
let objective = Arc::new(objective);
let mut handles = Vec::new();
while Instant::now() < deadline {
let permit = semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| crate::Error::TaskError(e.to_string()))?;
let trial = self.create_trial();
let objective = Arc::clone(&objective);
let handle = tokio::spawn(async move {
let result = objective(trial).await;
drop(permit);
result
});
handles.push(handle);
}
for handle in handles {
match handle
.await
.map_err(|e| crate::Error::TaskError(e.to_string()))?
{
Ok((trial, value)) => {
self.complete_trial(trial, value);
}
Err(e) => {
let _ = e.to_string();
}
}
}
let has_complete = self
.completed_trials
.read()
.iter()
.any(|t| t.state == TrialState::Complete);
if !has_complete {
return Err(crate::Error::NoCompletedTrials);
}
2026-01-30 16:02:42 +01:00
Ok(())
}
}
2026-02-11 17:27:43 +01:00
impl<V> Study<V>
where
V: PartialOrd + Clone + fmt::Display,
{
/// Returns a human-readable summary of the study.
///
/// The summary includes:
/// - Optimization direction and total trial count
/// - Breakdown by state (complete, pruned) when applicable
/// - Best trial value and parameters (if any completed trials exist)
///
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let x = FloatParam::new(0.0, 10.0).name("x");
///
/// let mut trial = study.create_trial();
/// let _ = x.suggest(&mut trial).unwrap();
/// study.complete_trial(trial, 0.42);
///
/// let summary = study.summary();
/// assert!(summary.contains("Minimize"));
/// assert!(summary.contains("0.42"));
/// ```
#[must_use]
pub fn summary(&self) -> String {
use fmt::Write;
let trials = self.completed_trials.read();
let n_complete = trials
.iter()
.filter(|t| t.state == TrialState::Complete)
.count();
let n_pruned = trials
.iter()
.filter(|t| t.state == TrialState::Pruned)
.count();
let direction_str = match self.direction {
Direction::Minimize => "Minimize",
Direction::Maximize => "Maximize",
};
let mut s = format!("Study: {direction_str} | {n} trials", n = trials.len());
if n_pruned > 0 {
let _ = write!(s, " ({n_complete} complete, {n_pruned} pruned)");
}
drop(trials);
if let Ok(best) = self.best_trial() {
let _ = write!(s, "\nBest value: {} (trial #{})", best.value, best.id);
if !best.params.is_empty() {
s.push_str("\nBest parameters:");
let mut params: Vec<_> = best.params.iter().collect();
params.sort_by_key(|(id, _)| *id);
for (id, value) in params {
let label = best.param_labels.get(id).map_or("?", String::as_str);
let _ = write!(s, "\n {label} = {value}");
}
}
}
s
}
}
impl<V> fmt::Display for Study<V>
where
V: PartialOrd + Clone + fmt::Display,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.summary())
}
}
2026-02-06 18:54:55 +01:00
// Specialized implementation for Study<f64> that provides deprecated `_with_sampler` aliases.
//
// For Study<f64>, the generic methods from `impl<V> Study<V>` (like `optimize()`,
// `create_trial()`) now automatically use the sampler via the `trial_factory`.
// The `_with_sampler` method names are deprecated in favor of the generic names.
#[allow(clippy::missing_errors_doc)]
2026-01-30 16:02:42 +01:00
impl Study<f64> {
2026-02-06 18:54:55 +01:00
/// Deprecated: use `create_trial()` instead.
///
/// The generic `create_trial()` now automatically integrates with the sampler
/// for `Study<f64>`.
#[deprecated(
since = "0.2.0",
note = "use `create_trial()` instead — it now uses the sampler automatically for Study<f64>"
)]
2026-01-30 16:02:42 +01:00
pub fn create_trial_with_sampler(&self) -> Trial {
2026-02-06 18:54:55 +01:00
self.create_trial()
2026-01-30 16:02:42 +01:00
}
2026-02-06 18:54:55 +01:00
/// Deprecated: use `optimize()` instead.
2026-01-30 16:02:42 +01:00
///
2026-02-06 18:54:55 +01:00
/// The generic `optimize()` now automatically integrates with the sampler
/// for `Study<f64>`.
#[deprecated(
since = "0.2.0",
note = "use `optimize()` instead — it now uses the sampler automatically for Study<f64>"
)]
pub fn optimize_with_sampler<F, E>(&self, n_trials: usize, objective: F) -> crate::Result<()>
2026-01-30 16:02:42 +01:00
where
2026-01-30 19:21:35 +01:00
F: FnMut(&mut Trial) -> core::result::Result<f64, E>,
E: ToString + 'static,
2026-01-30 16:02:42 +01:00
{
2026-02-06 18:54:55 +01:00
self.optimize(n_trials, objective)
2026-01-30 16:02:42 +01:00
}
2026-02-06 18:54:55 +01:00
/// Deprecated: use `optimize_with_callback()` instead.
2026-01-30 16:02:42 +01:00
///
2026-02-06 18:54:55 +01:00
/// The generic `optimize_with_callback()` now automatically integrates with the
/// sampler for `Study<f64>`.
#[deprecated(
since = "0.2.0",
note = "use `optimize_with_callback()` instead — it now uses the sampler automatically for Study<f64>"
)]
2026-01-30 16:02:42 +01:00
pub fn optimize_with_callback_sampler<F, C, E>(
&self,
n_trials: usize,
2026-02-06 18:54:55 +01:00
objective: F,
callback: C,
2026-01-30 16:02:42 +01:00
) -> crate::Result<()>
where
2026-01-30 19:21:35 +01:00
F: FnMut(&mut Trial) -> core::result::Result<f64, E>,
2026-01-30 16:02:42 +01:00
C: FnMut(&Study<f64>, &CompletedTrial<f64>) -> ControlFlow<()>,
E: ToString + 'static,
2026-01-30 16:02:42 +01:00
{
2026-02-06 18:54:55 +01:00
self.optimize_with_callback(n_trials, objective, callback)
2026-01-30 16:02:42 +01:00
}
2026-02-06 18:54:55 +01:00
/// Deprecated: use `optimize_async()` instead.
2026-01-30 16:02:42 +01:00
///
2026-02-06 18:54:55 +01:00
/// The generic `optimize_async()` now automatically integrates with the sampler
/// for `Study<f64>`.
2026-01-30 16:02:42 +01:00
#[cfg(feature = "async")]
2026-02-06 18:54:55 +01:00
#[deprecated(
since = "0.2.0",
note = "use `optimize_async()` instead — it now uses the sampler automatically for Study<f64>"
)]
2026-01-30 16:02:42 +01:00
pub async fn optimize_async_with_sampler<F, Fut, E>(
&self,
n_trials: usize,
objective: F,
) -> crate::Result<()>
where
F: Fn(Trial) -> Fut,
2026-01-30 19:21:35 +01:00
Fut: Future<Output = core::result::Result<(Trial, f64), E>>,
2026-01-30 16:02:42 +01:00
E: ToString,
{
2026-02-06 18:54:55 +01:00
self.optimize_async(n_trials, objective).await
2026-01-30 16:02:42 +01:00
}
2026-02-06 18:54:55 +01:00
/// Deprecated: use `optimize_parallel()` instead.
2026-01-30 16:02:42 +01:00
///
2026-02-06 18:54:55 +01:00
/// The generic `optimize_parallel()` now automatically integrates with the
/// sampler for `Study<f64>`.
2026-01-30 16:02:42 +01:00
#[cfg(feature = "async")]
2026-02-06 18:54:55 +01:00
#[deprecated(
since = "0.2.0",
note = "use `optimize_parallel()` instead — it now uses the sampler automatically for Study<f64>"
)]
2026-01-30 16:02:42 +01:00
pub async fn optimize_parallel_with_sampler<F, Fut, E>(
&self,
n_trials: usize,
concurrency: usize,
objective: F,
) -> crate::Result<()>
where
F: Fn(Trial) -> Fut + Send + Sync + 'static,
2026-01-30 19:21:35 +01:00
Fut: Future<Output = core::result::Result<(Trial, f64), E>> + Send,
2026-01-30 16:02:42 +01:00
E: ToString + Send + 'static,
{
2026-02-06 18:54:55 +01:00
self.optimize_parallel(n_trials, concurrency, objective)
.await
2026-01-30 16:02:42 +01:00
}
}
/// A serializable snapshot of a study's state.
///
/// Since [`Study`] contains non-serializable fields (samplers, atomics, etc.),
/// this struct captures the essential state needed to save and restore a study.
///
/// # Schema versioning
///
/// The `version` field enables future schema evolution without breaking existing files.
/// The current version is `1`.
///
/// # Sampler state
///
/// Sampler state is **not** included in the snapshot. After loading, the study
/// uses a default `RandomSampler`. Call [`Study::set_sampler`] to restore
/// the desired sampler configuration.
#[cfg(feature = "serde")]
#[derive(serde::Serialize, serde::Deserialize)]
pub struct StudySnapshot<V> {
/// Schema version for forward compatibility.
pub version: u32,
/// The optimization direction.
pub direction: Direction,
/// All completed (and pruned) trials.
pub trials: Vec<CompletedTrial<V>>,
/// The next trial ID to assign.
pub next_trial_id: u64,
/// Optional metadata (creation timestamp, sampler description, etc.).
pub metadata: HashMap<String, String>,
}
#[cfg(feature = "serde")]
impl<V: PartialOrd + Clone + serde::Serialize> Study<V> {
/// Saves the study state to a JSON file.
///
/// # Errors
///
/// Returns an I/O error if the file cannot be created or written.
pub fn save(&self, path: impl AsRef<std::path::Path>) -> std::io::Result<()> {
let snapshot = StudySnapshot {
version: 1,
direction: self.direction,
trials: self.trials(),
next_trial_id: self.next_trial_id.load(Ordering::Relaxed),
metadata: HashMap::new(),
};
let file = std::fs::File::create(path)?;
serde_json::to_writer_pretty(file, &snapshot).map_err(std::io::Error::other)
}
}
#[cfg(feature = "serde")]
impl<V: PartialOrd + Clone + serde::de::DeserializeOwned + 'static> Study<V> {
/// Loads a study from a JSON file.
///
/// The loaded study uses a `RandomSampler` by default. Call
/// [`set_sampler()`](Self::set_sampler) to restore the original sampler
/// configuration.
///
/// # Errors
///
/// Returns an I/O error if the file cannot be read or parsed.
pub fn load(path: impl AsRef<std::path::Path>) -> std::io::Result<Self> {
let file = std::fs::File::open(path)?;
let snapshot: StudySnapshot<V> = serde_json::from_reader(file)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
let study = Study::new(snapshot.direction);
*study.completed_trials.write() = snapshot.trials;
study
.next_trial_id
.store(snapshot.next_trial_id, Ordering::Relaxed);
Ok(study)
}
}
/// Returns `true` if the error represents a pruned trial.
///
/// Checks via `Any` downcasting whether `e` is `Error::TrialPruned` or
/// the standalone `TrialPruned` struct.
fn is_trial_pruned<E: 'static>(e: &E) -> bool {
let any: &dyn Any = e;
if let Some(err) = any.downcast_ref::<crate::Error>() {
matches!(err, crate::Error::TrialPruned)
} else {
any.downcast_ref::<crate::error::TrialPruned>().is_some()
}
}