use crate::{
Direction, HyperoptError, ObjectiveError, ObjectiveResult, Pruner, Sampler, Storage,
StudyMetadata, StudyState, Trial, TrialContext, TrialState,
};
use std::panic::{catch_unwind, AssertUnwindSafe};
use std::sync::Mutex;
pub struct Study {
name: String,
direction: Direction,
sampler: Mutex<Box<dyn Sampler>>,
pruner: Box<dyn Pruner>,
storage: Box<dyn Storage>,
}
impl Study {
pub fn new(
name: impl Into<String>,
direction: Direction,
sampler: Box<dyn Sampler>,
pruner: Box<dyn Pruner>,
storage: Box<dyn Storage>,
) -> Result<Self, HyperoptError> {
let name = name.into();
let direction = match storage.load_study_metadata(&name)? {
Some(meta) => meta.direction,
None => {
storage.save_study_metadata(&StudyMetadata {
study_name: name.clone(),
direction,
})?;
direction
}
};
Ok(Study {
name,
direction,
sampler: Mutex::new(sampler),
pruner,
storage,
})
}
pub fn name(&self) -> &str {
&self.name
}
pub fn direction(&self) -> Direction {
self.direction
}
pub fn trials(&self) -> Result<Vec<Trial>, HyperoptError> {
Ok(self.storage.load_trials(&self.name)?)
}
pub fn best_trial(&self) -> Result<Option<Trial>, HyperoptError> {
let trials = self.storage.load_trials(&self.name)?;
let state = StudyState::new(self.direction, trials);
Ok(state.best_trial().cloned())
}
pub fn best_value(&self) -> Result<Option<f64>, HyperoptError> {
Ok(self.best_trial()?.and_then(|t| t.value))
}
pub fn optimize<F>(&self, mut objective: F, n_trials: usize) -> Result<(), HyperoptError>
where
F: FnMut(&mut TrialContext) -> ObjectiveResult,
{
for _ in 0..n_trials {
let existing = self.storage.load_trials(&self.name)?;
let number = existing.len();
let state = StudyState::new(self.direction, existing);
let mut ctx =
TrialContext::new(number, &self.sampler, &state, self.pruner.as_ref());
let result = catch_unwind(AssertUnwindSafe(|| objective(&mut ctx)));
let trial = finish_trial(ctx.into_trial(), result);
self.storage.save_trial(&self.name, &trial)?;
}
Ok(())
}
#[cfg(feature = "parallel")]
pub fn optimize_parallel<F>(
&self,
objective: F,
n_trials: usize,
n_workers: usize,
) -> Result<(), HyperoptError>
where
F: Fn(&mut TrialContext) -> ObjectiveResult + Sync,
{
use rayon::prelude::*;
use std::sync::atomic::{AtomicUsize, Ordering};
let base = self.storage.load_trials(&self.name)?.len();
let counter = AtomicUsize::new(base);
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(n_workers)
.build()
.map_err(|e| HyperoptError::Storage(crate::StorageError::Backend(e.to_string())))?;
pool.install(|| {
(0..n_trials).into_par_iter().try_for_each(|_| {
let number = counter.fetch_add(1, Ordering::SeqCst);
let existing = self.storage.load_trials(&self.name)?;
let state = StudyState::new(self.direction, existing);
let mut ctx =
TrialContext::new(number, &self.sampler, &state, self.pruner.as_ref());
let result = catch_unwind(AssertUnwindSafe(|| objective(&mut ctx)));
let trial = finish_trial(ctx.into_trial(), result);
self.storage.save_trial(&self.name, &trial)?;
Ok::<(), HyperoptError>(())
})
})
}
}
fn finish_trial(
mut trial: Trial,
result: std::thread::Result<ObjectiveResult>,
) -> Trial {
match result {
Ok(Ok(value)) => {
trial.value = Some(value);
trial.state = TrialState::Complete;
}
Ok(Err(ObjectiveError::Pruned)) => {
trial.state = TrialState::Pruned;
}
Ok(Err(ObjectiveError::Failed(_))) => {
trial.state = TrialState::Failed;
}
Err(_panic) => {
trial.state = TrialState::Failed;
}
}
trial
}