refactor(tests): split integration.rs into focused subfolders
- Delete 20 duplicate tests already covered by parameter_tests.rs - Move 11 pure Trial unit tests into src/trial.rs - Split remaining 84 integration tests into tests/study/ (9 modules) - Group sampler tests into tests/sampler/ (7 modules) - Group pruner tests into tests/pruner/ (2 modules)
This commit is contained in:
@@ -0,0 +1,98 @@
|
||||
use optimizer::parameter::{FloatParam, Parameter};
|
||||
use optimizer::sampler::tpe::TpeSampler;
|
||||
use optimizer::{Direction, Study};
|
||||
|
||||
#[test]
|
||||
fn test_ask_and_tell_basic() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
for _ in 0..10 {
|
||||
let mut trial = study.ask();
|
||||
let x = x_param.suggest(&mut trial).unwrap();
|
||||
let value = x * x;
|
||||
study.tell(trial, Ok::<_, &str>(value));
|
||||
}
|
||||
|
||||
assert_eq!(study.n_trials(), 10);
|
||||
assert!(study.best_value().unwrap() >= 0.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ask_and_tell_with_failures() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
// Alternate success and failure
|
||||
for i in 0..10 {
|
||||
let mut trial = study.ask();
|
||||
let x = x_param.suggest(&mut trial).unwrap();
|
||||
if i % 2 == 0 {
|
||||
study.tell(trial, Ok::<_, &str>(x * x));
|
||||
} else {
|
||||
study.tell(trial, Err::<f64, _>("simulated failure"));
|
||||
}
|
||||
}
|
||||
|
||||
// Only successful trials are counted
|
||||
assert_eq!(study.n_trials(), 5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ask_and_tell_with_tpe_sampler() {
|
||||
let sampler = TpeSampler::builder()
|
||||
.seed(42)
|
||||
.n_startup_trials(5)
|
||||
.build()
|
||||
.unwrap();
|
||||
let study: Study<f64> = Study::minimize(sampler);
|
||||
let x_param = FloatParam::new(-10.0, 10.0);
|
||||
|
||||
for _ in 0..30 {
|
||||
let mut trial = study.ask();
|
||||
let x = x_param.suggest(&mut trial).unwrap();
|
||||
study.tell(trial, Ok::<_, &str>((x - 3.0).powi(2)));
|
||||
}
|
||||
|
||||
assert_eq!(study.n_trials(), 30);
|
||||
assert!(
|
||||
study.best_value().unwrap() < 5.0,
|
||||
"TPE ask-and-tell should find a reasonable value"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ask_and_tell_batch() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
// Ask a batch of trials
|
||||
let batch: Vec<_> = (0..5)
|
||||
.map(|_| {
|
||||
let mut t = study.ask();
|
||||
let x = x_param.suggest(&mut t).unwrap();
|
||||
(t, x)
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Tell results for the batch
|
||||
for (trial, x) in batch {
|
||||
study.tell(trial, Ok::<_, &str>(x * x));
|
||||
}
|
||||
|
||||
assert_eq!(study.n_trials(), 5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ask_and_tell_with_custom_value_type() {
|
||||
// Ask-and-tell works with non-f64 value types too
|
||||
let study: Study<i32> = Study::new(Direction::Maximize);
|
||||
|
||||
for i in 0..5 {
|
||||
let trial = study.ask();
|
||||
study.tell(trial, Ok::<_, &str>(i * 10));
|
||||
}
|
||||
|
||||
assert_eq!(study.n_trials(), 5);
|
||||
assert_eq!(study.best_value().unwrap(), 40);
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
use optimizer::parameter::{FloatParam, Parameter};
|
||||
use optimizer::sampler::random::RandomSampler;
|
||||
use optimizer::sampler::tpe::TpeSampler;
|
||||
use optimizer::{Direction, Error, Study};
|
||||
|
||||
#[test]
|
||||
fn test_builder_defaults() {
|
||||
let study: Study<f64> = Study::builder().build();
|
||||
assert_eq!(study.direction(), Direction::Minimize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_maximize() {
|
||||
let study: Study<f64> = Study::builder().maximize().build();
|
||||
assert_eq!(study.direction(), Direction::Maximize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_minimize() {
|
||||
let study: Study<f64> = Study::builder().minimize().build();
|
||||
assert_eq!(study.direction(), Direction::Minimize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_direction() {
|
||||
let study: Study<f64> = Study::builder().direction(Direction::Maximize).build();
|
||||
assert_eq!(study.direction(), Direction::Maximize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_with_sampler() {
|
||||
let x = FloatParam::new(-5.0, 5.0);
|
||||
let study: Study<f64> = Study::builder().sampler(TpeSampler::new()).build();
|
||||
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
let val = x.suggest(trial)?;
|
||||
Ok::<_, Error>(val * val)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(study.trials().len(), 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_with_pruner() {
|
||||
use optimizer::pruner::NopPruner;
|
||||
|
||||
let study: Study<f64> = Study::builder().pruner(NopPruner).build();
|
||||
|
||||
assert_eq!(study.direction(), Direction::Minimize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_chaining() {
|
||||
let study: Study<f64> = Study::builder()
|
||||
.maximize()
|
||||
.sampler(RandomSampler::with_seed(42))
|
||||
.pruner(optimizer::pruner::NopPruner)
|
||||
.build();
|
||||
|
||||
assert_eq!(study.direction(), Direction::Maximize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_with_custom_value_type() {
|
||||
let study: Study<i32> = Study::builder().maximize().build();
|
||||
assert_eq!(study.direction(), Direction::Maximize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_builder_optimizes_correctly() {
|
||||
let x = FloatParam::new(-10.0, 10.0);
|
||||
let study: Study<f64> = Study::builder()
|
||||
.minimize()
|
||||
.sampler(TpeSampler::builder().seed(42).build().unwrap())
|
||||
.build();
|
||||
|
||||
study
|
||||
.optimize(100, |trial| {
|
||||
let val = x.suggest(trial)?;
|
||||
Ok::<_, Error>((val - 3.0) * (val - 3.0))
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let best = study.best_trial().unwrap();
|
||||
assert!(
|
||||
best.value < 5.0,
|
||||
"best value should be < 5.0, got {}",
|
||||
best.value
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
use optimizer::{Direction, Study};
|
||||
|
||||
#[test]
|
||||
fn test_is_feasible_all_satisfied() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let mut trial = study.create_trial();
|
||||
trial.set_constraints(vec![-1.0, 0.0, -0.5]);
|
||||
study.complete_trial(trial, 1.0);
|
||||
|
||||
let completed = study.best_trial().unwrap();
|
||||
assert!(completed.is_feasible());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_feasible_one_violated() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let mut trial = study.create_trial();
|
||||
trial.set_constraints(vec![-1.0, 0.5, -0.5]);
|
||||
study.complete_trial(trial, 1.0);
|
||||
|
||||
let completed = study.best_trial().unwrap();
|
||||
assert!(!completed.is_feasible());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_feasible_empty_constraints() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let trial = study.create_trial();
|
||||
study.complete_trial(trial, 1.0);
|
||||
|
||||
let completed = study.best_trial().unwrap();
|
||||
assert!(completed.is_feasible());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_best_trial_prefers_feasible() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
// Infeasible trial with better objective
|
||||
let mut trial1 = study.create_trial();
|
||||
trial1.set_constraints(vec![1.0]);
|
||||
study.complete_trial(trial1, 0.1);
|
||||
|
||||
// Feasible trial with worse objective
|
||||
let mut trial2 = study.create_trial();
|
||||
trial2.set_constraints(vec![-1.0]);
|
||||
study.complete_trial(trial2, 100.0);
|
||||
|
||||
let best = study.best_trial().unwrap();
|
||||
assert_eq!(best.id, 1); // feasible trial wins
|
||||
assert_eq!(best.value, 100.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_best_trial_feasible_by_objective() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
// Feasible, worse objective
|
||||
let mut trial1 = study.create_trial();
|
||||
trial1.set_constraints(vec![-1.0]);
|
||||
study.complete_trial(trial1, 10.0);
|
||||
|
||||
// Feasible, better objective
|
||||
let mut trial2 = study.create_trial();
|
||||
trial2.set_constraints(vec![-0.5]);
|
||||
study.complete_trial(trial2, 2.0);
|
||||
|
||||
let best = study.best_trial().unwrap();
|
||||
assert_eq!(best.id, 1); // lower objective wins among feasible
|
||||
assert_eq!(best.value, 2.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_top_trials_ranks_feasible_above_infeasible() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
// Infeasible, low violation
|
||||
let mut t0 = study.create_trial();
|
||||
t0.set_constraints(vec![0.5]);
|
||||
study.complete_trial(t0, 1.0);
|
||||
|
||||
// Feasible, worst objective among feasible
|
||||
let mut t1 = study.create_trial();
|
||||
t1.set_constraints(vec![-1.0]);
|
||||
study.complete_trial(t1, 50.0);
|
||||
|
||||
// Feasible, best objective among feasible
|
||||
let mut t2 = study.create_trial();
|
||||
t2.set_constraints(vec![-0.1]);
|
||||
study.complete_trial(t2, 5.0);
|
||||
|
||||
// Infeasible, high violation
|
||||
let mut t3 = study.create_trial();
|
||||
t3.set_constraints(vec![3.0]);
|
||||
study.complete_trial(t3, 0.5);
|
||||
|
||||
let top = study.top_trials(4);
|
||||
let ids: Vec<u64> = top.iter().map(|t| t.id).collect();
|
||||
// Feasible sorted by objective first (5.0, 50.0), then infeasible by violation (0.5, 3.0)
|
||||
assert_eq!(ids, vec![2, 1, 0, 3]);
|
||||
}
|
||||
@@ -0,0 +1,189 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use optimizer::parameter::{FloatParam, IntParam, ParamValue, Parameter};
|
||||
use optimizer::sampler::random::RandomSampler;
|
||||
use optimizer::{Direction, Error, Study};
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_params_evaluated_first() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
let y = IntParam::new(1, 100);
|
||||
|
||||
// Enqueue a specific configuration
|
||||
study.enqueue(HashMap::from([
|
||||
(x.id(), ParamValue::Float(5.0)),
|
||||
(y.id(), ParamValue::Int(42)),
|
||||
]));
|
||||
|
||||
// The first trial should use the enqueued params
|
||||
let mut trial = study.ask();
|
||||
let x_val = x.suggest(&mut trial).unwrap();
|
||||
let y_val = y.suggest(&mut trial).unwrap();
|
||||
|
||||
assert_eq!(x_val, 5.0);
|
||||
assert_eq!(y_val, 42);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_fifo_order() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
|
||||
// Enqueue two configs
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(1.0))]));
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(2.0))]));
|
||||
|
||||
// First trial gets first enqueued value
|
||||
let mut trial1 = study.ask();
|
||||
assert_eq!(x.suggest(&mut trial1).unwrap(), 1.0);
|
||||
|
||||
// Second trial gets second enqueued value
|
||||
let mut trial2 = study.ask();
|
||||
assert_eq!(x.suggest(&mut trial2).unwrap(), 2.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_then_normal_sampling_resumes() {
|
||||
let sampler = RandomSampler::with_seed(42);
|
||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
|
||||
// Enqueue one config
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(5.0))]));
|
||||
|
||||
// First trial uses enqueued value
|
||||
let mut trial1 = study.ask();
|
||||
assert_eq!(x.suggest(&mut trial1).unwrap(), 5.0);
|
||||
study.tell(trial1, Ok::<_, &str>(25.0));
|
||||
|
||||
// Second trial uses normal sampling (not 5.0)
|
||||
let mut trial2 = study.ask();
|
||||
let x_val = x.suggest(&mut trial2).unwrap();
|
||||
// The sampled value should be in [0, 10] but extremely unlikely to be exactly 5.0
|
||||
assert!((0.0..=10.0).contains(&x_val));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_with_optimize() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
|
||||
// Enqueue two specific configs
|
||||
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();
|
||||
|
||||
study
|
||||
.optimize(5, |trial| {
|
||||
let x_val = x.suggest(trial)?;
|
||||
values.push(x_val);
|
||||
Ok::<_, Error>(x_val * x_val)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
// First two trials should use enqueued values
|
||||
assert_eq!(values[0], 1.0);
|
||||
assert_eq!(values[1], 2.0);
|
||||
// All 5 trials should have completed
|
||||
assert_eq!(study.n_trials(), 5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_partial_params_fall_back_to_sampling() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
let y = IntParam::new(1, 100);
|
||||
|
||||
// Enqueue only x, not y
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(3.0))]));
|
||||
|
||||
let mut trial = study.ask();
|
||||
let x_val = x.suggest(&mut trial).unwrap();
|
||||
let y_val = y.suggest(&mut trial).unwrap();
|
||||
|
||||
// x should be the enqueued value
|
||||
assert_eq!(x_val, 3.0);
|
||||
// y should be sampled (within range)
|
||||
assert!((1..=100).contains(&y_val));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_trials_appear_in_completed_trials() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(7.0))]));
|
||||
|
||||
study
|
||||
.optimize(1, |trial| {
|
||||
let x_val = x.suggest(trial)?;
|
||||
Ok::<_, Error>(x_val)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let trials = study.trials();
|
||||
assert_eq!(trials.len(), 1);
|
||||
assert_eq!(trials[0].value, 7.0);
|
||||
assert_eq!(
|
||||
*trials[0].params.get(&x.id()).unwrap(),
|
||||
ParamValue::Float(7.0)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_with_ask_and_tell() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(4.0))]));
|
||||
|
||||
let mut trial = study.ask();
|
||||
let x_val = x.suggest(&mut trial).unwrap();
|
||||
assert_eq!(x_val, 4.0);
|
||||
|
||||
study.tell(trial, Ok::<_, &str>(x_val * x_val));
|
||||
assert_eq!(study.n_trials(), 1);
|
||||
assert_eq!(study.best_value().unwrap(), 16.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_n_enqueued() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
|
||||
assert_eq!(study.n_enqueued(), 0);
|
||||
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(1.0))]));
|
||||
assert_eq!(study.n_enqueued(), 1);
|
||||
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(2.0))]));
|
||||
assert_eq!(study.n_enqueued(), 2);
|
||||
|
||||
// Creating a trial dequeues one
|
||||
let _ = study.ask();
|
||||
assert_eq!(study.n_enqueued(), 1);
|
||||
|
||||
let _ = study.ask();
|
||||
assert_eq!(study.n_enqueued(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_enqueue_counted_in_n_trials() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(1.0))]));
|
||||
study.enqueue(HashMap::from([(x.id(), ParamValue::Float(2.0))]));
|
||||
|
||||
study
|
||||
.optimize(5, |trial| {
|
||||
let x_val = x.suggest(trial)?;
|
||||
Ok::<_, Error>(x_val)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
// All 5 trials count, including the 2 enqueued ones
|
||||
assert_eq!(study.n_trials(), 5);
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
use optimizer::parameter::{FloatParam, Parameter};
|
||||
use optimizer::sampler::random::RandomSampler;
|
||||
use optimizer::{Direction, Study};
|
||||
|
||||
#[test]
|
||||
fn test_into_iterator_iterates_all_trials() {
|
||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(42));
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
for _ in 0..5 {
|
||||
let mut trial = study.create_trial();
|
||||
let x = x_param.suggest(&mut trial).unwrap();
|
||||
study.complete_trial(trial, x * x);
|
||||
}
|
||||
|
||||
let mut count = 0;
|
||||
for trial in &study {
|
||||
assert_eq!(trial.state, optimizer::TrialState::Complete);
|
||||
count += 1;
|
||||
}
|
||||
assert_eq!(count, 5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_into_iterator_empty_study() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
let count = (&study).into_iter().count();
|
||||
assert_eq!(count, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_into_iterator_preserves_insertion_order() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
for i in 0..3 {
|
||||
let trial = study.create_trial();
|
||||
study.complete_trial(trial, f64::from(i));
|
||||
}
|
||||
|
||||
let ids: Vec<u64> = (&study).into_iter().map(|t| t.id).collect();
|
||||
assert_eq!(ids, vec![0, 1, 2]);
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
#![allow(
|
||||
clippy::cast_sign_loss,
|
||||
clippy::cast_precision_loss,
|
||||
clippy::cast_possible_truncation
|
||||
)]
|
||||
|
||||
mod ask_tell;
|
||||
mod builder;
|
||||
mod constraints;
|
||||
mod enqueue;
|
||||
mod iterator;
|
||||
mod objective;
|
||||
mod summary;
|
||||
mod top_trials;
|
||||
mod workflow;
|
||||
@@ -0,0 +1,346 @@
|
||||
use optimizer::parameter::{FloatParam, Parameter};
|
||||
use optimizer::sampler::random::RandomSampler;
|
||||
use optimizer::{Direction, Error, Study, Trial};
|
||||
|
||||
#[test]
|
||||
fn test_callback_early_stopping() {
|
||||
use std::ops::ControlFlow;
|
||||
|
||||
use optimizer::Objective;
|
||||
use optimizer::sampler::CompletedTrial;
|
||||
|
||||
struct EarlyStopAfter5 {
|
||||
x_param: FloatParam,
|
||||
}
|
||||
|
||||
impl Objective<f64> for EarlyStopAfter5 {
|
||||
type Error = Error;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, Error> {
|
||||
let x = self.x_param.suggest(trial)?;
|
||||
Ok(x)
|
||||
}
|
||||
fn after_trial(&self, study: &Study<f64>, _trial: &CompletedTrial<f64>) -> ControlFlow<()> {
|
||||
if study.n_trials() >= 5 {
|
||||
ControlFlow::Break(())
|
||||
} else {
|
||||
ControlFlow::Continue(())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
study
|
||||
.optimize_with(
|
||||
100,
|
||||
EarlyStopAfter5 {
|
||||
x_param: FloatParam::new(0.0, 10.0),
|
||||
},
|
||||
)
|
||||
.expect("optimization should succeed");
|
||||
|
||||
assert_eq!(study.n_trials(), 5, "should have stopped after 5 trials");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_callback_early_stopping_on_first_trial() {
|
||||
use std::ops::ControlFlow;
|
||||
|
||||
use optimizer::Objective;
|
||||
use optimizer::sampler::CompletedTrial;
|
||||
|
||||
struct StopImmediately {
|
||||
x_param: FloatParam,
|
||||
}
|
||||
|
||||
impl Objective<f64> for StopImmediately {
|
||||
type Error = Error;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, Error> {
|
||||
let x = self.x_param.suggest(trial)?;
|
||||
Ok(x)
|
||||
}
|
||||
fn after_trial(
|
||||
&self,
|
||||
_study: &Study<f64>,
|
||||
_trial: &CompletedTrial<f64>,
|
||||
) -> ControlFlow<()> {
|
||||
ControlFlow::Break(())
|
||||
}
|
||||
}
|
||||
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
study
|
||||
.optimize_with(
|
||||
100,
|
||||
StopImmediately {
|
||||
x_param: FloatParam::new(0.0, 10.0),
|
||||
},
|
||||
)
|
||||
.expect("optimization should succeed");
|
||||
|
||||
assert_eq!(study.n_trials(), 1, "should have stopped after 1 trial");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_callback_sampler_early_stopping() {
|
||||
use std::ops::ControlFlow;
|
||||
|
||||
use optimizer::Objective;
|
||||
use optimizer::sampler::CompletedTrial;
|
||||
|
||||
struct StopAfter3 {
|
||||
x_param: FloatParam,
|
||||
}
|
||||
|
||||
impl Objective<f64> for StopAfter3 {
|
||||
type Error = Error;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, Error> {
|
||||
let x = self.x_param.suggest(trial)?;
|
||||
Ok(x)
|
||||
}
|
||||
fn after_trial(&self, study: &Study<f64>, _trial: &CompletedTrial<f64>) -> ControlFlow<()> {
|
||||
if study.n_trials() >= 3 {
|
||||
ControlFlow::Break(())
|
||||
} else {
|
||||
ControlFlow::Continue(())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let sampler = RandomSampler::with_seed(42);
|
||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
|
||||
study
|
||||
.optimize_with(
|
||||
100,
|
||||
StopAfter3 {
|
||||
x_param: FloatParam::new(0.0, 10.0),
|
||||
},
|
||||
)
|
||||
.expect("optimization should succeed");
|
||||
|
||||
assert_eq!(study.n_trials(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retries_successful_trials_not_retried() {
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicU32, Ordering};
|
||||
|
||||
use optimizer::Objective;
|
||||
|
||||
struct SuccessObj {
|
||||
x_param: FloatParam,
|
||||
call_count: Arc<AtomicU32>,
|
||||
}
|
||||
|
||||
impl Objective<f64> for SuccessObj {
|
||||
type Error = Error;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, Error> {
|
||||
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<f64> = Study::new(Direction::Minimize);
|
||||
let call_count = Arc::new(AtomicU32::new(0));
|
||||
let obj = SuccessObj {
|
||||
x_param: FloatParam::new(0.0, 10.0),
|
||||
call_count: Arc::clone(&call_count),
|
||||
};
|
||||
|
||||
study.optimize_with(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<AtomicU32>,
|
||||
}
|
||||
|
||||
impl Objective<f64> for AlwaysFailObj {
|
||||
type Error = String;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, String> {
|
||||
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<f64> = 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<f64> for AlwaysFailObj {
|
||||
type Error = String;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, String> {
|
||||
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<f64> = 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<Mutex<Vec<f64>>>,
|
||||
call_count: Arc<AtomicU32>,
|
||||
}
|
||||
|
||||
impl Objective<f64> for RetryObj {
|
||||
type Error = String;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, String> {
|
||||
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<f64> = 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<AtomicU32>,
|
||||
}
|
||||
|
||||
impl Objective<f64> for FailFirstObj {
|
||||
type Error = String;
|
||||
fn evaluate(&self, trial: &mut Trial) -> Result<f64, String> {
|
||||
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<f64> = 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<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
let call_count = std::cell::Cell::new(0u32);
|
||||
|
||||
study
|
||||
.optimize(5, |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);
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
use optimizer::parameter::{FloatParam, Parameter};
|
||||
use optimizer::sampler::random::RandomSampler;
|
||||
use optimizer::{Direction, Error, Study};
|
||||
|
||||
#[test]
|
||||
fn test_summary_with_completed_trials() {
|
||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(1));
|
||||
let x = FloatParam::new(0.0, 10.0).name("x");
|
||||
|
||||
study
|
||||
.optimize(5, |trial| {
|
||||
let val = x.suggest(trial)?;
|
||||
Ok::<_, Error>(val * val)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let summary = study.summary();
|
||||
assert!(summary.contains("Minimize"));
|
||||
assert!(summary.contains("5 trials"));
|
||||
assert!(summary.contains("Best value:"));
|
||||
assert!(summary.contains("x = "));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_summary_no_completed_trials() {
|
||||
let study: Study<f64> = Study::new(Direction::Maximize);
|
||||
let summary = study.summary();
|
||||
assert!(summary.contains("Maximize"));
|
||||
assert!(summary.contains("0 trials"));
|
||||
assert!(!summary.contains("Best value:"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_summary_with_pruned_trials() {
|
||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(1));
|
||||
let x = FloatParam::new(0.0, 10.0).name("x");
|
||||
|
||||
// Manually create some complete and pruned trials
|
||||
for _ in 0..3 {
|
||||
let mut trial = study.create_trial();
|
||||
let val = x.suggest(&mut trial).unwrap();
|
||||
study.complete_trial(trial, val);
|
||||
}
|
||||
for _ in 0..2 {
|
||||
let mut trial = study.create_trial();
|
||||
let _ = x.suggest(&mut trial).unwrap();
|
||||
study.prune_trial(trial);
|
||||
}
|
||||
|
||||
let summary = study.summary();
|
||||
// Should show breakdown when there are pruned trials
|
||||
if study.n_pruned_trials() > 0 {
|
||||
assert!(summary.contains("complete"));
|
||||
assert!(summary.contains("pruned"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_display_matches_summary() {
|
||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(1));
|
||||
let x = FloatParam::new(0.0, 10.0).name("x");
|
||||
|
||||
study
|
||||
.optimize(3, |trial| {
|
||||
let val = x.suggest(trial)?;
|
||||
Ok::<_, Error>(val)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(format!("{study}"), study.summary());
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
use optimizer::{Direction, Study};
|
||||
|
||||
#[test]
|
||||
fn test_top_trials_minimize() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
// Manually complete trials with known values
|
||||
for &val in &[5.0, 1.0, 3.0, 2.0, 4.0] {
|
||||
let trial = study.create_trial();
|
||||
study.complete_trial(trial, val);
|
||||
}
|
||||
|
||||
let top3 = study.top_trials(3);
|
||||
assert_eq!(top3.len(), 3);
|
||||
assert_eq!(top3[0].value, 1.0);
|
||||
assert_eq!(top3[1].value, 2.0);
|
||||
assert_eq!(top3[2].value, 3.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_top_trials_maximize() {
|
||||
let study: Study<f64> = Study::new(Direction::Maximize);
|
||||
|
||||
for &val in &[5.0, 1.0, 3.0, 2.0, 4.0] {
|
||||
let trial = study.create_trial();
|
||||
study.complete_trial(trial, val);
|
||||
}
|
||||
|
||||
let top3 = study.top_trials(3);
|
||||
assert_eq!(top3.len(), 3);
|
||||
assert_eq!(top3[0].value, 5.0);
|
||||
assert_eq!(top3[1].value, 4.0);
|
||||
assert_eq!(top3[2].value, 3.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_top_trials_n_greater_than_total() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
for &val in &[3.0, 1.0] {
|
||||
let trial = study.create_trial();
|
||||
study.complete_trial(trial, val);
|
||||
}
|
||||
|
||||
let top = study.top_trials(10);
|
||||
assert_eq!(top.len(), 2);
|
||||
assert_eq!(top[0].value, 1.0);
|
||||
assert_eq!(top[1].value, 3.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_top_trials_empty() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
let top = study.top_trials(5);
|
||||
assert!(top.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_top_trials_excludes_pruned() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
// Complete some trials
|
||||
for &val in &[5.0, 1.0, 3.0] {
|
||||
let trial = study.create_trial();
|
||||
study.complete_trial(trial, val);
|
||||
}
|
||||
|
||||
// Prune a trial (it gets a default value of 0.0 but should be excluded)
|
||||
let trial = study.create_trial();
|
||||
study.prune_trial(trial);
|
||||
|
||||
let top = study.top_trials(5);
|
||||
assert_eq!(top.len(), 3, "pruned trial should be excluded");
|
||||
assert_eq!(top[0].value, 1.0);
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
use optimizer::parameter::{BoolParam, FloatParam, IntParam, Parameter};
|
||||
use optimizer::sampler::tpe::TpeSampler;
|
||||
use optimizer::{Direction, Error, Study};
|
||||
|
||||
#[test]
|
||||
fn test_study_basic_workflow() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x)
|
||||
})
|
||||
.expect("optimization should succeed");
|
||||
|
||||
assert_eq!(study.n_trials(), 10);
|
||||
let best = study.best_trial().expect("should have best trial");
|
||||
assert!(best.value >= 0.0, "x^2 should be non-negative");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_study_with_failures() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
// Every other trial fails
|
||||
let mut counter = 0;
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
counter += 1;
|
||||
if counter % 2 == 0 {
|
||||
return Err::<f64, &str>("intentional failure");
|
||||
}
|
||||
let x = x_param.suggest(trial).map_err(|_| "param error")?;
|
||||
Ok(x * x)
|
||||
})
|
||||
.expect("optimization should succeed with some failures");
|
||||
|
||||
// Only half the trials should have succeeded
|
||||
assert_eq!(study.n_trials(), 5, "only 5 trials should have completed");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_no_completed_trials_error() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
let result = study.best_trial();
|
||||
assert!(matches!(result, Err(Error::NoCompletedTrials)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_study_direction() {
|
||||
let study_min: Study<f64> = Study::new(Direction::Minimize);
|
||||
assert_eq!(study_min.direction(), Direction::Minimize);
|
||||
|
||||
let study_max: Study<f64> = Study::new(Direction::Maximize);
|
||||
assert_eq!(study_max.direction(), Direction::Maximize);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_study_trials_iteration() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(0.0, 1.0);
|
||||
|
||||
study
|
||||
.optimize(5, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let trials = study.trials();
|
||||
assert_eq!(trials.len(), 5);
|
||||
|
||||
for trial in &trials {
|
||||
assert!(
|
||||
!trial.params.is_empty(),
|
||||
"each trial should have parameters"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_study_set_sampler() {
|
||||
let mut study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
let tpe = TpeSampler::builder()
|
||||
.seed(42)
|
||||
.n_startup_trials(5)
|
||||
.build()
|
||||
.unwrap();
|
||||
study.set_sampler(tpe);
|
||||
|
||||
let x_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x)
|
||||
})
|
||||
.expect("optimization should succeed with new sampler");
|
||||
|
||||
assert_eq!(study.n_trials(), 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_study_with_i32_value_type() {
|
||||
let study: Study<i32> = Study::new(Direction::Minimize);
|
||||
let x_param = IntParam::new(-10, 10);
|
||||
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x.abs() as i32)
|
||||
})
|
||||
.expect("optimization should succeed");
|
||||
|
||||
assert_eq!(study.n_trials(), 10);
|
||||
let best = study.best_trial().expect("should have best trial");
|
||||
assert!(best.value >= 0, "absolute value should be non-negative");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_optimize_all_trials_fail() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
let result = study.optimize(5, |_trial| Err::<f64, &str>("always fails"));
|
||||
|
||||
assert!(
|
||||
matches!(result, Err(Error::NoCompletedTrials)),
|
||||
"should return NoCompletedTrials when all trials fail"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_best_value() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let best_value = study.best_value().expect("should have best value");
|
||||
let best_trial = study.best_trial().expect("should have best trial");
|
||||
|
||||
assert_eq!(
|
||||
best_value, best_trial.value,
|
||||
"best_value should match best_trial.value"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_best_trial_with_nan_values() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize(5, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let best = study.best_trial();
|
||||
assert!(best.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_manual_trial_completion() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
// Manually create and complete trials
|
||||
let mut trial = study.create_trial();
|
||||
let x = x_param.suggest(&mut trial).unwrap();
|
||||
study.complete_trial(trial, x * x);
|
||||
|
||||
let mut trial2 = study.create_trial();
|
||||
let y = x_param.suggest(&mut trial2).unwrap();
|
||||
study.complete_trial(trial2, y * y);
|
||||
|
||||
// Manually fail a trial
|
||||
let trial3 = study.create_trial();
|
||||
study.fail_trial(trial3, "test failure");
|
||||
|
||||
// Only 2 completed trials
|
||||
assert_eq!(study.n_trials(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_multiple_params_in_optimization() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(-10.0, 10.0);
|
||||
let n_param = IntParam::new(1, 5);
|
||||
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let n = n_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x + n as f64)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(study.n_trials(), 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_suggest_bool_in_optimization() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let use_feature_param = BoolParam::new();
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize(10, |trial| {
|
||||
let use_feature = use_feature_param.suggest(trial)?;
|
||||
let x = x_param.suggest(trial)?;
|
||||
|
||||
let value = if use_feature { x } else { x * 2.0 };
|
||||
Ok::<_, Error>(value)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(study.n_trials(), 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_completed_trial_get() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(-10.0, 10.0).name("x");
|
||||
let n_param = IntParam::new(1, 10).name("n");
|
||||
|
||||
study
|
||||
.optimize(5, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let n = n_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x + n as f64)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let best = study.best_trial().unwrap();
|
||||
let x_val: f64 = best.get(&x_param).unwrap();
|
||||
let n_val: i64 = best.get(&n_param).unwrap();
|
||||
assert!((-10.0..=10.0).contains(&x_val));
|
||||
assert!((1..=10).contains(&n_val));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_single_value_int_range() {
|
||||
let param = IntParam::new(5, 5);
|
||||
let mut trial = optimizer::Trial::new(0);
|
||||
|
||||
let n = param.suggest(&mut trial).unwrap();
|
||||
assert_eq!(n, 5, "single-value range should return that value");
|
||||
}
|
||||
Reference in New Issue
Block a user